diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 2ca8e3daba3..694fa2b1d55 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -32,7 +32,7 @@ from litellm.proxy._types import ( ProxyErrorTypes, ProxyException, TeamModelAddRequest, - TeamModelDeleteRequest, + UpdateTeamRequest, UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -40,10 +40,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helpe from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin from litellm.proxy.management_endpoints.team_endpoints import ( team_model_add, - team_model_delete, -) -from litellm.proxy.management_endpoints.team_endpoints import ( - update_team as _legacy_update_team, + update_team, ) from litellm.proxy.management_helpers.audit_logs import create_object_audit_log from litellm.proxy.utils import PrismaClient @@ -61,14 +58,6 @@ from litellm.utils import get_utc_datetime router = APIRouter() -async def update_team(*args, **kwargs): - """ - Backward-compatible shim for tests/legacy call sites that patch this symbol. - Team model management now uses team_model_add/team_model_delete directly. - """ - return await _legacy_update_team(*args, **kwargs) - - class UpdatePublicModelGroupsRequest(BaseModel): """Request model for updating public model groups""" @@ -340,19 +329,12 @@ async def _add_team_model_to_db( _team_id = model_params.model_info.team_id if _team_id is None: return None - - # Capture the original public name FIRST, before any mutations original_model_name = model_params.model_name - - # Set team_public_model_name in model_info using the captured original_model_name - # This must happen BEFORE mutating model_params.model_name so _add_model_to_db - # serializes the correct team_public_model_name (not the internal UUID name) if original_model_name: model_params.model_info.team_public_model_name = original_model_name - # Generate and assign unique internal model_name LAST - # (after team_public_model_name is safely stored) unique_model_name = f"model_name_{_team_id}_{uuid.uuid4()}" + model_params.model_name = unique_model_name ## CREATE MODEL IN DB ## @@ -362,15 +344,14 @@ async def _add_team_model_to_db( prisma_client=prisma_client, ) - if original_model_name: - await team_model_add( - data=TeamModelAddRequest( - team_id=_team_id, - models=[original_model_name], - ), - http_request=Request(scope={"type": "http"}), - user_api_key_dict=user_api_key_dict, - ) + await team_model_add( + data=TeamModelAddRequest( + team_id=_team_id, + models=[original_model_name], + ), + http_request=Request(scope={"type": "http"}), + user_api_key_dict=user_api_key_dict, + ) return model_response @@ -436,7 +417,6 @@ async def _update_team_model_in_db( db_model=db_model, patch_data=patch_data, user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, ) return update_db_model(db_model=db_model, updated_patch=patch_data) @@ -476,85 +456,19 @@ async def _setup_new_team_model_assignment( ) -async def _get_team_deployments( - team_id: str, prisma_client: PrismaClient -) -> List[LiteLLM_ProxyModelTable]: - """ - Fetch all deployments for a given team_id from the database. - - Centralizes team deployment queries to ensure consistent filtering and error handling. - This is the established helper pattern for team deployment DB access in this module. - - Note: Direct Prisma call is intentional here as this IS the helper function that - encapsulates the DB access pattern for team deployments. - """ - response = await prisma_client.db.litellm_proxymodeltable.find_many( - where={ - "model_info": { - "path": ["team_id"], - "equals": team_id, - } - } - ) - return response if response else [] - - async def _update_existing_team_model_assignment( team_id: str, public_model_name: str, db_model: Deployment, patch_data: updateDeployment, user_api_key_dict: UserAPIKeyAuth, - prisma_client: Optional[PrismaClient], ) -> None: - """Update an existing team model if the public name changed. - - Note on DB scan: Prisma's JSON filtering does not support compound AND conditions - across multiple JSON paths, so we fetch all deployments for the team and filter - team_public_model_name in Python. For teams with many deployments this scan grows - linearly; if team deployment counts become large this should be revisited. - """ - - def _get_team_public_model_name( - model_info: Optional[Union[dict, str]] - ) -> Optional[str]: - if isinstance(model_info, dict): - value = model_info.get("team_public_model_name") - return value if isinstance(value, str) else None - if isinstance(model_info, str): - try: - parsed = json.loads(model_info) - except (TypeError, ValueError): - return None - if isinstance(parsed, dict): - value = parsed.get("team_public_model_name") - return value if isinstance(value, str) else None - return None - + """Update an existing team model if the public name changed.""" old_public_name = ( db_model.model_info.team_public_model_name if db_model.model_info else None ) if old_public_name and public_model_name != old_public_name: - # Clear user-supplied public name from patch before any early return so the - # caller does not overwrite the internal UUID-based model_name in the DB. - patch_data.model_name = None - if prisma_client is None: - verbose_proxy_logger.warning( - "prisma_client not initialized; skipping public name update entirely to avoid orphaned entries" - ) - return - - # Query DB for all team deployments to check for sibling deployments - team_deployments = await _get_team_deployments(team_id, prisma_client) - other_deployments_with_old_name = [ - d - for d in team_deployments - if d.model_name != db_model.model_name - and _get_team_public_model_name(d.model_info) == old_public_name - ] - - # Add new name first, then delete old name to prevent access loss on partial failure await team_model_add( data=TeamModelAddRequest( team_id=team_id, @@ -564,31 +478,6 @@ async def _update_existing_team_model_assignment( user_api_key_dict=user_api_key_dict, ) - if not other_deployments_with_old_name: - await team_model_delete( - data=TeamModelDeleteRequest( - team_id=team_id, - models=[old_public_name], - ), - http_request=Request(scope={"type": "http"}), - user_api_key_dict=user_api_key_dict, - ) - elif not old_public_name and public_model_name: - # First-time assignment of public name on an existing team deployment: - # ensure the team's models list is updated so team routing can resolve it. - await team_model_add( - data=TeamModelAddRequest( - team_id=team_id, - models=[public_model_name], - ), - http_request=Request(scope={"type": "http"}), - user_api_key_dict=user_api_key_dict, - ) - # else: old_public_name == public_model_name (no rename needed) - # No team_model_add/delete calls required; public name is already registered - - # Always clear patch_data.model_name to prevent caller from overwriting - # the internal UUID-based model_name in the DB with the user-supplied public name patch_data.model_name = None diff --git a/litellm/router.py b/litellm/router.py index b96d20926c8..e146d60e359 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -54,11 +54,7 @@ from litellm.caching.caching import ( RedisCache, RedisClusterCache, ) -from litellm.constants import ( - DEFAULT_HEALTH_CHECK_INTERVAL, - DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER, - DEFAULT_MAX_LRU_CACHE_SIZE, -) +from litellm.constants import DEFAULT_MAX_LRU_CACHE_SIZE from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import ( @@ -117,7 +113,6 @@ from litellm.router_utils.handle_error import ( async_raise_no_deployment_exception, send_llm_exception_alert, ) -from litellm.router_utils.health_state_cache import DeploymentHealthCache from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( DeploymentAffinityCheck, ) @@ -308,8 +303,6 @@ class Router: deployment_affinity_ttl_seconds: int = 3600, model_group_affinity_config: Optional[Dict[str, List[str]]] = None, ignore_invalid_deployments: bool = False, - enable_health_check_routing: bool = False, - health_check_staleness_threshold: Optional[int] = None, ) -> None: """ Initialize the Router class with the given parameters for caching, reliability, and routing strategy. @@ -474,8 +467,6 @@ class Router: # Initialize model name to deployment indices mapping for O(1) lookups # Maps model_name -> list of indices in model_list self.model_name_to_deployment_indices: Dict[str, List[int]] = {} - # Maps (team_id, team_public_model_name) -> list of indices in model_list - self.team_model_to_deployment_indices: Dict[Tuple[str, str], List[int]] = {} if model_list is not None: # set_model_list will build indices automatically @@ -500,13 +491,6 @@ class Router: cache=self.cache, default_cooldown_time=self.cooldown_time ) self.disable_cooldowns = disable_cooldowns - self.enable_health_check_routing = enable_health_check_routing - _staleness = health_check_staleness_threshold or ( - DEFAULT_HEALTH_CHECK_INTERVAL * DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER - ) - self.health_state_cache = DeploymentHealthCache( - cache=self.cache, staleness_threshold=float(_staleness) - ) self.failed_calls = ( InMemoryCache() ) # cache to track failed call per deployment, if num failed calls within 1 minute > allowed fails, then add it to cooldown @@ -5304,63 +5288,6 @@ class Router: if "fallback_depth" not in input_kwargs: input_kwargs["fallback_depth"] = 0 - # ORDER-BASED FALLBACKS: prepend higher order levels to the fallback list - # Skip for error types that have their own dedicated fallback handlers - _skip_order_fallback = isinstance( - e, - (litellm.ContextWindowExceededError, litellm.ContentPolicyViolationError), - ) - all_deployments = self._get_all_deployments(model_name=original_model_group) - _order_set: set = { - d.get("litellm_params", {}).get("order") - for d in all_deployments - if d.get("litellm_params", {}).get("order") is not None - } - order_values: list = sorted(_order_set) - if len(order_values) > 1 and not _skip_order_fallback: - # Determine which order levels have already been tried - current_target = kwargs.get("_target_order") - skip_up_to = ( - current_target if current_target is not None else order_values[0] - ) - # Build order-based fallback entries (skip already-tried levels) - order_fallback_entries: List = [ - {"model": original_model_group, "_target_order": o} - for o in order_values - if o > skip_up_to - ] - # Get external fallbacks — handle both standard and non-standard formats - external_fallback_group: Optional[List] = None - if fallbacks is not None and model_group is not None: - if _check_non_standard_fallback_format(fallbacks=fallbacks): - # Non-standard formats (e.g. ["claude-3-haiku"] or - # [{"model": "...", "messages": [...]}]) are passed through directly - external_fallback_group = fallbacks - else: - external_fallback_group, generic_idx = get_fallback_model_group( - fallbacks=fallbacks, - model_group=cast(str, model_group), - ) - if external_fallback_group is None and generic_idx is not None: - external_fallback_group = fallbacks[generic_idx]["*"] - # Combined list: order fallbacks first, then external - combined_fallbacks = order_fallback_entries + ( - external_fallback_group or [] - ) - - if combined_fallbacks: - input_kwargs.update( - { - "fallback_model_group": combined_fallbacks, - "original_model_group": original_model_group, - } - ) - response = await run_async_fallback( - *args, - **input_kwargs, - ) - return response - try: verbose_router_logger.info("Trying to fallback b/w models") @@ -6908,7 +6835,6 @@ class Router: self.model_list = [] self.model_id_to_deployment_index_map = {} # Reset the index self.model_name_to_deployment_indices = {} # Reset the model_name index - self.team_model_to_deployment_indices = {} # Reset the team_model index self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() # we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works @@ -7206,17 +7132,16 @@ class Router: # Update model_name_to_deployment_indices for model_name, indices in list(self.model_name_to_deployment_indices.items()): - # Build new list without mutating the original + # Remove the deleted index + if removal_idx in indices: + indices.remove(removal_idx) + + # Decrement all indices greater than removal_idx updated_indices = [] for idx in indices: - if idx == removal_idx: - # Skip the removed index - continue - elif idx > removal_idx: - # Decrement indices after removal + if idx > removal_idx: updated_indices.append(idx - 1) else: - # Keep indices before removal unchanged updated_indices.append(idx) # Update or remove the entry @@ -7225,46 +7150,6 @@ class Router: else: del self.model_name_to_deployment_indices[model_name] - # Update team_model_to_deployment_indices - for key, indices in list(self.team_model_to_deployment_indices.items()): - # Build new list without mutating the original - updated_indices = [] - for idx in indices: - if idx == removal_idx: - # Skip the removed index - continue - elif idx > removal_idx: - # Decrement indices after removal - updated_indices.append(idx - 1) - else: - # Keep indices before removal unchanged - updated_indices.append(idx) - - # Update or remove the entry - if len(updated_indices) > 0: - self.team_model_to_deployment_indices[key] = updated_indices - else: - del self.team_model_to_deployment_indices[key] - - def _update_team_model_index(self, model: dict, idx: int) -> None: - """ - Helper to update team_model_to_deployment_indices for a single deployment. - - Parameters: - - model: dict - the deployment to index - - idx: int - the index in model_list - """ - team_id = (model.get("model_info") or {}).get("team_id") - team_public_model_name = (model.get("model_info") or {}).get( - "team_public_model_name" - ) - if team_id and team_public_model_name: - key = (team_id, team_public_model_name) - if key not in self.team_model_to_deployment_indices: - self.team_model_to_deployment_indices[key] = [] - if idx not in self.team_model_to_deployment_indices[key]: - self.team_model_to_deployment_indices[key].append(idx) - def _add_model_to_list_and_index_map( self, model: dict, model_id: Optional[str] = None ) -> None: @@ -7293,9 +7178,6 @@ class Router: self.model_name_to_deployment_indices[model_name] = [] self.model_name_to_deployment_indices[model_name].append(idx) - # Update team_model index for O(1) team-scoped lookup - self._update_team_model_index(model, idx) - def upsert_deployment(self, deployment: Deployment) -> Optional[Deployment]: """ Add or update deployment @@ -7314,10 +7196,7 @@ class Router: ) if _deployment_on_router is not None: # deployment with this model_id exists on the router - if ( - deployment.litellm_params == _deployment_on_router.litellm_params - and deployment.model_info == _deployment_on_router.model_info - ): + if deployment.litellm_params == _deployment_on_router.litellm_params: # No need to update return None @@ -7807,8 +7686,8 @@ class Router: max_tokens=None, max_input_tokens=None, max_output_tokens=None, - input_cost_per_token=None, - output_cost_per_token=None, + input_cost_per_token=0, + output_cost_per_token=0, litellm_provider=llm_provider, mode=mode, supported_openai_params=supported_openai_params, @@ -8129,7 +8008,6 @@ class Router: instead of O(n) linear scan through the entire model_list. """ self.model_name_to_deployment_indices.clear() - self.team_model_to_deployment_indices.clear() for idx, model in enumerate(model_list): model_name = model.get("model_name") @@ -8138,8 +8016,6 @@ class Router: self.model_name_to_deployment_indices[model_name] = [] self.model_name_to_deployment_indices[model_name].append(idx) - self._update_team_model_index(model, idx) - def _build_model_id_to_deployment_index_map(self, model_list: list): """ Build model index from model list to enable O(1) lookups immediately. @@ -8289,8 +8165,6 @@ class Router: if model.get("model_info", {}).get("team_id") == team_id: return team_model_name - # No team-scoped deployment found; wildcard/pattern routes are - # handled downstream by the pattern_router in _common_checks_available_deployment. return None def should_include_deployment( @@ -8301,22 +8175,12 @@ class Router: """ if ( team_id is not None - and (model.get("model_info") or {}).get("team_id") == team_id - and model_name - == (model.get("model_info") or {}).get("team_public_model_name") + and model["model_info"].get("team_id") == team_id + and model_name == model["model_info"].get("team_public_model_name") ): return True elif model_name is not None and model["model_name"] == model_name: - # Fallback: check by internal model_name for non-team deployments - # or deployments that haven't been migrated to team_public_model_name yet - model_team_id = (model.get("model_info") or {}).get("team_id") - if ( - team_id is None # requester has no team constraint - or model_team_id is None # global deployment - accessible to all teams - or model_team_id == team_id # deployment belongs to requester's team - ): - return True - # No match: deployment is for a different team or doesn't match the requested model + return True return False def _get_all_deployments( @@ -8333,36 +8197,9 @@ class Router: if team_id specified, only return team-specific models Optimized with O(1) index lookup instead of O(n) linear scan. - - Note: when team_id is provided, O(1) lookup in - `team_model_to_deployment_indices` only applies when `model_name` is the - team public model name. If a caller passes an internal deployment model - name (for example, `model_name__`), this method falls back - to the standard model-name index / scan path. """ returned_models: List[DeploymentTypedDict] = [] - # O(1) lookup in team_model index when team_id is provided - if team_id is not None: - key = (team_id, model_name) - if key in self.team_model_to_deployment_indices: - indices = self.team_model_to_deployment_indices[key] - # O(k) where k = team deployments for this model_name (typically 1-10) - for idx in indices: - model = self.model_list[idx] - if not self.should_include_deployment( - model_name=model_name, model=model, team_id=team_id - ): - continue - if model_alias is not None: - alias_model = model.copy() - alias_model["model_name"] = model_alias - returned_models.append(alias_model) - else: - returned_models.append(model) - if returned_models: - return returned_models - # O(1) lookup in model_name index if model_name in self.model_name_to_deployment_indices: indices = self.model_name_to_deployment_indices[model_name] @@ -8957,6 +8794,12 @@ class Router: if i not in invalid_model_indices ] + ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) + if len(_returned_deployments) > 0: + _returned_deployments = litellm.utils._get_order_filtered_deployments( + _returned_deployments + ) + return _returned_deployments def _get_model_from_alias(self, model: str) -> Optional[str]: @@ -9027,14 +8870,36 @@ class Router: model = _model_from_alias if model not in self.model_names: - # Check for team-specific deployments by team_public_model_name. - # This intentionally takes priority over team pattern routers below, - # so that named team deployments shadow wildcard/pattern routes. + # Check for team-specific deployments by team_public_model_name if request_team_id is not None: team_deployments = self._get_all_deployments( model_name=model, team_id=request_team_id ) if team_deployments: + candidate_details = [] + for deployment in team_deployments: + deployment_info = deployment.get("model_info", {}) or {} + deployment_params = deployment.get("litellm_params", {}) or {} + candidate_details.append( + { + "model_name": deployment.get("model_name"), + "model_id": deployment_info.get("id"), + "team_public_model_name": deployment_info.get( + "team_public_model_name" + ), + "api_base": deployment_params.get("api_base"), + } + ) + verbose_router_logger.info( + "🔥 routing_candidates_before_lb " + f"model={model} count={len(team_deployments)} " + f"candidates={candidate_details}" + ) + if len(team_deployments) > 1: + verbose_router_logger.info( + "🔥 load_balancer_candidate_pool " + f"model={model} candidate_count={len(team_deployments)}" + ) return model, team_deployments # check if provider/ specific wildcard routing use pattern matching @@ -9069,14 +8934,38 @@ class Router: ## get healthy deployments ### get all deployments - healthy_deployments = self._get_all_deployments( - model_name=model, team_id=request_team_id - ) + healthy_deployments = self._get_all_deployments(model_name=model) if len(healthy_deployments) == 0: # check if the user sent in a deployment name instead healthy_deployments = self._get_deployment_by_litellm_model(model=model) + if isinstance(healthy_deployments, list) and len(healthy_deployments) > 0: + candidate_details = [] + for deployment in healthy_deployments: + deployment_info = deployment.get("model_info", {}) or {} + deployment_params = deployment.get("litellm_params", {}) or {} + candidate_details.append( + { + "model_name": deployment.get("model_name"), + "model_id": deployment_info.get("id"), + "team_public_model_name": deployment_info.get( + "team_public_model_name" + ), + "api_base": deployment_params.get("api_base"), + } + ) + verbose_router_logger.info( + "🔥 routing_candidates_before_lb " + f"model={model} count={len(healthy_deployments)} " + f"candidates={candidate_details}" + ) + if len(healthy_deployments) > 1: + verbose_router_logger.info( + "🔥 load_balancer_candidate_pool " + f"model={model} candidate_count={len(healthy_deployments)}" + ) + if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug( f"initial list of deployments: {healthy_deployments}" @@ -9092,9 +8981,7 @@ class Router: ) # Re-assign model to the fallback and try to get deployments again model = fallback_model - healthy_deployments = self._get_all_deployments( - model_name=model, team_id=request_team_id - ) + healthy_deployments = self._get_all_deployments(model_name=model) # If still no deployments after checking for fallbacks, raise an error if len(healthy_deployments) == 0: @@ -9171,14 +9058,6 @@ class Router: if isinstance(healthy_deployments, dict): return healthy_deployments - # Health-check-based filtering (before cooldown) - healthy_deployments = ( - await self._async_filter_health_check_unhealthy_deployments( - healthy_deployments=healthy_deployments, - parent_otel_span=parent_otel_span, - ) - ) - cooldown_deployments = await _async_get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) @@ -9217,12 +9096,6 @@ class Router: ), ) - ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) - _target_order = (request_kwargs or {}).pop("_target_order", None) - healthy_deployments = litellm.utils._get_order_filtered_deployments( - cast(List[Dict], healthy_deployments), target_order=_target_order - ) - if len(healthy_deployments) == 0: exception = await async_raise_no_deployment_exception( litellm_router_instance=self, @@ -9610,13 +9483,6 @@ class Router: parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs( request_kwargs ) - - # Health-check-based filtering (before cooldown) - healthy_deployments = self._filter_health_check_unhealthy_deployments( - healthy_deployments=healthy_deployments, - parent_otel_span=parent_otel_span, - ) - cooldown_deployments = _get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) @@ -9634,12 +9500,6 @@ class Router: request_kwargs=request_kwargs, ) - ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) - _target_order = (request_kwargs or {}).pop("_target_order", None) - healthy_deployments = litellm.utils._get_order_filtered_deployments( - healthy_deployments, target_order=_target_order - ) - if len(healthy_deployments) == 0: model_ids = self.get_model_ids(model_name=model) _cooldown_time = self.cooldown_cache.get_min_cooldown( @@ -9782,14 +9642,10 @@ class Router: llm_provider="", ) - # 4. Apply health-check and cooldown filtering + # 4. Apply cooldown filtering parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs( request_kwargs ) - pass_through_deployments = self._filter_health_check_unhealthy_deployments( - healthy_deployments=pass_through_deployments, - parent_otel_span=parent_otel_span, - ) cooldown_deployments = _get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) @@ -9911,67 +9767,6 @@ class Router: if deployment["model_info"]["id"] not in cooldown_set ] - async def _async_filter_health_check_unhealthy_deployments( - self, - healthy_deployments: List[Dict], - parent_otel_span: Optional[Span] = None, - ) -> List[Dict]: - """ - Filter out deployments marked unhealthy by background health checks. - No-op when enable_health_check_routing is False. - Returns all deployments if health state is unavailable, stale, or would - exclude every candidate (safety net). - """ - if not self.enable_health_check_routing: - return healthy_deployments - - unhealthy_ids = ( - await self.health_state_cache.async_get_unhealthy_deployment_ids( - parent_otel_span=parent_otel_span - ) - ) - if not unhealthy_ids: - return healthy_deployments - - filtered = [ - d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids - ] - - if not filtered: - verbose_router_logger.warning( - "All deployments marked unhealthy by health checks, bypassing health filter" - ) - return healthy_deployments - - return filtered - - def _filter_health_check_unhealthy_deployments( - self, - healthy_deployments: List[Dict], - parent_otel_span: Optional[Span] = None, - ) -> List[Dict]: - """Sync version of _async_filter_health_check_unhealthy_deployments.""" - if not self.enable_health_check_routing: - return healthy_deployments - - unhealthy_ids = self.health_state_cache.get_unhealthy_deployment_ids( - parent_otel_span=parent_otel_span - ) - if not unhealthy_ids: - return healthy_deployments - - filtered = [ - d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids - ] - - if not filtered: - verbose_router_logger.warning( - "All deployments marked unhealthy by health checks, bypassing health filter" - ) - return healthy_deployments - - return filtered - def _filter_pass_through_deployments( self, healthy_deployments: List[Dict] ) -> List[Dict]: diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index c3e3d1ecbd9..5ef7face1f2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -14,7 +14,6 @@ sys.path.insert( ) # Adds the parent directory to the system path from litellm.proxy._types import ( LiteLLM_ModelTable, - LiteLLM_ProxyModelTable, LiteLLM_TeamTable, LitellmUserRoles, Member, @@ -29,15 +28,9 @@ from litellm.types.router import Deployment, LiteLLM_Params, updateDeployment class MockPrismaClient: - def __init__( - self, - team_exists: bool = True, - user_admin: bool = True, - sibling_deployments: list = None, - ): + def __init__(self, team_exists: bool = True, user_admin: bool = True): self.team_exists = team_exists self.user_admin = user_admin - self.sibling_deployments = sibling_deployments or [] self.db = self async def find_unique(self, where): @@ -53,53 +46,10 @@ class MockPrismaClient: ) return None - async def find_many(self, where): - # Filter sibling deployments by team_id if where clause specifies it - if not self.sibling_deployments: - return [] - - # Extract team_id from where clause if present - team_id_filter = None - if where and "model_info" in where: - model_info_filter = where["model_info"] - if isinstance(model_info_filter, dict) and "path" in model_info_filter: - if ( - model_info_filter["path"] == ["team_id"] - and "equals" in model_info_filter - ): - team_id_filter = model_info_filter["equals"] - - # Filter deployments by team_id if specified - if team_id_filter: - - def _get_team_id(model_info): - if isinstance(model_info, dict): - return model_info.get("team_id") - if isinstance(model_info, str): - try: - parsed = json.loads(model_info) - except (TypeError, ValueError): - return None - if isinstance(parsed, dict): - return parsed.get("team_id") - return None - - return [ - d - for d in self.sibling_deployments - if _get_team_id(d.model_info) == team_id_filter - ] - - return self.sibling_deployments - @property def litellm_teamtable(self): return self - @property - def litellm_proxymodeltable(self): - return self - class MockLLMRouter: def __init__(self): @@ -636,6 +586,8 @@ class TestTeamModelSiblingRouting: team_id = "team_no_alias" public_name = "gpt-4.1-mini" + mock_update_team = AsyncMock() + async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client): return MagicMock(model_id=str(uuid.uuid4())) @@ -655,6 +607,9 @@ class TestTeamModelSiblingRouting: model_info=ModelInfo(team_id=team_id), ) with patch( + "litellm.proxy.management_endpoints.model_management_endpoints.update_team", + mock_update_team, + ), patch( "litellm.proxy.management_endpoints.model_management_endpoints._add_model_to_db", side_effect=mock_add_model_to_db, ), patch( @@ -667,6 +622,7 @@ class TestTeamModelSiblingRouting: prisma_client=prisma_client, ) + mock_update_team.assert_not_called() assert mock_team_model_add.call_count == 2 @pytest.mark.asyncio @@ -707,15 +663,6 @@ class TestTeamModelSiblingRouting: "team_public_model_name": public_name, }, }, - { - "model_name": "global-gpt-4o", - "litellm_params": { - "model": "azure/gpt-4o", - "api_key": "global-key", - "api_base": "https://global.openai.azure.com", - }, - "model_info": {}, # No team_id - global deployment - }, ], ) @@ -736,38 +683,6 @@ class TestTeamModelSiblingRouting: "https://westus.openai.azure.com", } - def test_global_deployments_accessible_to_teams(self): - """Test that global deployments (no team_id) are accessible to all teams""" - import litellm - - router = litellm.Router( - model_list=[ - { - "model_name": "global-gpt-4o", - "litellm_params": { - "model": "azure/gpt-4o", - "api_key": "global-key", - "api_base": "https://global.openai.azure.com", - }, - "model_info": {}, # No team_id - global deployment - }, - ], - ) - - # Global deployment should be accessible when team_id is provided - deployments = router._get_all_deployments( - model_name="global-gpt-4o", team_id="teamA" - ) - assert len(deployments) == 1 - assert deployments[0]["model_name"] == "global-gpt-4o" - - # should_include_deployment should return True for global deployments - assert router.should_include_deployment( - model_name="global-gpt-4o", - model={"model_name": "global-gpt-4o", "model_info": {}}, - team_id="teamA", - ) - class TestTeamModelUpdate: """Test team model update handles team_id consistently with model creation""" @@ -802,10 +717,10 @@ class TestTeamModelUpdate: "litellm.proxy.proxy_server.premium_user", True, ), patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_team_model_add, patch( "litellm.proxy.management_endpoints.model_management_endpoints.update_team" - ) as mock_update_team: + ) as mock_update_team, patch( + "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" + ) as mock_team_model_add: result = await _update_team_model_in_db( db_model=db_model, patch_data=patch_data, @@ -815,201 +730,8 @@ class TestTeamModelUpdate: assert result.get("model_name", "").startswith("model_name_test_team_123_") assert "team_public_model_name" in str(result.get("model_info", "")) - # team_model_add must be called to add public name to team's models list - mock_team_model_add.assert_called_once() - # update_team (model_aliases write) must NOT be called in the new implementation mock_update_team.assert_not_called() - - @pytest.mark.asyncio - async def test_rename_preserves_old_name_when_siblings_exist(self): - """Test that renaming a deployment preserves old public name when sibling deployments still use it""" - from unittest.mock import MagicMock - - from litellm.proxy.management_endpoints.model_management_endpoints import ( - _update_existing_team_model_assignment, - ) - from litellm.types.router import ModelInfo - - # Create a deployment being renamed - db_model = Deployment( - model_name="model_name_team_123_uuid1", - litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), - ) - - # Create a sibling deployment that still uses the old public name - sibling_deployment = MagicMock() - sibling_deployment.model_name = "model_name_team_123_uuid2" - sibling_deployment.model_info = { - "team_id": "team_123", - "team_public_model_name": "old-public-name", - } - - prisma_client = MockPrismaClient( - team_exists=True, sibling_deployments=[sibling_deployment] - ) - - patch_data = updateDeployment( - model_name="new-public-name", - model_info=ModelInfo(team_id="team_123"), - ) - - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user", - user_role=LitellmUserRoles.PROXY_ADMIN, - ) - - with patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add: - await _update_existing_team_model_assignment( - team_id="team_123", - public_model_name="new-public-name", - db_model=db_model, - patch_data=patch_data, - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, # type: ignore - ) - - # team_model_delete should NOT be called because sibling exists - mock_delete.assert_not_called() - # team_model_add should be called to add new public name - mock_add.assert_called_once() - - @pytest.mark.asyncio - async def test_first_time_public_name_assignment_adds_team_model(self): - """If existing team deployment had no public name, first assignment must call team_model_add.""" - from litellm.proxy.management_endpoints.model_management_endpoints import ( - _update_existing_team_model_assignment, - ) - from litellm.types.router import ModelInfo - - db_model = Deployment( - model_name="model_name_team_123_uuid1", - litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo(team_id="team_123"), - ) - - patch_data = updateDeployment( - model_name="new-public-name", - model_info=ModelInfo(team_id="team_123"), - ) - - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user", - user_role=LitellmUserRoles.PROXY_ADMIN, - ) - - with patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add: - await _update_existing_team_model_assignment( - team_id="team_123", - public_model_name="new-public-name", - db_model=db_model, - patch_data=patch_data, - user_api_key_dict=user_api_key_dict, - prisma_client=None, - ) - - mock_add.assert_called_once() - mock_delete.assert_not_called() - - @pytest.mark.asyncio - async def test_rename_with_prisma_none_clears_patch_model_name(self): - """Rename path must clear patch_data.model_name even when prisma is unavailable (P1).""" - from litellm.proxy.management_endpoints.model_management_endpoints import ( - _update_existing_team_model_assignment, - ) - from litellm.types.router import ModelInfo - - db_model = Deployment( - model_name="model_name_team_123_uuid1", - litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), - ) - patch_data = updateDeployment( - model_name="new-public-name", - model_info=ModelInfo(team_id="team_123"), - ) - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user", - user_role=LitellmUserRoles.PROXY_ADMIN, - ) - - await _update_existing_team_model_assignment( - team_id="team_123", - public_model_name="new-public-name", - db_model=db_model, - patch_data=patch_data, - user_api_key_dict=user_api_key_dict, - prisma_client=None, - ) - - assert patch_data.model_name is None - - @pytest.mark.asyncio - async def test_rename_handles_legacy_string_model_info(self): - """Test rename path handles legacy string-encoded model_info rows without crashing.""" - from unittest.mock import MagicMock - - from litellm.proxy.management_endpoints.model_management_endpoints import ( - _update_existing_team_model_assignment, - ) - from litellm.types.router import ModelInfo - - db_model = Deployment( - model_name="model_name_team_123_uuid1", - litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), - ) - - sibling_deployment = MagicMock() - sibling_deployment.model_name = "model_name_team_123_uuid2" - sibling_deployment.model_info = ( - '{"team_id":"team_123","team_public_model_name":"old-public-name"}' - ) - - prisma_client = MockPrismaClient( - team_exists=True, sibling_deployments=[sibling_deployment] - ) - - patch_data = updateDeployment( - model_name="new-public-name", - model_info=ModelInfo(team_id="team_123"), - ) - - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user", - user_role=LitellmUserRoles.PROXY_ADMIN, - ) - - with patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, patch( - "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add: - await _update_existing_team_model_assignment( - team_id="team_123", - public_model_name="new-public-name", - db_model=db_model, - patch_data=patch_data, - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, # type: ignore - ) - - mock_delete.assert_not_called() - mock_add.assert_called_once() + mock_team_model_add.assert_called_once() @pytest.mark.asyncio async def test_patch_model_with_team_id_validates_permissions(self): @@ -1177,102 +899,3 @@ class TestModelInfoEndpoint: assert result["id"] == "team-model-1" assert result["object"] == "model" assert result["owned_by"] == "custom" - - -class TestAddAndDeleteModelLifecycle: - """ - Mock replacement for test_add_and_delete_models in tests/test_models.py. - - The original integration test required a live proxy + OPENAI_API_KEY. - This test verifies the same lifecycle (add → delete → double-delete fails) - by calling the endpoint handlers directly with mocked DB. - """ - - @pytest.mark.asyncio - async def test_add_then_delete_model(self): - """ - - Add model via add_new_model → returns model_id - - Delete model via delete_model → returns success - - Delete same model again → raises (model not found) - """ - from litellm.proxy.management_endpoints.model_management_endpoints import ( - add_new_model, - delete_model as delete_model_endpoint, - ) - from litellm.proxy.management_endpoints.model_management_endpoints import ( - ModelInfoDelete, - ) - - model_id = "lifecycle-test-model-123" - admin_user = UserAPIKeyAuth( - user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) - - # Build a real LiteLLM_ProxyModelTable for the DB mock to return - db_row = LiteLLM_ProxyModelTable( - model_id=model_id, - model_name="lifecycle-model", - litellm_params={"model": "openai/gpt-4.1-nano"}, - model_info={"id": model_id}, - created_by="test-admin", - updated_by="test-admin", - ) - - mock_prisma = MagicMock() - mock_prisma.db = MagicMock() - mock_prisma.db.litellm_proxymodeltable = AsyncMock() - mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) - mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) - - mock_proxy_config = MagicMock() - mock_proxy_config.add_deployment = AsyncMock() - - mock_router = MagicMock() - mock_router.delete_deployment = MagicMock() - - _PS = "litellm.proxy.proxy_server" - _ENCRYPT = "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper" - with patch(f"{_PS}.prisma_client", mock_prisma), \ - patch(f"{_PS}.store_model_in_db", True), \ - patch(f"{_PS}.proxy_config", mock_proxy_config), \ - patch(f"{_PS}.proxy_logging_obj", MagicMock()), \ - patch(f"{_PS}.general_settings", {}), \ - patch(f"{_PS}.premium_user", True), \ - patch(f"{_PS}.llm_router", mock_router), \ - patch(_ENCRYPT, side_effect=lambda value, **kwargs: value): - - # --- ADD --- - add_result = await add_new_model( - model_params=Deployment( - model_name="lifecycle-model", - litellm_params=LiteLLM_Params( - model="openai/gpt-4.1-nano", api_key="fake-key" - ), - model_info={"id": model_id}, - ), - user_api_key_dict=admin_user, - ) - assert add_result.model_id == model_id - - # --- DELETE --- - delete_result = await delete_model_endpoint( - model_info=ModelInfoDelete(id=model_id), - user_api_key_dict=admin_user, - ) - assert "deleted successfully" in delete_result["message"] - - # --- DELETE again should fail (model not found) --- - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=None - ) - from litellm.proxy.proxy_server import ProxyException - - with pytest.raises(ProxyException) as exc_info: - await delete_model_endpoint( - model_info=ModelInfoDelete(id=model_id), - user_api_key_dict=admin_user, - ) - assert str(exc_info.value.code) == "400"