Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/policy_engine/policy_registry.py: 64%

374 statements  

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

1""" 

2Policy Registry - In-memory storage for policies. 

3 

4Handles storing, retrieving, and managing policies. 

5 

6Policies define WHAT guardrails to apply. WHERE they apply is defined 

7by policy_attachments (see AttachmentRegistry). 

8""" 

9 

10import json 

11from collections.abc import Mapping, Sequence 

12from datetime import datetime, timezone 

13from typing import ( 

14 TYPE_CHECKING, 

15 Any, 

16 Final, 

17 Literal, 

18 Optional, 

19 Protocol, 

20 TypedDict, 

21 Union, 

22 cast, # noqa: TID251 # prisma types the condition/pipeline Json columns as str, but reads return decoded values 

23) 

24 

25from litellm._logging import verbose_proxy_logger 

26from litellm.repositories.prisma_protocols import TableActions 

27from litellm.repositories.table_repositories import PolicyRepository 

28from litellm.types.proxy.policy_engine import ( 

29 GuardrailPipeline, 

30 PipelineStep, 

31 Policy, 

32 PolicyCondition, 

33 PolicyCreateRequest, 

34 PolicyDBResponse, 

35 PolicyGuardrails, 

36 PolicyUpdateRequest, 

37 PolicyVersionCompareResponse, 

38 PolicyVersionListResponse, 

39) 

40 

41if TYPE_CHECKING: 41 ↛ 42line 41 didn't jump to line 42 because the condition on line 41 was never true

42 from litellm.proxy.utils import PrismaClient 

43 

44# Prefix for policy version IDs in request body. Use policy_<uuid> to execute a specific version. 

45POLICY_VERSION_ID_PREFIX: Final = "policy_" 

46 

47 

48class _RawPipelineStep(TypedDict): 

49 guardrail: str 

50 

51 

52class _RawPipelineConfig(TypedDict, total=False): 

53 mode: str 

54 steps: Sequence[Union[PipelineStep, "_RawPipelineStep"]] 

55 

56 

57class _PolicyRow(Protocol): 

58 policy_id: str 

59 policy_name: str 

60 version_number: int 

61 version_status: str 

62 parent_version_id: str | None 

63 is_latest: bool 

64 published_at: datetime | None 

65 production_at: datetime | None 

66 inherit: str | None 

67 description: str | None 

68 guardrails_add: list[str] | None 

69 guardrails_remove: list[str] | None 

70 condition: dict[str, object] | None 

71 pipeline: dict[str, object] | None 

72 created_at: datetime 

73 updated_at: datetime 

74 created_by: str | None 

75 updated_by: str | None 

76 

77 

78class _PolicyVersionSourceRow(Protocol): 

79 @property 

80 def policy_id(self) -> str: ... 80 ↛ exitline 80 didn't return from function 'policy_id' because

81 

82 @property 

83 def policy_name(self) -> str: ... 83 ↛ exitline 83 didn't return from function 'policy_name' because

84 

85 @property 

86 def version_number(self) -> int: ... 86 ↛ exitline 86 didn't return from function 'version_number' because

87 

88 @property 

89 def inherit(self) -> str | None: ... 89 ↛ exitline 89 didn't return from function 'inherit' because

90 

91 @property 

92 def description(self) -> str | None: ... 92 ↛ exitline 92 didn't return from function 'description' because

93 

94 @property 

95 def guardrails_add(self) -> Sequence[str] | None: ... 95 ↛ exitline 95 didn't return from function 'guardrails_add' because

96 

97 @property 

98 def guardrails_remove(self) -> Sequence[str] | None: ... 98 ↛ exitline 98 didn't return from function 'guardrails_remove' because

99 

100 @property 

101 def condition(self) -> Mapping[str, object] | str | None: ... 101 ↛ exitline 101 didn't return from function 'condition' because

102 

103 @property 

104 def pipeline(self) -> Mapping[str, object] | str | None: ... 104 ↛ exitline 104 didn't return from function 'pipeline' because

105 

106 

107class _PolicyTableClient(Protocol): 

108 async def create(self, data: Mapping[str, object]) -> _PolicyRow: ... 108 ↛ exitline 108 didn't return from function 'create' because

109 

110 async def find_unique(self, where: Mapping[str, object]) -> _PolicyRow | None: ... 110 ↛ exitline 110 didn't return from function 'find_unique' because

111 

112 async def find_many( 112 ↛ exitline 112 didn't return from function 'find_many' because

113 self, 

114 where: Mapping[str, object] | None = None, 

115 order: Mapping[str, str] | None = None, 

116 ) -> Sequence[_PolicyRow]: ... 

117 

118 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _PolicyRow: ... 118 ↛ exitline 118 didn't return from function 'update' because

119 

120 async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... 120 ↛ exitline 120 didn't return from function 'update_many' because

121 

122 async def delete(self, where: Mapping[str, object]) -> _PolicyRow | None: ... 122 ↛ exitline 122 didn't return from function 'delete' because

123 

124 async def delete_many(self, where: Mapping[str, object]) -> int: ... 124 ↛ exitline 124 didn't return from function 'delete_many' because

125 

126 

127def _policy_table(prisma_client: "PrismaClient") -> _PolicyTableClient: 

