From 323c54d11730c0acc816ff1d51a63987b142f019 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 23 May 2025 20:58:36 -0700 Subject: [PATCH] fix(route_llm_request.py): map team model from list in route llm request remove unnecessary proxy model table lookup for model alias Fixes issue where aliases weren't being consistently written to model table --- litellm/proxy/route_llm_request.py | 31 +++++++++++++++++++++++++++++- litellm/router.py | 20 +++++++++++++++++++ 2 files changed, 50 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index eea6f37189a..ad570f2f09f 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -33,6 +33,25 @@ class ProxyModelNotFoundError(HTTPException): super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail) +def get_team_id_from_data(data: dict) -> Optional[str]: + """ + Get the team id from the data's metadata or litellm_metadata params. + """ + if ( + "metadata" in data + and data["metadata"] is not None + and "user_api_key_team_id" in data["metadata"] + ): + return data["metadata"].get("user_api_key_team_id") + elif ( + "litellm_metadata" in data + and data["litellm_metadata"] is not None + and "user_api_key_team_id" in data["litellm_metadata"] + ): + return data["litellm_metadata"].get("user_api_key_team_id") + return None + + async def route_request( data: dict, llm_router: Optional[LitellmRouter], @@ -55,6 +74,7 @@ async def route_request( """ Common helper to route the request """ + team_id = get_team_id_from_data(data) router_model_names = llm_router.model_names if llm_router is not None else [] if "api_key" in data or "api_base" in data: return getattr(llm_router, f"{route_type}")(**data) @@ -78,7 +98,16 @@ async def route_request( models = [model.strip() for model in data.pop("model").split(",")] return llm_router.abatch_completion(models=models, **data) elif llm_router is not None: - if ( + team_model_name = ( + llm_router.map_team_model(data["model"], team_id) + if team_id is not None + else None + ) + if team_model_name is not None: + data["model"] = team_model_name + return getattr(llm_router, f"{route_type}")(**data) + + elif ( data["model"] in router_model_names or data["model"] in llm_router.get_model_ids() ): diff --git a/litellm/router.py b/litellm/router.py index 38f187d9a8b..90b72b69f9d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -5407,6 +5407,26 @@ class Router: ids.append(id) return ids + def map_team_model(self, team_model_name: str, team_id: str) -> Optional[str]: + """ + Map a team model name to a team-specific model name. + + Returns: + - team_model_name: str - the team-specific model name + - None: if no team-specific model name is found + """ + for model in self.model_list: + model_team_id = model["model_info"].get("team_id") + model_team_public_model_name = model["model_info"].get( + "team_public_model_name" + ) + if ( + model_team_id == team_id + and model_team_public_model_name == team_model_name + ): + return model["model_name"] + return None + def _get_all_deployments( self, model_name: str, model_alias: Optional[str] = None ) -> List[DeploymentTypedDict]: