Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/policy_engine/policy_validator.py: 60%
131 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"""
2Policy Validator - Validates policy configurations.
4Validates:
5- Guardrail names exist in the guardrail registry
6- Non-wildcard team aliases exist in the database
7- Non-wildcard key aliases exist in the database
8- Non-wildcard model names exist in the router or match a wildcard route
9- Inheritance chains are valid (no cycles, parents exist)
10"""
12import asyncio
13from typing import TYPE_CHECKING, Any, Final, Optional
15from litellm._logging import verbose_proxy_logger
16from litellm.proxy.auth.route_checks import RouteChecks
17from litellm.repositories.team_repository import TeamRepository
18from litellm.repositories.verification_token_repository import (
19 VerificationTokenRepository,
20)
21from litellm.types.proxy.policy_engine import (
22 Policy,
23 PolicyValidationError,
24 PolicyValidationErrorType,
25 PolicyValidationResponse,
26)
28if TYPE_CHECKING: 28 ↛ 29line 28 didn't jump to line 29 because the condition on line 28 was never true
29 from litellm.proxy.utils import PrismaClient
30 from litellm.router import Router
33class PolicyValidator:
34 """
35 Validates policy configurations against actual data.
36 """
38 def __init__(
39 self,
40 prisma_client: Optional["PrismaClient"] = None,
41 llm_router: Optional["Router"] = None,
42 ):
43 """
44 Initialize the validator.
46 Args:
47 prisma_client: Optional Prisma client for database validation
48 llm_router: Optional LLM router for model validation
49 """
50 self.prisma_client = prisma_client
51 self.llm_router = llm_router
53 @staticmethod
54 def is_wildcard_pattern(pattern: str) -> bool:
55 """
56 Check if a pattern contains wildcards.
58 Args:
59 pattern: The pattern to check
61 Returns:
62 True if the pattern contains wildcard characters
63 """
64 return "*" in pattern or "?" in pattern
66 def get_available_guardrails(self) -> set[str]:
67 """
68 Get set of available guardrail names from the guardrail registry.
70 Returns:
71 Set of guardrail names
72 """
73 try:
74 from litellm.proxy.guardrails.guardrail_registry import (
75 IN_MEMORY_GUARDRAIL_HANDLER,
76 )
78 guardrails: Final = IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails()
79 return {g.get("guardrail_name", "") for g in guardrails if g.get("guardrail_name")}
80 except Exception as e:
81 verbose_proxy_logger.warning("Could not get guardrails from registry: %s", e)
82 return set()
84 async def check_team_alias_exists(self, team_alias: str) -> bool:
85 """
86 Check if a specific team alias exists in the database.
88 Args:
89 team_alias: The team alias to check
91 Returns:
92 True if the team alias exists
93 """
94 if self.prisma_client is None: 94 ↛ 95line 94 didn't jump to line 95 because the condition on line 94 was never true
95 return True # Can't validate without DB, assume valid
97 try:
98 team: Final = await TeamRepository(self.prisma_client).table.find_first(
99 where={"team_alias": team_alias},
100 )
101 return team is not None
102 except Exception as e:
103 verbose_proxy_logger.warning("Could not check team alias '%s': %s", team_alias, e)
104 return True # Assume valid on error
106 async def check_key_alias_exists(self, key_alias: str) -> bool:
107 """
108 Check if a specific key alias exists in the database.
110 Args:
111 key_alias: The key alias to check
113 Returns:
114 True if the key alias exists
115 """
116 if self.prisma_client is None: 116 ↛ 117line 116 didn't jump to line 117 because the condition on line 116 was never true
117 return True # Can't validate without DB, assume valid
119 try:
120 key: Final = await VerificationTokenRepository(self.prisma_client).table.find_first(
121 where={"key_alias": key_alias},
122 )
123 return key is not None
124 except Exception as e:
125 verbose_proxy_logger.warning("Could not check key alias '%s': %s", key_alias, e)
126 return True # Assume valid on error
128 def check_model_exists(self, model: str) -> bool:
129 """
130 Check if a model exists in the router or matches a wildcard pattern.
132 Args:
133 model: The model name to check
135 Returns:
136 True if the model exists or matches a pattern in the router
137 """
138 if self.llm_router is None: 138 ↛ 139line 138 didn't jump to line 139 because the condition on line 138 was never true
139 return True # Can't validate without router, assume valid
141 try:
142 # Check if model is in router's model names
143 if model in self.llm_router.model_names: 143 ↛ 144line 143 didn't jump to line 144 because the condition on line 143 was never true
144 return True
146 # Check if model matches any pattern via pattern router
147 if hasattr(self.llm_router, "pattern_router"): 147 ↛ 152line 147 didn't jump to line 152 because the condition on line 147 was always true
148 pattern_deployments: Final = self.llm_router.pattern_router.get_deployments_by_pattern(model=model)
149 if pattern_deployments: 149 ↛ 150line 149 didn't jump to line 150 because the condition on line 149 was never true
150 return True
152 return False
153 except Exception as e:
154 verbose_proxy_logger.warning("Could not check model '%s': %s", model, e)
155 return True # Assume valid on error
157 @staticmethod
158 def _scope_error(
159 policy_name: str,
160 error_type: PolicyValidationErrorType,
161 field: str,
162 value: str,
163 label: str,
164 ) -> PolicyValidationError:
165 return PolicyValidationError(
166 policy_name=policy_name,
167 error_type=error_type,
168 message=(
169 f"{label.capitalize()} '{value}' does not exist. Reference an existing "
170 f"{label} or use a wildcard pattern (e.g. '{value}*') to match by prefix."
171 ),
172 field=field,
173 value=value,
174 )
176 async def find_invalid_scope_entries(
177 self,
178 policy_name: str,
179 teams: list[str] | None = None,
180 keys: list[str] | None = None,
181 models: list[str] | None = None,
182 ) -> list[PolicyValidationError]:
183 """
184 Validate the concrete scope entries of a policy attachment.
186 Returns an error for every non-wildcard entry that does not resolve to an
187 existing team alias, key alias, or model. Wildcard patterns are always
188 accepted: a pattern like "healthcare-*" may match zero entities today and
189 match ones created later, so it cannot be validated by existence. Tags are
190 intentionally not checked - they are free-form labels with no registry to
191 validate against.
192 """
193 # A concrete entry is one the request-time matcher compares by exact equality;
194 # only a trailing "*" is a wildcard (RouteChecks._is_wildcard_pattern), and those
195 # are left unvalidated since they may match zero entities today and more later.
196 is_pattern: Final = RouteChecks._is_wildcard_pattern
197 concrete_teams: Final = [t for t in (teams or []) if not is_pattern(pattern=t)]
198 concrete_keys: Final = [k for k in (keys or []) if not is_pattern(pattern=k)]
199 concrete_models: Final = [m for m in (models or []) if not is_pattern(pattern=m)]
201 team_exists: Final = await asyncio.gather(*(self.check_team_alias_exists(t) for t in concrete_teams))
202 key_exists: Final = await asyncio.gather(*(self.check_key_alias_exists(k) for k in concrete_keys))
204 return [
205 *(
206 self._scope_error(policy_name, PolicyValidationErrorType.INVALID_TEAM, "teams", team, "team")
207 for team, exists in zip(concrete_teams, team_exists)
208 if not exists
209 ),
210 *(
211 self._scope_error(policy_name, PolicyValidationErrorType.INVALID_KEY, "keys", key, "key")
212 for key, exists in zip(concrete_keys, key_exists)
213 if not exists
214 ),
215 *(
216 self._scope_error(policy_name, PolicyValidationErrorType.INVALID_MODEL, "models", model, "model")
217 for model in concrete_models
218 if not self.check_model_exists(model)
219 ),
220 ]
222 def _validate_inheritance_chain(
223 self,
224 policy_name: str,
225 policies: dict[str, Policy],
226 visited: set[str] | None = None,
227 max_depth: int = 100,
228 ) -> list[PolicyValidationError]:
229 """
230 Validate the inheritance chain for a policy.
232 Checks for:
233 - Parent policy exists
234 - No circular inheritance
235 - Max depth not exceeded
237 Args:
238 policy_name: Name of the policy to validate
239 policies: All policies
240 visited: Set of already visited policy names (for cycle detection)
241 max_depth: Maximum recursion depth to prevent infinite loops
243 Returns:
244 List of validation errors
245 """
246 errors: Final[list[PolicyValidationError]] = []
248 # Prevent infinite recursion
249 if max_depth <= 0: 249 ↛ 250line 249 didn't jump to line 250 because the condition on line 249 was never true
250 errors.append(
251 PolicyValidationError(
252 policy_name=policy_name,
253 error_type=PolicyValidationErrorType.CIRCULAR_INHERITANCE,
254 message="Inheritance chain too deep (exceeded max depth of 100)",
255 field="inherit",
256 )
257 )
258 return errors
260 if visited is None: 260 ↛ 263line 260 didn't jump to line 263 because the condition on line 260 was always true
261 visited = set()
263 if policy_name in visited: 263 ↛ 264line 263 didn't jump to line 264 because the condition on line 263 was never true
264 errors.append(
265 PolicyValidationError(
266 policy_name=policy_name,
267 error_type=PolicyValidationErrorType.CIRCULAR_INHERITANCE,
268 message=f"Circular inheritance detected: {' -> '.join(visited)} -> {policy_name}",
269 field="inherit",
270 )
271 )
272 return errors
274 policy: Final = policies.get(policy_name)
275 if policy is None: 275 ↛ 276line 275 didn't jump to line 276 because the condition on line 275 was never true
276 return errors
278 if policy.inherit: 278 ↛ 279line 278 didn't jump to line 279 because the condition on line 278 was never true
279 if policy.inherit not in policies:
280 errors.append(
281 PolicyValidationError(
282 policy_name=policy_name,
283 error_type=PolicyValidationErrorType.INVALID_INHERITANCE,
284 message=f"Parent policy '{policy.inherit}' not found",
285 field="inherit",
286 value=policy.inherit,
287 )
288 )
289 else:
290 # Recursively check parent with decremented depth
291 visited.add(policy_name)
292 errors.extend(self._validate_inheritance_chain(policy.inherit, policies, visited, max_depth - 1))
294 return errors
296 async def validate_policies(
297 self,
298 policies: dict[str, Policy],
299 validate_db: bool = True,
300 ) -> PolicyValidationResponse:
301 """
302 Validate a set of policies.
304 Args:
305 policies: Dictionary mapping policy names to Policy objects
306 validate_db: Whether to validate against database (teams, keys)
308 Returns:
309 PolicyValidationResponse with errors and warnings
310 """
311 errors: Final[list[PolicyValidationError]] = []
312 warnings: Final[list[PolicyValidationError]] = []
314 # Get available guardrails
315 available_guardrails: Final = self.get_available_guardrails()
317 for policy_name, policy in policies.items():
318 # Validate guardrails
319 for guardrail in policy.guardrails.get_add(): 319 ↛ 320line 319 didn't jump to line 320 because the loop on line 319 never started
320 if available_guardrails and guardrail not in available_guardrails:
321 errors.append(
322 PolicyValidationError(
323 policy_name=policy_name,
324 error_type=PolicyValidationErrorType.INVALID_GUARDRAIL,
325 message=f"Guardrail '{guardrail}' not found in guardrail registry",
326 field="guardrails.add",
327 value=guardrail,
328 )
329 )
331 for guardrail in policy.guardrails.get_remove(): 331 ↛ 332line 331 didn't jump to line 332 because the loop on line 331 never started
332 if available_guardrails and guardrail not in available_guardrails:
333 warnings.append(
334 PolicyValidationError(
335 policy_name=policy_name,
336 error_type=PolicyValidationErrorType.INVALID_GUARDRAIL,
337 message=f"Guardrail '{guardrail}' in remove list not found in guardrail registry",
338 field="guardrails.remove",
339 value=guardrail,
340 )
341 )
343 # Validate pipeline if present
344 if policy.pipeline is not None: 344 ↛ 345line 344 didn't jump to line 345 because the condition on line 344 was never true
345 pipeline_errors = PolicyValidator._validate_pipeline(
346 policy_name=policy_name,
347 policy=policy,
348 available_guardrails=available_guardrails,
349 )
350 errors.extend(pipeline_errors)
352 # Validate inheritance
353 inheritance_errors = self._validate_inheritance_chain(policy_name=policy_name, policies=policies)
354 errors.extend(inheritance_errors)
356 return PolicyValidationResponse(
357 valid=len(errors) == 0,
358 errors=errors,
359 warnings=warnings,
360 )
362 @staticmethod
363 def _validate_pipeline(
364 policy_name: str,
365 policy: Policy,
366 available_guardrails: set[str],
367 ) -> list[PolicyValidationError]:
368 """Validate a policy's pipeline configuration."""
369 errors: Final[list[PolicyValidationError]] = []
370 pipeline: Final = policy.pipeline
371 if pipeline is None:
372 return errors
374 guardrails_add: Final = set(policy.guardrails.get_add())
376 for i, step in enumerate(pipeline.steps):
377 # Check guardrail is in policy's guardrails.add
378 if step.guardrail not in guardrails_add:
379 errors.append(
380 PolicyValidationError(
381 policy_name=policy_name,
382 error_type=PolicyValidationErrorType.INVALID_GUARDRAIL,
383 message=(
384 f"Pipeline step {i} guardrail '{step.guardrail}' is not in the policy's guardrails.add list"
385 ),
386 field="pipeline.steps",
387 value=step.guardrail,
388 )
389 )
391 # Check guardrail exists in registry
392 if available_guardrails and step.guardrail not in available_guardrails:
393 errors.append(
394 PolicyValidationError(
395 policy_name=policy_name,
396 error_type=PolicyValidationErrorType.INVALID_GUARDRAIL,
397 message=(f"Pipeline step {i} guardrail '{step.guardrail}' not found in guardrail registry"),
398 field="pipeline.steps",
399 value=step.guardrail,
400 )
401 )
403 return errors
405 async def validate_policy_config(
406 self,
407 policy_config: dict[str, Any],
408 validate_db: bool = True,
409 ) -> PolicyValidationResponse:
410 """
411 Validate a raw policy configuration dictionary.
413 This parses the config and then validates it.
415 Args:
416 policy_config: Raw policy configuration from YAML
417 validate_db: Whether to validate against database
419 Returns:
420 PolicyValidationResponse with errors and warnings
421 """
422 from litellm.proxy.policy_engine.policy_registry import PolicyRegistry
424 # First, try to parse the policies
425 errors: Final[list[PolicyValidationError]] = []
426 policies: Final[dict[str, Policy]] = {}
428 temp_registry: Final = PolicyRegistry()
430 for policy_name, policy_data in policy_config.items():
431 try:
432 policy = temp_registry._parse_policy(policy_name, policy_data)
433 policies[policy_name] = policy
434 except Exception as e:
435 errors.append(
436 PolicyValidationError(
437 policy_name=policy_name,
438 error_type=PolicyValidationErrorType.INVALID_SYNTAX,
439 message=f"Failed to parse policy: {e}",
440 )
441 )
443 # If there were parsing errors, return early
444 if errors:
445 return PolicyValidationResponse(
446 valid=False,
447 errors=errors,
448 warnings=[],
449 )
451 # Validate the parsed policies
452 return await self.validate_policies(policies, validate_db=validate_db)