128 table: Final = PolicyRepository(prisma_client).table 

129 return cast( # cast-ok: prisma types Json columns as str; the client hands back the decoded condition/pipeline 

130 "_PolicyTableClient", table 

131 ) 

132 

133 

134def _policy_version_source_table(prisma_client: "PrismaClient") -> "TableActions[_PolicyVersionSourceRow]": 

135 table: Final[TableActions[_PolicyVersionSourceRow]] = PolicyRepository(prisma_client).table 

136 return table 

137 

138 

139def _row_to_policy_db_response(row: _PolicyRow) -> PolicyDBResponse: 

140 """Build PolicyDBResponse from a Prisma LiteLLM_PolicyTable row.""" 

141 return PolicyDBResponse( 

142 policy_id=row.policy_id, 

143 policy_name=row.policy_name, 

144 version_number=getattr(row, "version_number", 1), 

145 version_status=getattr(row, "version_status", "production"), 

146 parent_version_id=getattr(row, "parent_version_id", None), 

147 is_latest=getattr(row, "is_latest", True), 

148 published_at=getattr(row, "published_at", None), 

149 production_at=getattr(row, "production_at", None), 

150 inherit=row.inherit, 

151 description=row.description, 

152 guardrails_add=row.guardrails_add or [], 

153 guardrails_remove=row.guardrails_remove or [], 

154 condition=row.condition, 

155 pipeline=row.pipeline, 

156 created_at=row.created_at, 

157 updated_at=row.updated_at, 

158 created_by=row.created_by, 

159 updated_by=row.updated_by, 

160 ) 

161 

162 

163class PolicyRegistry: 

164 """ 

165 In-memory registry for storing and managing policies. 

166 

167 This is a singleton that holds all loaded policies and provides 

168 methods to access them. 

169 

170 Policies define WHAT guardrails to apply: 

171 - Base guardrails via guardrails.add/remove 

172 - Inheritance via inherit field 

173 - Conditional guardrails via condition.model 

174 """ 

175 

176 def __init__(self): 

177 self._policies: dict[str, Policy] = {} 

178 self._config_policies: Mapping[str, Policy] = {} 

179 self._sources: Mapping[str, Literal["db", "config"]] = {} 

180 self._policies_by_id: dict[str, tuple[str, Policy]] = {} 

181 self._initialized: bool = False 

182 

183 def load_policies(self, policies_config: Mapping[str, dict[str, object]]) -> None: 

184 """ 

185 Load policies from a configuration dictionary. 

186 

187 Args: 

188 policies_config: Dictionary mapping policy names to policy definitions. 

189 This is the raw config from the YAML file. 

190 """ 

191 self._policies = {} 

192 self._config_policies = {} 

193 self._sources = {} 

194 self._policies_by_id = {} 

195 

196 for policy_name, policy_data in policies_config.items(): 

197 try: 

198 policy = self._parse_policy(policy_name, policy_data) 

199 self._policies[policy_name] = policy 

200 verbose_proxy_logger.debug("Loaded policy: %s", policy_name) 

201 except Exception as e: 

202 verbose_proxy_logger.error("Error loading policy '%s': %s", policy_name, e) 

203 raise ValueError(f"Invalid policy '{policy_name}': {e}") from e 

204 

205 self._config_policies = dict(self._policies) 

206 self._sources = {policy_name: "config" for policy_name in self._policies} 

207 self._initialized = True 

208 verbose_proxy_logger.info("Loaded %s policies", len(self._policies)) 

209 

210 def _parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy: 

211 """ 

212 Parse a policy from raw configuration data. 

213 

214 Args: 

215 policy_name: Name of the policy 

216 policy_data: Raw policy configuration 

217 

218 Returns: 

219 Parsed Policy object 

220 """ 

221 # Parse guardrails 

222 guardrails_data: Final = policy_data.get("guardrails", {}) 

223 if isinstance(guardrails_data, dict): 223 ↛ 230line 223 didn't jump to line 230 because the condition on line 223 was always true

224 guardrails = PolicyGuardrails( 

225 add=guardrails_data.get("add"), 

226 remove=guardrails_data.get("remove"), 

227 ) 

228 else: 

229 # Handle legacy format where guardrails might be a list 

230 guardrails = PolicyGuardrails(add=guardrails_data if guardrails_data else None) 

231 

232 # Parse condition (simple model-based condition) 

233 condition = None 

234 condition_data: Final = policy_data.get("condition") 

235 if condition_data: 

236 condition = PolicyCondition(model=condition_data.get("model")) 

237 

238 # Parse pipeline (optional ordered guardrail execution) 

239 pipeline: Final = PolicyRegistry._parse_pipeline(policy_data.get("pipeline")) 

240 

241 return Policy( 

242 inherit=policy_data.get("inherit"), 

243 description=policy_data.get("description"), 

244 guardrails=guardrails, 

245 condition=condition, 

246 pipeline=pipeline, 

247 ) 

248 

249 @staticmethod 

250 def _parse_pipeline( 

251 pipeline_data: Optional["_RawPipelineConfig"], 

252 ) -> GuardrailPipeline | None: 

253 """Parse a pipeline configuration from raw data.""" 

