fix(route_llm_request.py): map team model from list in route llm request (#11111)

* 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

* test: add unit tests
This commit is contained in:
Krish Dholakia
2025-05-24 08:57:47 -07:00
committed by GitHub
parent 26dd540a30
commit dc79424e6b
3 changed files with 73 additions and 2 deletions
+30 -1
View File
@@ -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()
):
+20
View File
@@ -5465,6 +5465,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]:
+23 -1
View File
@@ -298,12 +298,34 @@ async def test_router_amoderation_with_credential_name(mock_amoderation):
assert call_kwargs["model"] == "text-moderation-stable"
def test_router_test_team_model():
"""
Test that router.test_team_model returns the correct model
"""
router = litellm.Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "gpt-3.5-turbo"},
"model_info": {
"team_id": "test-team",
"team_public_model_name": "test-model",
},
},
],
)
result = router.map_team_model(team_model_name="test-model", team_id="test-team")
assert result is not None
def test_router_ignore_invalid_deployments():
"""
Test that router.ignore_invalid_deployments is set to True
"""
from litellm.types.router import Deployment
router = litellm.Router(
model_list=[
{
@@ -326,4 +348,4 @@ def test_router_ignore_invalid_deployments():
)
)
assert router.get_model_list() == []
assert router.get_model_list() == []