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

1from abc import abstractmethod 

2from collections.abc import Mapping 

3from typing import Annotated, Final, Protocol 

4 

5from fastapi import HTTPException 

6from pydantic import BaseModel, BeforeValidator, ValidationError 

7 

8from litellm.repositories.prisma_protocols import TableActions 

9from litellm.types.router_weights import RouterWeights 

10 

11 

12class _StoredModel(Protocol): 

13 @property 

14 @abstractmethod 

15 def model_id(self) -> str: 

16 pass 

17 

18 

19class _ModelDb(Protocol): 

20 @property 

21 @abstractmethod 

22 def litellm_proxymodeltable(self) -> TableActions[_StoredModel]: 

23 pass 

24 

25 

26class _PrismaClient(Protocol): 

27 @property 

28 @abstractmethod 

29 def db(self) -> _ModelDb: 

30 pass 

31 

32 

33class _Router(Protocol): 

34 @abstractmethod 

35 def get_deployment(self, model_id: str) -> object | None: 

36 pass 

37 

38 

39class _RouterWeightSettings(BaseModel): 

40 weights: RouterWeights | None = None 

41 

42 

43class _RouterWeightModelInfo(BaseModel): 

44 team_id: str | None = None 

45 db_model: bool | None = None 

46 team_public_model_name: str | None = None 

47 

48 

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) 

53 

54 

55class _RouterWeightDeployment(BaseModel): 

56 model_name: str 

57 model_info: Annotated[_RouterWeightModelInfo, BeforeValidator(_router_weight_model_info)] 

58 

59 

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 ) 

88 

89 

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 )