254 if pipeline_data is None: 254 ↛ 257line 254 didn't jump to line 257 because the condition on line 254 was always true

255 return None 

256 

257 steps_data: Final[Sequence[PipelineStep | _RawPipelineStep]] = pipeline_data.get("steps", []) 

258 steps = [PipelineStep(**step_data) if isinstance(step_data, dict) else step_data for step_data in steps_data] 

259 

260 return GuardrailPipeline( 

261 mode=pipeline_data.get("mode", "pre_call"), 

262 steps=steps, 

263 ) 

264 

265 def get_policy(self, policy_name: str) -> Policy | None: 

266 """ 

267 Get a policy by name. 

268 

269 Args: 

270 policy_name: Name of the policy to retrieve 

271 

272 Returns: 

273 Policy object if found, None otherwise 

274 """ 

275 return self._policies.get(policy_name) 

276 

277 def get_all_policies(self) -> dict[str, Policy]: 

278 """ 

279 Get all loaded policies. 

280 

281 Returns: 

282 Dictionary mapping policy names to Policy objects 

283 """ 

284 return self._policies.copy() 

285 

286 def get_policy_names(self) -> list[str]: 

287 """ 

288 Get list of all policy names. 

289 

290 Returns: 

291 List of policy names 

292 """ 

293 return list(self._policies.keys()) 

294 

295 def has_policy(self, policy_name: str) -> bool: 

296 """ 

297 Check if a policy exists. 

298 

299 Args: 

300 policy_name: Name of the policy to check 

301 

302 Returns: 

303 True if policy exists, False otherwise 

304 """ 

305 return policy_name in self._policies 

306 

307 def is_initialized(self) -> bool: 

308 """ 

309 Check if the registry has been initialized with policies. 

310 

311 Returns: 

312 True if policies have been loaded, False otherwise 

313 """ 

314 return self._initialized 

315 

316 def clear(self) -> None: 

317 """ 

318 Clear all policies from the registry. 

319 """ 

320 self._policies = {} 

321 self._config_policies = {} 

322 self._sources = {} 

323 self._initialized = False 

324 

325 def get_source(self, policy_name: str) -> Literal["db", "config"] | None: 

326 """ 

327 Return the provenance of an in-memory policy, or None if unknown. 

328 """ 

329 return self._sources.get(policy_name) 

330 

331 def list_config_policies(self) -> Mapping[str, Policy]: 

332 """ 

333 Return the policies loaded from config.yaml, keyed by policy name. 

334 """ 

335 return dict(self._config_policies) 

336 

337 def add_policy(self, policy_name: str, policy: Policy, source: Literal["db", "config"] = "db") -> None: 

338 """ 

339 Add or update a single policy. 

340 

341 Args: 

342 policy_name: Name of the policy 

343 policy: Policy object to add 

344 source: Provenance of the policy ("db" or "config") 

345 """ 

346 self._policies[policy_name] = policy 

347 self._sources = {**self._sources, policy_name: source} 

348 if source == "config": 348 ↛ 349line 348 didn't jump to line 349 because the condition on line 348 was never true

349 self._config_policies = {**self._config_policies, policy_name: policy} 

350 self._initialized = True 

351 verbose_proxy_logger.debug("Added/updated policy: %s", policy_name) 

352 

353 def remove_policy(self, policy_name: str) -> bool: 

354 """ 

355 Remove a policy by name. If a config-defined policy shares the name, 

356 it is restored immediately instead of waiting for the next DB sync. 

357 

358 Args: 

359 policy_name: Name of the policy to remove 

360 

361 Returns: 

362 True if policy was removed, False if it didn't exist 

363 """ 

364 if policy_name not in self._policies: 

365 return False 

366 config_fallback: Final = self._config_policies.get(policy_name) 

367 if config_fallback is not None: 367 ↛ 368line 367 didn't jump to line 368 because the condition on line 367 was never true

368 self._policies[policy_name] = config_fallback 

369 self._sources = {**self._sources, policy_name: "config"} 

370 verbose_proxy_logger.debug("Removed policy: %s; restored config-defined version", policy_name) 

371 return True 

372 del self._policies[policy_name] 

373 self._sources = {name: source for name, source in self._sources.items() if name != policy_name} 

374 verbose_proxy_logger.debug("Removed policy: %s", policy_name) 

375 return True 

376 

377 # ───────────────────────────────────────────────────────────────────────── 

378 # Database CRUD Methods 

379 # ───────────────────────────────────────────────────────────────────────── 

380 

381 async def add_policy_to_db( 

382 self, 

383 policy_request: PolicyCreateRequest, 

384 prisma_client: "PrismaClient", 

385 created_by: str | None = None, 

386 ) -> PolicyDBResponse: 

387 """ 

388 Add a policy to the database. 

389 

390 Args: 

391 policy_request: The policy creation request 

392 prisma_client: The Prisma client instance 

393 created_by: User who created the policy 

394 

395 Returns: 

396 PolicyDBResponse with the created policy 

397 """ 

398 try: 

399 now: Final = datetime.now(timezone.utc) 

400 # Build data dict; new policy is v1 production 

