Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/router_weights.py: 71%
63 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from abc import abstractmethod
2from collections.abc import Mapping
3from typing import Annotated, Final, Protocol
5from fastapi import HTTPException
6from pydantic import BaseModel, BeforeValidator, ValidationError
8from litellm.repositories.prisma_protocols import TableActions
9from litellm.types.router_weights import RouterWeights
12class _StoredModel(Protocol):
13 @property
14 @abstractmethod
15 def model_id(self) -> str:
16 pass
19class _ModelDb(Protocol):
20 @property
21 @abstractmethod
22 def litellm_proxymodeltable(self) -> TableActions[_StoredModel]:
23 pass
26class _PrismaClient(Protocol):
27 @property
28 @abstractmethod
29 def db(self) -> _ModelDb:
30 pass
33class _Router(Protocol):
34 @abstractmethod
35 def get_deployment(self, model_id: str) -> object | None:
36 pass
39class _RouterWeightSettings(BaseModel):
40 weights: RouterWeights | None = None
43class _RouterWeightModelInfo(BaseModel):
44 team_id: str | None = None
45 db_model: bool | None = None
46 team_public_model_name: str | None = None
49def _router_weight_model_info(value: object) -> _RouterWeightModelInfo:
50 if isinstance(value, str):
51 return _RouterWeightModelInfo.model_validate_json(value)
52 return _RouterWeightModelInfo.model_validate(value or {}, from_attributes=True)
55class _RouterWeightDeployment(BaseModel):
56 model_name: str
57 model_info: Annotated[_RouterWeightModelInfo, BeforeValidator(_router_weight_model_info)]
60def _validate_router_weight_reference(
61 model_group: str,
62 deployment_id: str,
63 team_id: str | None,
64 stored: _RouterWeightDeployment | None,
65 configured: object | None,
66) -> None:
67 reference: Final = (
68 stored
69 if stored is not None
70 else (
71 _RouterWeightDeployment.model_validate(configured, from_attributes=True) if configured is not None else None
72 )
73 )
74 if ( 74 ↛ 80line 74 didn't jump to line 80 because the condition on line 74 was always true
75 reference is None
76 or (stored is None and reference.model_info.db_model)
77 or (reference.model_info.team_id is not None and reference.model_info.team_id != team_id)
78 ):
79 raise HTTPException(status_code=400, detail=f"Unknown deployment ID in router weights: {deployment_id}")
80 canonical_group: Final = (
81 reference.model_info.team_public_model_name if reference.model_info.team_id is not None else None
82 ) or reference.model_name
83 if model_group != canonical_group:
84 raise HTTPException(
85 status_code=400,
86 detail=f"Deployment {deployment_id} does not belong to model group {model_group}",
87 )
90async def validate_router_settings_weights(
91 router_settings: BaseModel | Mapping[str, object] | None,
92 *,
93 team_id: str | None,
94 prisma_client: _PrismaClient | None,
95 llm_router: _Router | None,
96) -> None:
97 try:
98 weights: Final = (
99 _RouterWeightSettings.model_validate(router_settings, from_attributes=True).weights
100 if router_settings is not None
101 else None
102 )
103 except ValidationError:
104 raise HTTPException(
105 status_code=400,
106 detail="Invalid router weights. Replace or clear router_settings.weights.",
107 ) from None
108 if not weights:
109 return
110 deployment_ids: Final = frozenset(deployment_id for group in weights.values() for deployment_id in group)
111 if not deployment_ids: 111 ↛ 112line 111 didn't jump to line 112 because the condition on line 111 was never true
112 return
113 if prisma_client is None: 113 ↛ 114line 113 didn't jump to line 114 because the condition on line 113 was never true
114 raise HTTPException(status_code=503, detail="Database unavailable while validating router weights")
115 stored_models: Final = await prisma_client.db.litellm_proxymodeltable.find_many(
116 where={"model_id": {"in": list(deployment_ids)}}
117 )
118 stored_by_id: Final = {
119 row.model_id: _RouterWeightDeployment.model_validate(row, from_attributes=True) for row in stored_models
120 }
121 for model_group, group_weights in weights.items(): 121 ↛ exitline 121 didn't return from function 'validate_router_settings_weights' because the loop on line 121 didn't complete
122 for deployment_id in group_weights: 122 ↛ 121line 122 didn't jump to line 121 because the loop on line 122 didn't complete
123 _validate_router_weight_reference(
124 model_group,
125 deployment_id,
126 team_id,
127 stored_by_id.get(deployment_id),
128 llm_router.get_deployment(model_id=deployment_id) if llm_router is not None else None,
129 )