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
« 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.
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"""
12import asyncio
13import inspect
14from collections.abc import Awaitable, Mapping
15from types import MappingProxyType
16from typing import Final, Literal, Protocol
18from fastapi import HTTPException, status
19from pydantic import BaseModel, JsonValue, TypeAdapter
21from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
22from litellm.types.proxy.management_endpoints.team_endpoints import (
23 TeamMetadataFieldSchema,
24)
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."
33class TeamMetadataRequester(BaseModel):
34 user_id: str | None = None
35 user_email: str | None = None
36 user_role: str | None = None
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
48class TeamMetadataValidationResult(BaseModel):
49 valid: bool
50 error_message: str | None = None
53_EMPTY_METADATA: Final[Mapping[str, JsonValue]] = MappingProxyType({})
56class TeamMetadataValidator(Protocol):
57 def __call__(self, payload: TeamMetadataValidationPayload, /) -> Awaitable[TeamMetadataValidationResult]: ... 57 ↛ exitline 57 didn't return from function '__call__' because
60class TeamMetadataValidatorRegistry:
61 def __init__(self) -> None:
62 self._validator: TeamMetadataValidator | None = None
64 def set(self, validator: TeamMetadataValidator | None) -> None:
65 self._validator = validator
67 def get(self) -> TeamMetadataValidator | None:
68 return self._validator
71TEAM_METADATA_VALIDATOR_REGISTRY: Final = TeamMetadataValidatorRegistry()
73_TEAM_METADATA_SCHEMA_ADAPTER: Final[TypeAdapter[tuple[TeamMetadataFieldSchema, ...]]] = TypeAdapter(
74 tuple[TeamMetadataFieldSchema, ...]
75)
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
90class TeamMetadataSchemaRegistry:
91 def __init__(self) -> None:
92 self._fields: tuple[TeamMetadataFieldSchema, ...] = ()
94 def set(self, fields: tuple[TeamMetadataFieldSchema, ...]) -> None:
95 self._fields = fields
97 def get(self) -> tuple[TeamMetadataFieldSchema, ...]:
98 return self._fields
101TEAM_METADATA_SCHEMA_REGISTRY: Final = TeamMetadataSchemaRegistry()
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 )
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 )
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 )
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
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
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
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
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 )