401 data: Final[dict[str, object]] = { 

402 "policy_name": policy_request.policy_name, 

403 "version_number": 1, 

404 "version_status": "production", 

405 "is_latest": True, 

406 "production_at": now, 

407 "guardrails_add": policy_request.guardrails_add or [], 

408 "guardrails_remove": policy_request.guardrails_remove or [], 

409 "created_at": now, 

410 "updated_at": now, 

411 } 

412 

413 # Only add optional fields if they have values 

414 if policy_request.inherit is not None: 

415 data["inherit"] = policy_request.inherit 

416 if policy_request.description is not None: 

417 data["description"] = policy_request.description 

418 if created_by is not None: 418 ↛ 421line 418 didn't jump to line 421 because the condition on line 418 was always true

419 data["created_by"] = created_by 

420 data["updated_by"] = created_by 

421 if policy_request.condition is not None: 

422 data["condition"] = json.dumps(policy_request.condition.model_dump()) 

423 if policy_request.pipeline is not None: 

424 validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline) 

425 data["pipeline"] = json.dumps(validated_pipeline.model_dump()) 

426 

427 created_policy: Final = await _policy_table(prisma_client).create(data=data) 

428 

429 # Also add to in-memory registry 

430 policy: Final = self._parse_policy( 

431 policy_request.policy_name, 

432 { 

433 "inherit": policy_request.inherit, 

434 "description": policy_request.description, 

435 "guardrails": { 

436 "add": policy_request.guardrails_add, 

437 "remove": policy_request.guardrails_remove, 

438 }, 

439 "condition": (policy_request.condition.model_dump() if policy_request.condition else None), 

440 "pipeline": policy_request.pipeline, 

441 }, 

442 ) 

443 self.add_policy(policy_request.policy_name, policy) 

444 

445 return _row_to_policy_db_response(created_policy) 

446 except Exception as e: 

447 verbose_proxy_logger.exception("Error adding policy to DB: %s", e) 

448 raise Exception(f"Error adding policy to DB: {e}") 

449 

450 async def update_policy_in_db( 

451 self, 

452 policy_id: str, 

453 policy_request: PolicyUpdateRequest, 

454 prisma_client: "PrismaClient", 

455 updated_by: str | None = None, 

456 ) -> PolicyDBResponse: 

457 """ 

458 Update a policy in the database. Only draft versions can be updated. 

459 

460 Args: 

461 policy_id: The ID of the policy to update 

462 policy_request: The policy update request 

463 prisma_client: The Prisma client instance 

464 updated_by: User who updated the policy 

465 

466 Returns: 

467 PolicyDBResponse with the updated policy 

468 

469 Raises: 

470 Exception: If policy is not in draft status (only drafts are editable). 

471 """ 

472 try: 

473 existing: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id}) 

474 if existing is None: 

475 raise Exception(f"Policy with ID {policy_id} not found") 

476 version_status: Final = getattr(existing, "version_status", "production") 

477 if version_status != "draft": 

478 raise Exception(f"Only draft versions can be updated. This policy has status '{version_status}'.") 

479 

480 # Build update data - only include fields that are set 

481 update_data: Final[dict[str, object]] = { 

482 "updated_at": datetime.now(timezone.utc), 

483 "updated_by": updated_by, 

484 } 

485 

486 if policy_request.policy_name is not None: 

487 update_data["policy_name"] = policy_request.policy_name 

488 if policy_request.inherit is not None: 

489 update_data["inherit"] = policy_request.inherit 

490 if policy_request.description is not None: 

491 update_data["description"] = policy_request.description 

492 if policy_request.guardrails_add is not None: 

493 update_data["guardrails_add"] = policy_request.guardrails_add 

494 if policy_request.guardrails_remove is not None: 

495 update_data["guardrails_remove"] = policy_request.guardrails_remove 

496 if policy_request.condition is not None: 

497 update_data["condition"] = json.dumps(policy_request.condition.model_dump()) 

498 if policy_request.pipeline is not None: 

499 validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline) 

500 update_data["pipeline"] = json.dumps(validated_pipeline.model_dump()) 

501 

502 updated_policy: Final = await _policy_table(prisma_client).update( 

503 where={"policy_id": policy_id}, 

504 data=update_data, 

505 ) 

506 

507 # Do NOT update in-memory registry: drafts are not loaded into memory. 

508 

509 return _row_to_policy_db_response(updated_policy) 

510 except Exception as e: 

511 verbose_proxy_logger.exception("Error updating policy in DB: %s", e) 

512 raise Exception(f"Error updating policy in DB: {e}") 

513 

514 async def delete_policy_from_db( 

515 self, 

516 policy_id: str, 

517 prisma_client: "PrismaClient", 

518 ) -> Mapping[str, str]: 

519 """ 

520 Delete a policy version from the database. 

521 

522 If the deleted version was production, it is removed from the in-memory 

523 registry. No other version is auto-promoted; admin must explicitly promote. 

524 

525 Args: 

526 policy_id: The ID of the policy version to delete 

527 prisma_client: The Prisma client instance 

528 

529 Returns: 

530 Dict with "message" and optional "warning" if production was deleted. 

531 """ 

532 try: 

533 policy: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id}) 

534 

535 if policy is None: 535 ↛ 536line 535 didn't jump to line 536 because the condition on line 535 was never true

536 raise Exception(f"Policy with ID {policy_id} not found") 

