mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
b7cb9b6c21
commit
bff657ddfc
3 changed files with 96 additions and 789 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue