Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_helpers/team_metadata_validation.py: 57%

86 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1"""Custom validation of team metadata on team create/update. 

2 

3Operators point `general_settings.custom_team_metadata_validate` at an async 

4Python function (loaded via `get_instance_fn`, like `custom_key_generate`). 

5The function receives a `TeamMetadataValidationPayload` and returns a 

6`TeamMetadataValidationResult`. The proxy awaits it before committing a team 

7write and fails closed: a rejected value surfaces the function's own message 

8(HTTP 400), while any raised exception or timeout blocks the write with a 

9generic message (HTTP 503). 

10""" 

11 

12import asyncio 

13import inspect 

14from collections.abc import Awaitable, Mapping 

15from types import MappingProxyType 

16from typing import Final, Literal, Protocol 

17 

18from fastapi import HTTPException, status 

19from pydantic import BaseModel, JsonValue, TypeAdapter 

20 

21from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth 

22from litellm.types.proxy.management_endpoints.team_endpoints import ( 

23 TeamMetadataFieldSchema, 

24) 

25 

26DEFAULT_TEAM_METADATA_VALIDATION_TIMEOUT_SECONDS: Final = 5.0 

27DEFAULT_TEAM_METADATA_VALIDATION_UNAVAILABLE_MESSAGE: Final = ( 

28 "Team metadata validation is currently unavailable, so the team was not saved. Contact your proxy admin." 

29) 

30DEFAULT_TEAM_METADATA_VALIDATION_REJECTED_MESSAGE: Final = "Team metadata failed validation." 

31 

32 

33class TeamMetadataRequester(BaseModel): 

34 user_id: str | None = None 

35 user_email: str | None = None 

36 user_role: str | None = None 

37 

38 

39class TeamMetadataValidationPayload(BaseModel): 

40 operation: Literal["create", "update"] 

41 metadata: Mapping[str, JsonValue] 

42 existing_metadata: Mapping[str, JsonValue] | None = None 

43 team_id: str | None = None 

44 team_alias: str | None = None 

45 requester: TeamMetadataRequester 

46 

47 

48class TeamMetadataValidationResult(BaseModel): 

49 valid: bool 

50 error_message: str | None = None 

51 

52 

53_EMPTY_METADATA: Final[Mapping[str, JsonValue]] = MappingProxyType({}) 

54 

55 

56class TeamMetadataValidator(Protocol): 

57 def __call__(self, payload: TeamMetadataValidationPayload, /) -> Awaitable[TeamMetadataValidationResult]: ... 57 ↛ exitline 57 didn't return from function '__call__' because

58 

59 

60class TeamMetadataValidatorRegistry: 

61 def __init__(self) -> None: 

62 self._validator: TeamMetadataValidator | None = None 

63 

64 def set(self, validator: TeamMetadataValidator | None) -> None: 

65 self._validator = validator 

66 

67 def get(self) -> TeamMetadataValidator | None: 

68 return self._validator 

69 

70 

71TEAM_METADATA_VALIDATOR_REGISTRY: Final = TeamMetadataValidatorRegistry() 

72 

73_TEAM_METADATA_SCHEMA_ADAPTER: Final[TypeAdapter[tuple[TeamMetadataFieldSchema, ...]]] = TypeAdapter( 

74 tuple[TeamMetadataFieldSchema, ...] 

75) 

76 

77 

78def parse_team_metadata_schema(raw_schema: object) -> tuple[TeamMetadataFieldSchema, ...]: 

79 """Parse ``general_settings.team_metadata_schema``; raises on a malformed schema so config load fails fast.""" 

80 if raw_schema is None: 80 ↛ 82line 80 didn't jump to line 82 because the condition on line 80 was always true

81 return () 

82 fields: Final = _TEAM_METADATA_SCHEMA_ADAPTER.validate_python(raw_schema) 

83 keys: Final = tuple(field.key for field in fields) 

84 duplicate_keys: Final = sorted(frozenset(key for key in keys if keys.count(key) > 1)) 

85 if duplicate_keys: 

86 raise ValueError(f"team_metadata_schema contains duplicate keys: {', '.join(duplicate_keys)}") 

87 return fields 

88 

89 

90class TeamMetadataSchemaRegistry: 

91 def __init__(self) -> None: 

92 self._fields: tuple[TeamMetadataFieldSchema, ...] = () 

93 

94 def set(self, fields: tuple[TeamMetadataFieldSchema, ...]) -> None: 

95 self._fields = fields 

96 

97 def get(self) -> tuple[TeamMetadataFieldSchema, ...]: 

98 return self._fields 

99 

100 

101TEAM_METADATA_SCHEMA_REGISTRY: Final = TeamMetadataSchemaRegistry() 

102 

103 

104async def run_team_metadata_validation( 

105 validator: TeamMetadataValidator, 

106 payload: TeamMetadataValidationPayload, 

107 premium_user: bool, 

108 timeout_seconds: float, 

109 unavailable_message: str, 

110) -> None: 

111 if premium_user is not True: 

112 raise HTTPException( 

113 status_code=status.HTTP_400_BAD_REQUEST, 

114 detail={ # mutable-ok: HTTPException.detail has no immutable form 

115 "error": f"custom_team_metadata_validate is an Enterprise feature. {CommonProxyErrors.not_premium_user.value}" 

116 }, 

117 ) 

118 if not inspect.iscoroutinefunction(validator): 

119 validator_call: Final = getattr(validator, "__call__", None) # noqa: B004 # value unwrap for the functor check 

120 if not inspect.iscoroutinefunction(validator_call): 

121 raise HTTPException( 

122 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, 

123 detail={ # mutable-ok: HTTPException.detail has no immutable form 

124 "error": "custom_team_metadata_validate must be an async function" 

125 }, 

126 ) 

127 

128 try: 

129 raw_result: Final = await asyncio.wait_for(validator(payload), timeout=timeout_seconds) 

130 result: Final = TeamMetadataValidationResult.model_validate(raw_result) 

131 except Exception: # noqa: BLE001 # fail closed: any validator failure must block the team write 

132 raise HTTPException( 

133 status_code=status.HTTP_503_SERVICE_UNAVAILABLE, 

134 detail={"error": unavailable_message}, # mutable-ok: HTTPException.detail has no immutable form 

135 ) 

136 

137 if not result.valid: 

138 raise HTTPException( 

139 status_code=status.HTTP_400_BAD_REQUEST, 

140 detail={ # mutable-ok: HTTPException.detail has no immutable form 

141 "error": result.error_message or DEFAULT_TEAM_METADATA_VALIDATION_REJECTED_MESSAGE 

142 }, 

143 ) 

144 

145 

146def _read_timeout_seconds(general_settings: Mapping[str, object]) -> float: 

147 raw_timeout: Final = general_settings.get("team_metadata_validation_timeout") 

148 if isinstance(raw_timeout, (int, float)) and not isinstance(raw_timeout, bool) and raw_timeout > 0: 

149 return float(raw_timeout) 

150 return DEFAULT_TEAM_METADATA_VALIDATION_TIMEOUT_SECONDS 

151 

152 

153def _read_unavailable_message(general_settings: Mapping[str, object]) -> str: 

154 raw_message: Final = general_settings.get("team_metadata_validation_error_message") 

155 if isinstance(raw_message, str) and raw_message.strip(): 

156 return raw_message 

157 return DEFAULT_TEAM_METADATA_VALIDATION_UNAVAILABLE_MESSAGE 

158 

159 

160async def validate_team_metadata_if_configured( 

161 operation: Literal["create", "update"], 

162 metadata: Mapping[str, JsonValue] | None, 

163 existing_metadata: Mapping[str, JsonValue] | None, 

164 team_id: str | None, 

165 team_alias: str | None, 

166 user_api_key_dict: UserAPIKeyAuth, 

167 registry: TeamMetadataValidatorRegistry = TEAM_METADATA_VALIDATOR_REGISTRY, 

168) -> None: 

169 from litellm.proxy.proxy_server import general_settings, premium_user 

170 

171 validator: Final = registry.get() 

172 if validator is None: 172 ↛ 175line 172 didn't jump to line 175 because the condition on line 172 was always true

173 return 

174 

175 payload: Final = TeamMetadataValidationPayload( 

176 operation=operation, 

177 metadata=metadata if isinstance(metadata, dict) else _EMPTY_METADATA, 

178 existing_metadata=existing_metadata if isinstance(existing_metadata, dict) else None, 

179 team_id=team_id, 

180 team_alias=team_alias, 

181 requester=TeamMetadataRequester( 

182 user_id=user_api_key_dict.user_id, 

183 user_email=user_api_key_dict.user_email, 

184 user_role=user_api_key_dict.user_role.value if user_api_key_dict.user_role is not None else None, 

185 ), 

186 ) 

187 await run_team_metadata_validation( 

188 validator=validator, 

189 payload=payload, 

190 premium_user=premium_user, 

191 timeout_seconds=_read_timeout_seconds(general_settings), 

192 unavailable_message=_read_unavailable_message(general_settings), 

193 )