537 

538 version_status: Final = getattr(policy, "version_status", "production") 

539 policy_name: Final = policy.policy_name 

540 

541 # Delete from DB 

542 await _policy_table(prisma_client).delete(where={"policy_id": policy_id}) 

543 

544 result: Final[dict[str, str]] = {"message": f"Policy {policy_id} deleted successfully"} 

545 

546 # Remove from in-memory registry only if this was the production version 

547 if version_status == "production": 547 ↛ 559line 547 didn't jump to line 559 because the condition on line 547 was always true

548 self.remove_policy(policy_name) 

549 if self.get_source(policy_name) == "config": 549 ↛ 550line 549 didn't jump to line 550 because the condition on line 549 was never true

550 result["warning"] = ( 

551 "Production version was deleted. The config-defined policy with the same name is active again." 

552 ) 

553 else: 

554 result["warning"] = ( 

555 "Production version was deleted. No other version was promoted. " 

556 "Promote another version to production if this policy should remain active." 

557 ) 

558 

559 return result 

560 except Exception as e: 

561 verbose_proxy_logger.exception("Error deleting policy from DB: %s", e) 

562 raise Exception(f"Error deleting policy from DB: {e}") 

563 

564 async def get_policy_by_id_from_db( 

565 self, 

566 policy_id: str, 

567 prisma_client: "PrismaClient", 

568 ) -> PolicyDBResponse | None: 

569 """ 

570 Get a policy by ID from the database. 

571 

572 Args: 

573 policy_id: The ID of the policy to retrieve 

574 prisma_client: The Prisma client instance 

575 

576 Returns: 

577 PolicyDBResponse if found, None otherwise 

578 """ 

579 try: 

580 policy: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id}) 

581 

582 if policy is None: 

583 return None 

584 

585 return _row_to_policy_db_response(policy) 

586 except Exception as e: 

587 verbose_proxy_logger.exception("Error getting policy from DB: %s", e) 

588 raise Exception(f"Error getting policy from DB: {e}") 

589 

590 def get_policy_by_id_for_request(self, policy_id: str) -> tuple[str, Policy] | None: 

591 """ 

592 Return a policy version by ID from in-memory cache (no DB access). 

593 

594 Used when the request body specifies policy_<uuid> to execute a specific version 

595 (e.g. published or draft). The cache is populated by sync_policies_from_db, 

596 which loads draft and published versions keyed by policy_id. 

597 

598 Args: 

599 policy_id: The policy version ID (raw UUID, no prefix) 

600 

601 Returns: 

602 (policy_name, Policy) if found, None otherwise 

603 """ 

604 return self._policies_by_id.get(policy_id) 

605 

606 async def get_all_policies_from_db( 

607 self, 

608 prisma_client: "PrismaClient", 

609 version_status: str | None = None, 

610 ) -> list[PolicyDBResponse]: 

611 """ 

612 Get all policies from the database, optionally filtered by version_status. 

613 

614 Args: 

615 prisma_client: The Prisma client instance 

616 version_status: If set, only return policies with this status 

617 ("draft", "published", "production"). 

618 

619 Returns: 

620 List of PolicyDBResponse objects 

621 """ 

622 try: 

623 where: Final[dict[str, str]] = {} 

624 if version_status is not None: 

625 where["version_status"] = version_status 

626 

627 policies: Final = await _policy_table(prisma_client).find_many( 

628 where=where if where else None, 

629 order={"created_at": "desc"}, 

630 ) 

631 

632 return [_row_to_policy_db_response(p) for p in policies] 

633 except Exception as e: 

634 verbose_proxy_logger.exception("Error getting policies from DB: %s", e) 

635 raise Exception(f"Error getting policies from DB: {e}") 

636 

637 async def sync_policies_from_db( 

638 self, 

639 prisma_client: "PrismaClient", 

640 ) -> None: 

641 """ 

642 Sync policies from the database to in-memory registry. 

643 - Production versions are loaded into _policies (by policy name) for resolution. 

644 - Config-loaded policies are preserved; on a name conflict the DB version wins. 

645 - Draft and published versions are loaded into _policies_by_id so request-body 

646 policy_<uuid> overrides can be resolved without DB access in the hot path. 

647 """ 

648 try: 

649 production: Final = await self.get_all_policies_from_db(prisma_client, version_status="production") 

650 db_policies: Final = { 

651 policy_response.policy_name: self._parse_policy( 

652 policy_response.policy_name, 

653 { 

654 "inherit": policy_response.inherit, 

655 "description": policy_response.description, 

656 "guardrails": { 

657 "add": policy_response.guardrails_add, 

658 "remove": policy_response.guardrails_remove, 

659 }, 

660 "condition": policy_response.condition, 

661 "pipeline": policy_response.pipeline, 

662 }, 

663 ) 

664 for policy_response in production 

665 } 

666 for policy_name in set(db_policies) & set(self._config_policies): 666 ↛ 667line 666 didn't jump to line 667 because the loop on line 666 never started

667 verbose_proxy_logger.warning( 

668 "Policy '%s' is defined in both config.yaml and the DB; the DB version takes precedence", 

669 policy_name, 

670 ) 

671 config_sources: Mapping[str, Literal["db", "config"]] = {name: "config" for name in self._config_policies} 

672 db_sources: Final[Mapping[str, Literal["db", "config"]]] = {name: "db" for name in db_policies} 

673 self._policies = {**self._config_policies, **db_policies} 

674 self._sources = {**config_sources, **db_sources} 

675 

676 self._policies_by_id = {} 

677 non_production: Final = await _policy_table(prisma_client).find_many( 

678 where={"version_status": {"in": ["draft", "published"]}}, 

679 order={"created_at": "desc"}, 

680 ) 

681 for row in non_production: 

682 policy = self._parse_policy( 

683 row.policy_name, 

684 { 

685 "inherit": row.inherit, 

686 "description": row.description, 

687 "guardrails": { 

688 "add": row.guardrails_add or [], 

689 "remove": row.guardrails_remove or [], 

690 }, 

691 "condition": row.condition, 

692 "pipeline": row.pipeline, 

693 }, 

694 ) 

695 self._policies_by_id[row.policy_id] = (row.policy_name, policy) 

696 

697 self._initialized = True 

698 verbose_proxy_logger.info( 

699 "Synced %s production policies and %s draft/published (by ID) from DB to in-memory registry (%s config-defined policies preserved)", 

700 len(production), 

701 len(non_production), 

702 len(self._config_policies), 

703 ) 

704 except Exception as e: 

705 verbose_proxy_logger.exception("Error syncing policies from DB: %s", e) 

706 raise Exception(f"Error syncing policies from DB: {e}") 

707 

708 async def resolve_guardrails_from_db( 

709 self, 

710 policy_name: str, 

711 prisma_client: "PrismaClient", 

712 ) -> list[str]: 

713 """ 

714 Resolve all guardrails for a policy from the database. 

715 

716 Uses the existing PolicyResolver to handle inheritance chain resolution. 

717 

718 Args: 

719 policy_name: Name of the policy to resolve 

720 prisma_client: The Prisma client instance 

721 

722 Returns: 

723 List of resolved guardrail names 

724 """ 

725 from litellm.proxy.policy_engine.policy_resolver import PolicyResolver 

726 

727 try: 

728 # Load only production versions so inheritance resolves against production 

729 policies: Final = await self.get_all_policies_from_db(prisma_client, version_status="production") 

730 

731 # Build a temporary in-memory map for resolution 

732 temp_policies: Final = {} 

733 for policy_response in policies: 

734 policy = self._parse_policy( 

735 policy_response.policy_name, 

736 { 

737 "inherit": policy_response.inherit, 

738 "description": policy_response.description, 

739 "guardrails": { 

740 "add": policy_response.guardrails_add, 

741 "remove": policy_response.guardrails_remove, 

742 }, 

743 "condition": policy_response.condition, 

744 "pipeline": policy_response.pipeline, 

745 }, 

746 ) 

747 temp_policies[policy_response.policy_name] = policy 

748 

749 # Use the existing PolicyResolver to resolve guardrails 

750 resolved_policy: Final = PolicyResolver.resolve_policy_guardrails( 

751 policy_name=policy_name, 

752 policies=temp_policies, 

753 context=None, # No context needed for simple resolution 

754 ) 

755 

756 return sorted(resolved_policy.guardrails) 

757 except Exception as e: 

758 verbose_proxy_logger.exception("Error resolving guardrails from DB: %s", e) 

759 raise Exception(f"Error resolving guardrails from DB: {e}") 

760 

761 async def get_versions_by_policy_name( 

762 self, 

763 policy_name: str, 

764 prisma_client: "PrismaClient", 

765 ) -> PolicyVersionListResponse: 

766 """ 

767 Get all versions of a policy by name, ordered by version_number descending. 

768 

769 Args: 

770 policy_name: Name of the policy 

771 prisma_client: The Prisma client instance 

772 

773 Returns: 

774 PolicyVersionListResponse with policy_name and list of versions 

775 """ 

776 try: 

777 rows: Final = await _policy_table(prisma_client).find_many( 

778 where={"policy_name": policy_name}, 

779 order={"version_number": "desc"}, 

780 ) 

781 versions: Final = [_row_to_policy_db_response(r) for r in rows] 

782 return PolicyVersionListResponse( 

783 policy_name=policy_name, 

784 versions=versions, 

785 total_count=len(versions), 

786 ) 

787 except Exception as e: 

788 verbose_proxy_logger.exception("Error getting versions: %s", e) 

789 raise Exception(f"Error getting versions: {e}") 

790 

791 async def create_new_version( 

792 self, 

793 policy_name: str, 

794 prisma_client: "PrismaClient", 

795 source_policy_id: str | None = None, 

796 created_by: str | None = None, 

797 ) -> PolicyDBResponse: 

798 """ 

799 Create a new draft version of a policy. Copies all fields from the source. 

800 Source is current production if source_policy_id is None. 

801 

802 Args: 

803 policy_name: Name of the policy 

804 prisma_client: The Prisma client instance 

805 source_policy_id: Policy ID to clone from; if None, use current production 

806 created_by: User who created the version 

807 

808 Returns: 

809 PolicyDBResponse for the new draft version 

810 """ 

811 try: 

812 if source_policy_id is not None: 

813 source = await _policy_version_source_table(prisma_client).find_unique( 

814 where={"policy_id": source_policy_id} 

815 ) 

816 if source is None: 816 ↛ 818line 816 didn't jump to line 818 because the condition on line 816 was always true

817 raise Exception(f"Source policy {source_policy_id} not found") 

818 if source.policy_name != policy_name: 

819 raise Exception(f"Source policy name '{source.policy_name}' does not match '{policy_name}'") 

820 else: 

821 # Find current production version for this policy_name 

822 prod: Final = await _policy_version_source_table(prisma_client).find_first( 

823 where={ 

824 "policy_name": policy_name, 

825 "version_status": "production", 

826 } 

827 ) 

828 if prod is None: 

829 raise Exception(f"No production version found for policy '{policy_name}'") 

830 source = prod 

831 

832 # Next version number 

833 latest: Final = await _policy_version_source_table(prisma_client).find_first( 

834 where={"policy_name": policy_name}, 

835 order={"version_number": "desc"}, 

836 ) 

837 next_num: Final = (latest.version_number + 1) if latest else 1 

838 

839 now: Final = datetime.now(timezone.utc) 

840 # Set is_latest=False on all existing versions for this policy_name 

841 await _policy_table(prisma_client).update_many( 

842 where={"policy_name": policy_name}, 

843 data={"is_latest": False}, 

844 ) 

845 

846 data: Final[dict[str, object]] = { 

847 "policy_name": policy_name, 

848 "version_number": next_num, 

849 "version_status": "draft", 

850 "parent_version_id": source.policy_id, 

851 "is_latest": True, 

852 "published_at": None, 

853 "production_at": None, 

854 "inherit": source.inherit, 

855 "description": source.description, 

856 "guardrails_add": source.guardrails_add or [], 

857 "guardrails_remove": source.guardrails_remove or [], 

858 "created_at": now, 

859 "updated_at": now, 

860 "created_by": created_by, 

861 "updated_by": created_by, 

862 } 

863 # Prisma expects Json fields as JSON strings on create (same as add_policy_to_db) 

864 if source.condition is not None: 864 ↛ 868line 864 didn't jump to line 868 because the condition on line 864 was always true

865 data["condition"] = ( 

866 json.dumps(source.condition) if isinstance(source.condition, dict) else source.condition 

867 ) 

868 if source.pipeline is not None: 868 ↛ 869line 868 didn't jump to line 869 because the condition on line 868 was never true

869 data["pipeline"] = json.dumps(source.pipeline) if isinstance(source.pipeline, dict) else source.pipeline 

870 

871 created: Final = await _policy_table(prisma_client).create(data=data) 

872 return _row_to_policy_db_response(created) 

873 except Exception as e: 

874 verbose_proxy_logger.exception("Error creating new version: %s", e) 

875 raise Exception(f"Error creating new version: {e}") 

876 

877 async def update_version_status( 

878 self, 

879 policy_id: str, 

880 new_status: str, 

881 prisma_client: "PrismaClient", 

882 updated_by: str | None = None, 

883 ) -> PolicyDBResponse: 

884 """ 

885 Update a policy version's status. Valid transitions: 

886 - draft -> published (sets published_at) 

887 - published -> production (sets production_at, demotes current production to published, updates in-memory) 

888 - production -> published (demotes, removes from in-memory) 

889 - draft -> production: NOT allowed (must publish first) 

890 - published -> draft: NOT allowed 

891 

892 Args: 

893 policy_id: The policy version ID 

894 new_status: "published" or "production" 

895 prisma_client: The Prisma client instance 

896 updated_by: User who updated 

897 

898 Returns: 

899 PolicyDBResponse for the updated version 

900 """ 

901 try: 

902 if new_status not in ("published", "production"): 

903 raise Exception(f"Invalid status '{new_status}'. Use 'published' or 'production'.") 

904 

905 row: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id}) 

906 if row is None: 

907 raise Exception(f"Policy with ID {policy_id} not found") 

908 

909 current: Final = getattr(row, "version_status", "production") 

910 policy_name: Final = row.policy_name 

911 now: Final = datetime.now(timezone.utc) 

912 

913 if new_status == "published": 913 ↛ 914line 913 didn't jump to line 914 because the condition on line 913 was never true

914 if current != "draft": 

915 raise Exception(f"Only draft versions can be published. Current status: '{current}'.") 

916 updated = await _policy_table(prisma_client).update( 

917 where={"policy_id": policy_id}, 

918 data={ 

919 "version_status": "published", 

920 "published_at": now, 

921 "updated_at": now, 

922 "updated_by": updated_by, 

923 }, 

924 ) 

925 return _row_to_policy_db_response(updated) 

926 

927 # new_status == "production" 

928 if current not in ("draft", "published"): 928 ↛ 933line 928 didn't jump to line 933 because the condition on line 928 was always true

929 raise Exception( 

930 f"Only draft or published versions can be promoted to production. Current: '{current}'." 

931 ) 

932 # Plan: "draft -> production" NOT allowed 

933 if current == "draft": 

934 raise Exception("Cannot promote draft directly to production. Publish the version first.") 

935 

936 # Demote current production to published 

937 await _policy_table(prisma_client).update_many( 

938 where={ 

939 "policy_name": policy_name, 

940 "version_status": "production", 

941 }, 

942 data={ 

943 "version_status": "published", 

944 "updated_at": now, 

945 "updated_by": updated_by, 

946 }, 

947 ) 

948 

949 # Promote this version to production 

950 updated = await _policy_table(prisma_client).update( 

951 where={"policy_id": policy_id}, 

952 data={ 

953 "version_status": "production", 

954 "production_at": now, 

955 "updated_at": now, 

956 "updated_by": updated_by, 

957 }, 

958 ) 

959 

960 # Update in-memory registry: remove old production (by name), add this one 

961 self.remove_policy(policy_name) 

962 policy: Final = self._parse_policy( 

963 policy_name, 

964 { 

965 "inherit": updated.inherit, 

966 "description": updated.description, 

967 "guardrails": { 

968 "add": updated.guardrails_add or [], 

969 "remove": updated.guardrails_remove or [], 

970 }, 

971 "condition": updated.condition, 

972 "pipeline": updated.pipeline, 

973 }, 

974 ) 

975 self.add_policy(policy_name, policy) 

976 

977 return _row_to_policy_db_response(updated) 

978 except Exception as e: 

979 verbose_proxy_logger.exception("Error updating version status: %s", e) 

980 raise Exception(f"Error updating version status: {e}") 

981 

982 async def compare_versions( 

983 self, 

984 policy_id_a: str, 

985 policy_id_b: str, 

986 prisma_client: "PrismaClient", 

987 ) -> PolicyVersionCompareResponse: 

988 """ 

989 Compare two policy versions and return field-by-field diffs. 

990 

991 Args: 

992 policy_id_a: First policy version ID 

993 policy_id_b: Second policy version ID 

994 prisma_client: The Prisma client instance 

995 

996 Returns: 

997 PolicyVersionCompareResponse with both versions and field_diffs 

998 """ 

999 try: 

1000 a: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id_a}) 

1001 b: Final = await _policy_table(prisma_client).find_unique(where={"policy_id": policy_id_b}) 

1002 if a is None: 1002 ↛ 1004line 1002 didn't jump to line 1004 because the condition on line 1002 was always true

1003 raise Exception(f"Policy {policy_id_a} not found") 

1004 if b is None: 

1005 raise Exception(f"Policy {policy_id_b} not found") 

1006 

1007 resp_a: Final = _row_to_policy_db_response(a) 

1008 resp_b: Final = _row_to_policy_db_response(b) 

1009 

1010 # Compare fields that are part of policy content (not metadata) 

1011 compare_fields: Final = ( 

1012 "inherit", 

1013 "description", 

1014 "guardrails_add", 

1015 "guardrails_remove", 

1016 "condition", 

1017 "pipeline", 

1018 ) 

1019 field_diffs: Final[dict[str, dict[str, object]]] = {} 

1020 for field in compare_fields: 

1021 val_a = getattr(resp_a, field) 

1022 val_b = getattr(resp_b, field) 

1023 if val_a != val_b: 

1024 field_diffs[field] = {"version_a": val_a, "version_b": val_b} 

1025 

1026 return PolicyVersionCompareResponse( 

1027 version_a=resp_a, 

1028 version_b=resp_b, 

1029 field_diffs=field_diffs, 

1030 ) 

1031 except Exception as e: 

1032 verbose_proxy_logger.exception("Error comparing versions: %s", e) 

1033 raise Exception(f"Error comparing versions: {e}") 

1034 

1035 async def delete_all_versions( 

1036 self, 

1037 policy_name: str, 

1038 prisma_client: "PrismaClient", 

1039 ) -> Mapping[str, str]: 

1040 """ 

1041 Delete all versions of a policy. Also removes from in-memory registry. 

1042 

1043 Args: 

1044 policy_name: Name of the policy 

1045 prisma_client: The Prisma client instance 

1046 

1047 Returns: 

1048 Dict with "message" and optional "warning" if a config-defined policy took over. 

1049 """ 

1050 try: 

1051 await _policy_table(prisma_client).delete_many(where={"policy_name": policy_name}) 

1052 self.remove_policy(policy_name) 

1053 message: Final = f"All versions of policy '{policy_name}' deleted successfully" 

1054 if self.get_source(policy_name) == "config": 1054 ↛ 1055line 1054 didn't jump to line 1055 because the condition on line 1054 was never true

1055 return { 

1056 "message": message, 

1057 "warning": ( 

1058 "All DB versions were deleted. The config-defined policy with the same name is active again." 

1059 ), 

1060 } 

1061 return {"message": message} 

1062 except Exception as e: 

1063 verbose_proxy_logger.exception("Error deleting all versions: %s", e) 

1064 raise Exception(f"Error deleting all versions: {e}") 

1065 

1066 

1067# Global singleton instance 

1068_policy_registry: PolicyRegistry | None = None 

1069 

1070 

1071def get_policy_registry() -> PolicyRegistry: 

1072 """ 

1073 Get the global PolicyRegistry singleton. 

1074 

1075 Returns: 

1076 The global PolicyRegistry instance 

1077 """ 

1078 global _policy_registry 

1079 if _policy_registry is None: 

1080 _policy_registry = PolicyRegistry() 

1081 return _policy_registry