fix(team-routing): keep team model routing on public names

Remove team model_alias rewrites and resolve team deployments by team_public_model_name with team_id so sibling deployments stay in the routing candidate pool, with explicit logs showing candidate selection before load balancing.

Made-with: Cursor
This commit is contained in:
Sameer Kankute 2026-03-23 16:12:27 +05:30
parent b7cb9b6c21
commit bff657ddfc
No known key found for this signature in database
3 changed files with 96 additions and 789 deletions

View file

@ -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

View file

@ -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_<team_id>_<uuid>`), 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]:

View file

@ -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"