mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: args issue & refactor into helper function to reduce bloat for both(#19441)
This commit is contained in:
parent
e41b9c29a8
commit
6df0406cf6
4 changed files with 218 additions and 102 deletions
|
|
@ -550,9 +550,9 @@ except ImportError:
|
|||
server_root_path = get_server_root_path()
|
||||
_license_check = LicenseCheck()
|
||||
premium_user: bool = _license_check.is_premium()
|
||||
premium_user_data: Optional["EnterpriseLicenseData"] = (
|
||||
_license_check.airgapped_license_data
|
||||
)
|
||||
premium_user_data: Optional[
|
||||
"EnterpriseLicenseData"
|
||||
] = _license_check.airgapped_license_data
|
||||
global_max_parallel_request_retries_env: Optional[str] = os.getenv(
|
||||
"LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES"
|
||||
)
|
||||
|
|
@ -1209,9 +1209,9 @@ master_key: Optional[str] = None
|
|||
config_agents: Optional[List[AgentConfig]] = None
|
||||
otel_logging = False
|
||||
prisma_client: Optional[PrismaClient] = None
|
||||
shared_aiohttp_session: Optional["ClientSession"] = (
|
||||
None # Global shared session for connection reuse
|
||||
)
|
||||
shared_aiohttp_session: Optional[
|
||||
"ClientSession"
|
||||
] = None # Global shared session for connection reuse
|
||||
user_api_key_cache = DualCache(
|
||||
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
|
||||
)
|
||||
|
|
@ -1219,9 +1219,9 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(
|
|||
dual_cache=user_api_key_cache
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
|
||||
redis_usage_cache: Optional[RedisCache] = (
|
||||
None # redis cache used for tracking spend, tpm/rpm limits
|
||||
)
|
||||
redis_usage_cache: Optional[
|
||||
RedisCache
|
||||
] = None # redis cache used for tracking spend, tpm/rpm limits
|
||||
polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False
|
||||
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
|
||||
user_custom_auth = None
|
||||
|
|
@ -1560,9 +1560,9 @@ async def update_cache( # noqa: PLR0915
|
|||
_id = "team_id:{}".format(team_id)
|
||||
try:
|
||||
# Fetch the existing cost for the given user
|
||||
existing_spend_obj: Optional[LiteLLM_TeamTable] = (
|
||||
await user_api_key_cache.async_get_cache(key=_id)
|
||||
)
|
||||
existing_spend_obj: Optional[
|
||||
LiteLLM_TeamTable
|
||||
] = await user_api_key_cache.async_get_cache(key=_id)
|
||||
if existing_spend_obj is None:
|
||||
# do nothing if team not in api key cache
|
||||
return
|
||||
|
|
@ -2726,15 +2726,15 @@ class ProxyConfig:
|
|||
router_settings = config.get("router_settings", None)
|
||||
|
||||
if router_settings and isinstance(router_settings, dict):
|
||||
arg_spec = inspect.getfullargspec(litellm.Router)
|
||||
# model list and search_tools already set
|
||||
exclude_args = {
|
||||
"self",
|
||||
"model_list",
|
||||
"search_tools",
|
||||
}
|
||||
|
||||
available_args = [x for x in arg_spec.args if x not in exclude_args]
|
||||
available_args = [
|
||||
x for x in litellm.Router.get_valid_args() if x not in exclude_args
|
||||
]
|
||||
|
||||
for k, v in router_settings.items():
|
||||
if k in available_args:
|
||||
|
|
@ -2856,6 +2856,7 @@ class ProxyConfig:
|
|||
|
||||
from litellm.proxy.policy_engine.init_policies import init_policies
|
||||
from litellm.proxy.policy_engine.policy_validator import PolicyValidator
|
||||
|
||||
if config is None:
|
||||
verbose_proxy_logger.debug("Policy engine: config is None, skipping")
|
||||
return
|
||||
|
|
@ -2867,7 +2868,9 @@ class ProxyConfig:
|
|||
|
||||
policy_attachments_config = config.get("policy_attachments", None)
|
||||
|
||||
verbose_proxy_logger.info(f"Policy engine: found {len(policies_config)} policies in config")
|
||||
verbose_proxy_logger.info(
|
||||
f"Policy engine: found {len(policies_config)} policies in config"
|
||||
)
|
||||
|
||||
# Initialize policies
|
||||
await init_policies(
|
||||
|
|
@ -4009,10 +4012,10 @@ class ProxyConfig:
|
|||
)
|
||||
|
||||
try:
|
||||
guardrails_in_db: List[Guardrail] = (
|
||||
await GuardrailRegistry.get_all_guardrails_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
guardrails_in_db: List[
|
||||
Guardrail
|
||||
] = await GuardrailRegistry.get_all_guardrails_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"guardrails from the DB %s", str(guardrails_in_db)
|
||||
|
|
@ -4046,7 +4049,9 @@ class ProxyConfig:
|
|||
await policy_registry.sync_policies_from_db(prisma_client=prisma_client)
|
||||
|
||||
# Sync attachments from DB to in-memory registry
|
||||
await attachment_registry.sync_attachments_from_db(prisma_client=prisma_client)
|
||||
await attachment_registry.sync_attachments_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Successfully synced policies and attachments from DB"
|
||||
|
|
@ -4369,9 +4374,9 @@ async def initialize( # noqa: PLR0915
|
|||
user_api_base = api_base
|
||||
dynamic_config[user_model]["api_base"] = api_base
|
||||
if api_version:
|
||||
os.environ["AZURE_API_VERSION"] = (
|
||||
api_version # set this for azure - litellm can read this from the env
|
||||
)
|
||||
os.environ[
|
||||
"AZURE_API_VERSION"
|
||||
] = api_version # set this for azure - litellm can read this from the env
|
||||
if max_tokens: # model-specific param
|
||||
dynamic_config[user_model]["max_tokens"] = max_tokens
|
||||
if temperature: # model-specific param
|
||||
|
|
@ -5217,7 +5222,9 @@ async def model_list(
|
|||
|
||||
# Include model access groups if requested
|
||||
if include_model_access_groups:
|
||||
proxy_model_list = list(set(proxy_model_list + list(model_access_groups.keys())))
|
||||
proxy_model_list = list(
|
||||
set(proxy_model_list + list(model_access_groups.keys()))
|
||||
)
|
||||
|
||||
# Get complete model list including wildcard routes if requested
|
||||
from litellm.proxy.auth.model_checks import get_complete_model_list
|
||||
|
|
@ -7674,12 +7681,12 @@ def _enrich_model_info_with_litellm_data(
|
|||
"""
|
||||
Enrich a model dictionary with litellm model info (pricing, context window, etc.)
|
||||
and remove sensitive information.
|
||||
|
||||
|
||||
Args:
|
||||
model: Model dictionary to enrich
|
||||
debug: Whether to include debug information like openai_client
|
||||
llm_router: Optional router instance for debug info
|
||||
|
||||
|
||||
Returns:
|
||||
Enriched model dictionary with sensitive info removed
|
||||
"""
|
||||
|
|
@ -7689,9 +7696,7 @@ def _enrich_model_info_with_litellm_data(
|
|||
_openai_client = "None"
|
||||
if llm_router is not None:
|
||||
_openai_client = (
|
||||
llm_router._get_client(
|
||||
deployment=model, kwargs={}, client_type="async"
|
||||
)
|
||||
llm_router._get_client(deployment=model, kwargs={}, client_type="async")
|
||||
or "None"
|
||||
)
|
||||
else:
|
||||
|
|
@ -7749,7 +7754,7 @@ async def _apply_search_filter_to_models(
|
|||
) -> Tuple[List[Dict[str, Any]], Optional[int]]:
|
||||
"""
|
||||
Apply search filter to models, querying database for additional matching models.
|
||||
|
||||
|
||||
Args:
|
||||
all_models: List of models to filter
|
||||
search: Search term (case-insensitive)
|
||||
|
|
@ -7757,44 +7762,43 @@ async def _apply_search_filter_to_models(
|
|||
size: Page size
|
||||
prisma_client: Prisma client for database queries
|
||||
proxy_config: Proxy config for decrypting models
|
||||
|
||||
|
||||
Returns:
|
||||
Tuple of (filtered_models, total_count). total_count is None if not searching.
|
||||
"""
|
||||
if not search or not search.strip():
|
||||
return all_models, None
|
||||
|
||||
|
||||
search_lower = search.lower().strip()
|
||||
|
||||
|
||||
# Filter models in router by search term
|
||||
filtered_router_models = [
|
||||
m for m in all_models
|
||||
if search_lower in m.get("model_name", "").lower()
|
||||
m for m in all_models if search_lower in m.get("model_name", "").lower()
|
||||
]
|
||||
|
||||
|
||||
# Separate filtered models into config vs db models, and track db model IDs
|
||||
filtered_config_models = []
|
||||
db_model_ids_in_router = set()
|
||||
|
||||
|
||||
for m in filtered_router_models:
|
||||
model_info = m.get("model_info", {})
|
||||
is_db_model = model_info.get("db_model", False)
|
||||
model_id = model_info.get("id")
|
||||
|
||||
|
||||
if is_db_model and model_id:
|
||||
db_model_ids_in_router.add(model_id)
|
||||
else:
|
||||
filtered_config_models.append(m)
|
||||
|
||||
|
||||
config_models_count = len(filtered_config_models)
|
||||
db_models_in_router_count = len(db_model_ids_in_router)
|
||||
router_models_count = config_models_count + db_models_in_router_count
|
||||
|
||||
|
||||
# Query database for additional models with search term
|
||||
db_models = []
|
||||
db_models_total_count = 0
|
||||
models_needed_for_page = size * page
|
||||
|
||||
|
||||
# Only query database if prisma_client is available
|
||||
if prisma_client is not None:
|
||||
try:
|
||||
|
|
@ -7810,31 +7814,36 @@ async def _apply_search_filter_to_models(
|
|||
db_where_condition["model_id"] = {
|
||||
"not": {"in": list(db_model_ids_in_router)}
|
||||
}
|
||||
|
||||
|
||||
# Get total count of matching database models
|
||||
db_models_total_count = await prisma_client.db.litellm_proxymodeltable.count(
|
||||
where=db_where_condition
|
||||
db_models_total_count = (
|
||||
await prisma_client.db.litellm_proxymodeltable.count(
|
||||
where=db_where_condition
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# Calculate total count for search results
|
||||
search_total_count = router_models_count + db_models_total_count
|
||||
|
||||
|
||||
# Fetch database models if we need more for the current page
|
||||
if router_models_count < models_needed_for_page:
|
||||
models_to_fetch = min(
|
||||
models_needed_for_page - router_models_count,
|
||||
db_models_total_count
|
||||
models_needed_for_page - router_models_count, db_models_total_count
|
||||
)
|
||||
|
||||
|
||||
if models_to_fetch > 0:
|
||||
db_models_raw = await prisma_client.db.litellm_proxymodeltable.find_many(
|
||||
where=db_where_condition,
|
||||
take=models_to_fetch,
|
||||
db_models_raw = (
|
||||
await prisma_client.db.litellm_proxymodeltable.find_many(
|
||||
where=db_where_condition,
|
||||
take=models_to_fetch,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# Convert database models to router format
|
||||
for db_model in db_models_raw:
|
||||
decrypted_models = proxy_config.decrypt_model_list_from_db([db_model])
|
||||
decrypted_models = proxy_config.decrypt_model_list_from_db(
|
||||
[db_model]
|
||||
)
|
||||
if decrypted_models:
|
||||
db_models.extend(decrypted_models)
|
||||
except Exception as e:
|
||||
|
|
@ -7846,7 +7855,7 @@ async def _apply_search_filter_to_models(
|
|||
else:
|
||||
# If no prisma_client, only use router models
|
||||
search_total_count = router_models_count
|
||||
|
||||
|
||||
# Combine all models
|
||||
filtered_models = filtered_router_models + db_models
|
||||
return filtered_models, search_total_count
|
||||
|
|
@ -7861,28 +7870,28 @@ def _paginate_models_response(
|
|||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Paginate models and return response dictionary.
|
||||
|
||||
|
||||
Args:
|
||||
all_models: List of all models
|
||||
page: Current page number
|
||||
size: Page size
|
||||
total_count: Total count (if None, uses len(all_models))
|
||||
search: Search term (for logging)
|
||||
|
||||
|
||||
Returns:
|
||||
Paginated response dictionary
|
||||
"""
|
||||
if total_count is None:
|
||||
total_count = len(all_models)
|
||||
|
||||
|
||||
skip = (page - 1) * size
|
||||
total_pages = -(-total_count // size) if total_count > 0 else 0
|
||||
paginated_models = all_models[skip : skip + size]
|
||||
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Pagination: skip={skip}, take={size}, total_count={total_count}, total_pages={total_pages}, search={search}"
|
||||
)
|
||||
|
||||
|
||||
return {
|
||||
"data": paginated_models,
|
||||
"total_count": total_count,
|
||||
|
|
@ -7902,15 +7911,15 @@ async def _filter_models_by_team_id(
|
|||
Filter models by team ID. Returns models where:
|
||||
- direct_access is True, OR
|
||||
- team_id is in access_via_team_ids
|
||||
|
||||
|
||||
Also searches config and database for models accessible to the team.
|
||||
|
||||
|
||||
Args:
|
||||
all_models: List of models to filter
|
||||
team_id: Team ID to filter by
|
||||
prisma_client: Prisma client for database queries
|
||||
llm_router: Router instance for config queries
|
||||
|
||||
|
||||
Returns:
|
||||
Filtered list of models
|
||||
"""
|
||||
|
|
@ -7923,15 +7932,15 @@ async def _filter_models_by_team_id(
|
|||
verbose_proxy_logger.warning(f"Team {team_id} not found in database")
|
||||
# If team doesn't exist, return empty list
|
||||
return []
|
||||
|
||||
|
||||
team_object = LiteLLM_TeamTable(**team_db_object.model_dump())
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}")
|
||||
return []
|
||||
|
||||
|
||||
# Get models accessible to this team (similar to _add_team_models_to_all_models)
|
||||
team_accessible_model_ids: Set[str] = set()
|
||||
|
||||
|
||||
if (
|
||||
len(team_object.models) == 0 # empty list = all model access
|
||||
or SpecialModelNames.all_proxy_models.value in team_object.models
|
||||
|
|
@ -7950,25 +7959,30 @@ async def _filter_models_by_team_id(
|
|||
can_add_model = True
|
||||
elif team_model_id == team_id:
|
||||
can_add_model = True
|
||||
|
||||
|
||||
if can_add_model:
|
||||
team_accessible_model_ids.add(model_id)
|
||||
else:
|
||||
# Team has access to specific models
|
||||
for model_name in team_object.models:
|
||||
_models = llm_router.get_model_list(
|
||||
model_name=model_name, team_id=team_id
|
||||
) if llm_router else []
|
||||
_models = (
|
||||
llm_router.get_model_list(model_name=model_name, team_id=team_id)
|
||||
if llm_router
|
||||
else []
|
||||
)
|
||||
if _models is not None:
|
||||
for model in _models:
|
||||
model_id = model.get("model_info", {}).get("id", None)
|
||||
if model_id is not None:
|
||||
team_accessible_model_ids.add(model_id)
|
||||
|
||||
|
||||
# Also search database for models accessible to this team
|
||||
# This complements the config search done above
|
||||
try:
|
||||
if team_object.models and SpecialModelNames.all_proxy_models.value not in team_object.models:
|
||||
if (
|
||||
team_object.models
|
||||
and SpecialModelNames.all_proxy_models.value not in team_object.models
|
||||
):
|
||||
# Team has specific models - check database for those model names
|
||||
db_models = await prisma_client.db.litellm_proxymodeltable.find_many(
|
||||
where={"model_name": {"in": team_object.models}}
|
||||
|
|
@ -7978,31 +7992,33 @@ async def _filter_models_by_team_id(
|
|||
if model_id:
|
||||
team_accessible_model_ids.add(model_id)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(f"Error querying database models for team {team_id}: {str(e)}")
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Error querying database models for team {team_id}: {str(e)}"
|
||||
)
|
||||
|
||||
# Filter models based on direct_access or access_via_team_ids
|
||||
# Models are already enriched with these fields before this function is called
|
||||
filtered_models = []
|
||||
for _model in all_models:
|
||||
model_info = _model.get("model_info", {})
|
||||
model_id = model_info.get("id", None)
|
||||
|
||||
|
||||
# Include if direct_access is True
|
||||
if model_info.get("direct_access", False):
|
||||
filtered_models.append(_model)
|
||||
continue
|
||||
|
||||
|
||||
# Include if team_id is in access_via_team_ids
|
||||
access_via_team_ids = model_info.get("access_via_team_ids", [])
|
||||
if isinstance(access_via_team_ids, list) and team_id in access_via_team_ids:
|
||||
filtered_models.append(_model)
|
||||
continue
|
||||
|
||||
|
||||
# Also include if model_id is in team_accessible_model_ids (from config/db search)
|
||||
# This catches models that might not have been enriched with access_via_team_ids yet
|
||||
if model_id and model_id in team_accessible_model_ids:
|
||||
filtered_models.append(_model)
|
||||
|
||||
|
||||
return filtered_models
|
||||
|
||||
|
||||
|
|
@ -8034,7 +8050,8 @@ async def model_info_v2(
|
|||
None, description="Search for a specific model by its unique ID"
|
||||
),
|
||||
teamId: Optional[str] = fastapi.Query(
|
||||
None, description="Filter models by team ID. Returns models with direct_access=True or teamId in access_via_team_ids"
|
||||
None,
|
||||
description="Filter models by team ID. Returns models with direct_access=True or teamId in access_via_team_ids",
|
||||
),
|
||||
):
|
||||
"""
|
||||
|
|
@ -8064,13 +8081,13 @@ async def model_info_v2(
|
|||
# If modelId is provided, search for the specific model
|
||||
if modelId is not None:
|
||||
found_model = None
|
||||
|
||||
|
||||
# First, search in config
|
||||
if llm_router is not None:
|
||||
found_model = llm_router.get_model_info(id=modelId)
|
||||
if found_model:
|
||||
found_model = copy.deepcopy(found_model)
|
||||
|
||||
|
||||
# If not found in config, search in database
|
||||
if found_model is None:
|
||||
try:
|
||||
|
|
@ -8079,14 +8096,16 @@ async def model_info_v2(
|
|||
)
|
||||
if db_model:
|
||||
# Convert database model to router format
|
||||
decrypted_models = proxy_config.decrypt_model_list_from_db([db_model])
|
||||
decrypted_models = proxy_config.decrypt_model_list_from_db(
|
||||
[db_model]
|
||||
)
|
||||
if decrypted_models:
|
||||
found_model = decrypted_models[0]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
f"Error querying database for modelId {modelId}: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
# If model found, verify search filter if provided
|
||||
if found_model is not None:
|
||||
if search is not None and search.strip():
|
||||
|
|
@ -8095,7 +8114,7 @@ async def model_info_v2(
|
|||
if search_lower not in model_name.lower():
|
||||
# Model found but doesn't match search filter
|
||||
found_model = None
|
||||
|
||||
|
||||
# Set all_models to the found model or empty list
|
||||
all_models = [found_model] if found_model is not None else []
|
||||
search_total_count: Optional[int] = len(all_models)
|
||||
|
|
@ -8135,14 +8154,16 @@ async def model_info_v2(
|
|||
llm_router=llm_router,
|
||||
all_models=all_models,
|
||||
)
|
||||
|
||||
|
||||
# Fill in model info based on config.yaml and litellm model_prices_and_context_window.json
|
||||
# This must happen before teamId filtering so that direct_access and access_via_team_ids are populated
|
||||
for i, _model in enumerate(all_models):
|
||||
all_models[i] = _enrich_model_info_with_litellm_data(
|
||||
model=_model, debug=debug if debug is not None else False, llm_router=llm_router
|
||||
model=_model,
|
||||
debug=debug if debug is not None else False,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
|
||||
# Apply teamId filter if provided
|
||||
if teamId is not None and teamId.strip():
|
||||
all_models = await _filter_models_by_team_id(
|
||||
|
|
@ -8153,14 +8174,14 @@ async def model_info_v2(
|
|||
)
|
||||
# Update search_total_count after teamId filter is applied
|
||||
search_total_count = len(all_models)
|
||||
|
||||
|
||||
# If modelId was provided, update search_total_count after filters are applied
|
||||
# to ensure pagination reflects the final filtered result (0 or 1)
|
||||
if modelId is not None:
|
||||
search_total_count = len(all_models)
|
||||
|
||||
verbose_proxy_logger.debug("all_models: %s", all_models)
|
||||
|
||||
|
||||
return _paginate_models_response(
|
||||
all_models=all_models,
|
||||
page=page,
|
||||
|
|
@ -10278,9 +10299,9 @@ async def get_config_list(
|
|||
hasattr(sub_field_info, "description")
|
||||
and sub_field_info.description is not None
|
||||
):
|
||||
nested_fields[idx].field_description = (
|
||||
sub_field_info.description
|
||||
)
|
||||
nested_fields[
|
||||
idx
|
||||
].field_description = sub_field_info.description
|
||||
idx += 1
|
||||
|
||||
_stored_in_db = None
|
||||
|
|
|
|||
|
|
@ -197,7 +197,13 @@ async def route_request(
|
|||
|
||||
elif "user_config" in data:
|
||||
router_config = data.pop("user_config")
|
||||
user_router = litellm.Router(**router_config)
|
||||
|
||||
# Filter router_config to only include valid Router.__init__ arguments
|
||||
# This prevents TypeError when invalid parameters are stored in the database
|
||||
valid_args = litellm.Router.get_valid_args()
|
||||
filtered_config = {k: v for k, v in router_config.items() if k in valid_args}
|
||||
|
||||
user_router = litellm.Router(**filtered_config)
|
||||
ret_val = getattr(user_router, f"{route_type}")(**data)
|
||||
user_router.discard()
|
||||
return ret_val
|
||||
|
|
|
|||
|
|
@ -643,6 +643,17 @@ class Router:
|
|||
self.initialize_router_endpoints()
|
||||
self.apply_default_settings()
|
||||
|
||||
@staticmethod
|
||||
def get_valid_args() -> List[str]:
|
||||
"""
|
||||
Returns a list of valid arguments for the Router.__init__ method.
|
||||
"""
|
||||
arg_spec = inspect.getfullargspec(Router.__init__)
|
||||
valid_args = arg_spec.args + arg_spec.kwonlyargs
|
||||
if "self" in valid_args:
|
||||
valid_args.remove("self")
|
||||
return valid_args
|
||||
|
||||
def apply_default_settings(self):
|
||||
"""
|
||||
Apply the default settings to the router.
|
||||
|
|
@ -4692,9 +4703,12 @@ class Router:
|
|||
# get num_retries from retry policy
|
||||
# Use the model_group captured at the start of the function, or get it from metadata
|
||||
# kwargs.get("model") at this point is the deployment model, not the model_group
|
||||
_model_group_for_retry_policy = model_group or _metadata.get("model_group") or kwargs.get("model")
|
||||
_model_group_for_retry_policy = (
|
||||
model_group or _metadata.get("model_group") or kwargs.get("model")
|
||||
)
|
||||
_retry_policy_retries = self.get_num_retries_from_retry_policy(
|
||||
exception=original_exception, model_group=_model_group_for_retry_policy
|
||||
exception=original_exception,
|
||||
model_group=_model_group_for_retry_policy,
|
||||
)
|
||||
if _retry_policy_retries is not None:
|
||||
num_retries = _retry_policy_retries
|
||||
|
|
@ -5879,7 +5893,10 @@ class Router:
|
|||
)
|
||||
# done reading model["litellm_params"]
|
||||
# Check if provider is supported: either in enum or JSON-configured
|
||||
if custom_llm_provider not in litellm.provider_list and not JSONProviderRegistry.exists(custom_llm_provider):
|
||||
if (
|
||||
custom_llm_provider not in litellm.provider_list
|
||||
and not JSONProviderRegistry.exists(custom_llm_provider)
|
||||
):
|
||||
raise Exception(f"Unsupported provider - {custom_llm_provider}")
|
||||
|
||||
#### DEPLOYMENT NAMES INIT ########
|
||||
|
|
|
|||
|
|
@ -1,9 +1,7 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
|
|
@ -49,7 +47,13 @@ async def test_route_request_dynamic_credentials(route_type):
|
|||
@pytest.mark.asyncio
|
||||
async def test_route_request_no_model_required():
|
||||
"""Test route types that don't require model parameter"""
|
||||
test_cases = ["amoderation", "aget_responses", "adelete_responses", "avector_store_create", "avector_store_search"]
|
||||
test_cases = [
|
||||
"amoderation",
|
||||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"avector_store_create",
|
||||
"avector_store_search",
|
||||
]
|
||||
|
||||
for route_type in test_cases:
|
||||
# Test data without model parameter
|
||||
|
|
@ -72,7 +76,13 @@ async def test_route_request_no_model_required():
|
|||
@pytest.mark.asyncio
|
||||
async def test_route_request_no_model_required_with_router_settings():
|
||||
"""Test route types that don't require model parameter with router settings"""
|
||||
test_cases = ["amoderation", "aget_responses", "adelete_responses", "avector_store_create", "avector_store_search"]
|
||||
test_cases = [
|
||||
"amoderation",
|
||||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"avector_store_create",
|
||||
"avector_store_search",
|
||||
]
|
||||
|
||||
for route_type in test_cases:
|
||||
# Test data with model parameter (it will be ignored for these route types)
|
||||
|
|
@ -121,6 +131,68 @@ async def test_route_request_no_model_required_with_router_settings_and_no_route
|
|||
with patch.object(
|
||||
litellm, "acompletion", return_value="fake_response"
|
||||
) as mock_completion:
|
||||
response = await route_request(data, None, "gpt-3.5-turbo", "acompletion")
|
||||
await route_request(data, None, "gpt-3.5-turbo", "acompletion")
|
||||
|
||||
mock_completion.assert_called_once_with(**data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_with_invalid_router_params():
|
||||
"""
|
||||
Test that route_request filters out invalid Router init params from 'user_config'.
|
||||
This covers the fix for https://github.com/BerriAI/litellm/issues/19693
|
||||
"""
|
||||
import litellm
|
||||
from litellm.router import Router
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
# Mock data with user_config containing invalid keys (simulating DB entry)
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"user_config": {
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo", "api_key": "test"},
|
||||
}
|
||||
],
|
||||
"model_alias_map": {"alias": "real_model"}, # INVALID PARAM
|
||||
"invalid_garbage_key": "crash_me", # INVALID PARAM
|
||||
},
|
||||
}
|
||||
|
||||
# We expect Router(**config) to succeed because of the filtering.
|
||||
# If filtering fails, this will raise TypeError and fail the test.
|
||||
try:
|
||||
# route_request calls getattr(user_router, route_type)(**data)
|
||||
# We'll mock the internal call to avoid making real network requests
|
||||
with pytest.MonkeyPatch.context() as m:
|
||||
# Mock the method that gets called on the router instance
|
||||
# We don't easily have access to the instance created INSIDE existing route_request
|
||||
# So we will wrap litellm.Router to spy on it or verify it doesn't crash
|
||||
|
||||
original_router_init = litellm.Router.__init__
|
||||
|
||||
def safe_router_init(self, **kwargs):
|
||||
# Verify that invalid keys are NOT present in kwargs
|
||||
assert "model_alias_map" not in kwargs
|
||||
assert "invalid_garbage_key" not in kwargs
|
||||
# Call original init (which would raise TypeError if invalid keys were present)
|
||||
original_router_init(self, **kwargs)
|
||||
|
||||
m.setattr(litellm.Router, "__init__", safe_router_init)
|
||||
|
||||
# Use 'acompletion' as the route_type
|
||||
# We also need to mock the completion method to avoid real calls
|
||||
m.setattr(Router, "acompletion", AsyncMock(return_value="success"))
|
||||
|
||||
response = await route_request(data, None, None, "acompletion")
|
||||
assert response == "success"
|
||||
|
||||
except TypeError as e:
|
||||
pytest.fail(
|
||||
f"route_request raised TypeError, implying invalid params were passed to Router: {e}"
|
||||
)
|
||||
except Exception:
|
||||
# Other exceptions might happen (e.g. valid config issues) but we care about TypeError here
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue