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

1""" 

2Policy Validator - Validates policy configurations. 

3 

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""" 

11 

12import asyncio 

13from typing import TYPE_CHECKING, Any, Final, Optional 

14 

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) 

27 

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 

31 

32 

33class PolicyValidator: 

34 """ 

35 Validates policy configurations against actual data. 

36 """ 

37 

38 def __init__( 

39 self, 

40 prisma_client: Optional["PrismaClient"] = None, 

41 llm_router: Optional["Router"] = None, 

42 ): 

43 """ 

44 Initialize the validator. 

45 

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 

52 

53 @staticmethod 

54 def is_wildcard_pattern(pattern: str) -> bool: 

55 """ 

56 Check if a pattern contains wildcards. 

57 

58 Args: 

59 pattern: The pattern to check 

60 

61 Returns: 

62 True if the pattern contains wildcard characters 

63 """ 

64 return "*" in pattern or "?" in pattern 

65 

66 def get_available_guardrails(self) -> set[str]: 

67 """ 

68 Get set of available guardrail names from the guardrail registry. 

69 

70 Returns: 

71 Set of guardrail names 

72 """ 

73 try: 

74 from litellm.proxy.guardrails.guardrail_registry import ( 

75 IN_MEMORY_GUARDRAIL_HANDLER, 

76 ) 

77 

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() 

83 

84 async def check_team_alias_exists(self, team_alias: str) -> bool: 

85 """ 

86 Check if a specific team alias exists in the database. 

87 

88 Args: 

89 team_alias: The team alias to check 

90 

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 

96 

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 

105 

106 async def check_key_alias_exists(self, key_alias: str) -> bool: 

107 """ 

108 Check if a specific key alias exists in the database. 

109 

110 Args: 

111 key_alias: The key alias to check 

112 

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 

118 

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 

127 

128 def check_model_exists(self, model: str) -> bool: 

129 """ 

130 Check if a model exists in the router or matches a wildcard pattern. 

131 

132 Args: 

133 model: The model name to check 

134 

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 

140 

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 

145 

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 

151 

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 

156 

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 ) 

175 

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. 

185 

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)] 

200 

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)) 

203 

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 ] 

221 

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. 

231 

232 Checks for: 

233 - Parent policy exists 

234 - No circular inheritance 

235 - Max depth not exceeded 

236 

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 

242 

243 Returns: 

244 List of validation errors 

245 """ 

246 errors: Final[list[PolicyValidationError]] = [] 

247 

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 

259 

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() 

262 

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 

273 

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 

277 

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)) 

293 

294 return errors 

295 

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. 

303 

304 Args: 

305 policies: Dictionary mapping policy names to Policy objects 

306 validate_db: Whether to validate against database (teams, keys) 

307 

308 Returns: 

309 PolicyValidationResponse with errors and warnings 

310 """ 

311 errors: Final[list[PolicyValidationError]] = [] 

312 warnings: Final[list[PolicyValidationError]] = [] 

313 

314 # Get available guardrails 

315 available_guardrails: Final = self.get_available_guardrails() 

316 

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 ) 

330 

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 ) 

342 

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) 

351 

352 # Validate inheritance 

353 inheritance_errors = self._validate_inheritance_chain(policy_name=policy_name, policies=policies) 

354 errors.extend(inheritance_errors) 

355 

356 return PolicyValidationResponse( 

357 valid=len(errors) == 0, 

358 errors=errors, 

359 warnings=warnings, 

360 ) 

361 

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 

373 

374 guardrails_add: Final = set(policy.guardrails.get_add()) 

375 

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 ) 

390 

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 ) 

402 

403 return errors 

404 

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. 

412 

413 This parses the config and then validates it. 

414 

415 Args: 

416 policy_config: Raw policy configuration from YAML 

417 validate_db: Whether to validate against database 

418 

419 Returns: 

420 PolicyValidationResponse with errors and warnings 

421 """ 

422 from litellm.proxy.policy_engine.policy_registry import PolicyRegistry 

423 

424 # First, try to parse the policies 

425 errors: Final[list[PolicyValidationError]] = [] 

426 policies: Final[dict[str, Policy]] = {} 

427 

428 temp_registry: Final = PolicyRegistry() 

429 

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 ) 

442 

443 # If there were parsing errors, return early 

444 if errors: 

445 return PolicyValidationResponse( 

446 valid=False, 

447 errors=errors, 

448 warnings=[], 

449 ) 

450 

451 # Validate the parsed policies 

452 return await self.validate_policies(policies, validate_db=validate_db)