Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_registry.py: 43%

383 statements  

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

1# litellm/proxy/guardrails/guardrail_registry.py 

2 

3import asyncio 

4import importlib 

5import os 

6from collections.abc import Callable, Iterator, Mapping, Sequence 

7from datetime import datetime, timezone 

8from itertools import chain, count 

9from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, cast 

10 

11from pydantic import ValidationError 

12 

13import litellm 

14from litellm import Router 

15from litellm._logging import verbose_proxy_logger 

16from litellm._uuid import uuid 

17from litellm.integrations.custom_guardrail import CustomGuardrail 

18from litellm.litellm_core_utils.safe_json_dumps import safe_dumps 

19from litellm.llms.base_llm.guardrail_translation.utils import ( 

20 effective_scan_only_tool_results_for_guardrail, 

21 effective_skip_tool_message_for_guardrail, 

22) 

23from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( 

24 BedrockGuardrail, 

25) 

26from litellm.proxy.guardrails.guardrail_hooks.grayswan import ( 

27 GraySwanGuardrail, 

28) 

29from litellm.proxy.guardrails.guardrail_hooks.grayswan import ( 

30 initialize_guardrail as initialize_grayswan, 

31) 

32from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import lakeraAI_Moderation 

33from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import LakeraAIGuardrail 

34from litellm.proxy.guardrails.guardrail_hooks.presidio import ( 

35 _OPTIONAL_PresidioPIIMasking, 

36) 

37from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( 

38 ToolPermissionGuardrail, 

39) 

40from litellm.proxy.types_utils.utils import get_instance_fn 

41from litellm.proxy.utils import PrismaClient 

42from litellm.repositories.prisma_protocols import TableActions 

43from litellm.repositories.table_repositories import GuardrailsRepository 

44from litellm.secret_managers.main import get_secret 

45from litellm.types.guardrails import ( 

46 Guardrail, 

47 GuardrailEventHooks, 

48 LakeraCategoryThresholds, 

49 LitellmParams, 

50 SupportedGuardrailIntegrations, 

51) 

52 

53from .guardrail_hooks.llm_as_a_judge import ( 

54 initialize_guardrail as initialize_llm_as_a_judge, 

55) 

56from .guardrail_initializers import ( 

57 initialize_bedrock, 

58 initialize_hide_secrets, 

59 initialize_lakera, 

60 initialize_lakera_v2, 

61 initialize_presidio, 

62 initialize_tool_permission, 

63) 

64 

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

66 from prisma import models as prisma_models 

67 

68 

69class _GuardrailRowLike(Protocol): 

70 @property 

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

72 def __iter__(self) -> Iterator[tuple[str, object]]: ... 72 ↛ exitline 72 didn't return from function '__iter__' because

73 

74 

75def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_GuardrailsTable]": 

76 """Typed view of the guardrails table actions exposed by the Prisma repository.""" 

77 return GuardrailsRepository(prisma_client).table 

78 

79 

80guardrail_initializer_registry: Final = { 

81 SupportedGuardrailIntegrations.BEDROCK.value: initialize_bedrock, 

82 SupportedGuardrailIntegrations.LAKERA.value: initialize_lakera, 

83 SupportedGuardrailIntegrations.LAKERA_V2.value: initialize_lakera_v2, 

84 SupportedGuardrailIntegrations.PRESIDIO.value: initialize_presidio, 

85 SupportedGuardrailIntegrations.HIDE_SECRETS.value: initialize_hide_secrets, 

86 SupportedGuardrailIntegrations.TOOL_PERMISSION.value: initialize_tool_permission, 

87 SupportedGuardrailIntegrations.GRAYSWAN.value: initialize_grayswan, 

88 SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_llm_as_a_judge, 

89} 

90 

91CONFIG_GUARDRAIL_ID_NAMESPACE: Final = uuid.UUID("625f63f4-935a-50e5-98b5-fbe77babc74a") 

92 

93GuardrailCallbacks: TypeAlias = tuple[CustomGuardrail, ...] 

94 

95guardrail_class_registry: Final[dict[str, type[CustomGuardrail]]] = { 

96 SupportedGuardrailIntegrations.BEDROCK.value: BedrockGuardrail, 

97 SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail, 

98 SupportedGuardrailIntegrations.LAKERA.value: lakeraAI_Moderation, 

99 SupportedGuardrailIntegrations.LAKERA_V2.value: LakeraAIGuardrail, 

100 SupportedGuardrailIntegrations.PRESIDIO.value: _OPTIONAL_PresidioPIIMasking, 

101 SupportedGuardrailIntegrations.TOOL_PERMISSION.value: ToolPermissionGuardrail, 

102} 

103 

104 

105def get_guardrail_initializer_from_hooks(): 

106 """ 

107 Get guardrail initializers by discovering them from the guardrail_hooks directory structure. 

108 

109 Scans the guardrail_hooks directory for subdirectories containing __init__.py files 

110 with either guardrail_initializer_registry or initialize_guardrail functions. 

111 

112 Returns: 

113 Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions 

114 """ 

115 discovered_initializers: Final = {} 

116 

117 try: 

118 # Get the path to the guardrail_hooks directory 

119 current_dir: Final = os.path.dirname(__file__) 

120 hooks_dir: Final = os.path.join(current_dir, "guardrail_hooks") 

121 

122 if not os.path.exists(hooks_dir): 122 ↛ 123line 122 didn't jump to line 123 because the condition on line 122 was never true

123 verbose_proxy_logger.debug("guardrail_hooks directory not found") 

124 return discovered_initializers 

125 

126 # Scan each subdirectory in guardrail_hooks 

127 for item in os.listdir(hooks_dir): 

128 item_path = os.path.join(hooks_dir, item) 

129 

130 # Skip files and __pycache__ directories 

131 if not os.path.isdir(item_path) or item.startswith("__"): 

132 continue 

133 

134 # Check if the directory has an __init__.py file 

135 init_file = os.path.join(item_path, "__init__.py") 

136 if not os.path.exists(init_file): 136 ↛ 137line 136 didn't jump to line 137 because the condition on line 136 was never true

137 continue 

138 

139 module_path = f"litellm.proxy.guardrails.guardrail_hooks.{item}" 

140 try: 

141 # Import the module 

142 verbose_proxy_logger.debug("Discovering guardrails in: %s", module_path) 

143 

144 module = importlib.import_module(module_path) 

145 

146 # Check for guardrail_initializer_registry dictionary 

147 if hasattr(module, "guardrail_initializer_registry"): 

148 registry: Mapping[str, Callable[..., CustomGuardrail]] | None = getattr( 

149 module, "guardrail_initializer_registry", None 

150 ) 

151 if isinstance(registry, dict): 151 ↛ 127line 151 didn't jump to line 127 because the condition on line 151 was always true

152 discovered_initializers.update(registry) 

153 verbose_proxy_logger.debug( 

154 "Found guardrail_initializer_registry in %s: %s", module_path, list(registry.keys()) 

155 ) 

156 

157 # Check for standalone initialize_guardrail function (fallback for directory-based guardrails) 

158 elif hasattr(module, "initialize_guardrail"): 

159 # For directories with just initialize_guardrail, use the directory name as the key 

160 initialize_fn: Callable[..., CustomGuardrail] | None = getattr(module, "initialize_guardrail", None) 

161 discovered_initializers[item] = initialize_fn 

162 verbose_proxy_logger.debug("Found initialize_guardrail function in %s", module_path) 

163 

164 except ImportError as e: 

165 verbose_proxy_logger.error("Could not import %s: %s", module_path, e) 

166 continue 

167 except Exception as e: 

168 verbose_proxy_logger.error("Error processing %s: %s", module_path, e) 

169 continue 

170 

171 verbose_proxy_logger.debug( 

172 "Discovered %s guardrail initializers: %s", 

173 len(discovered_initializers), 

174 list(discovered_initializers.keys()), 

175 ) 

176 

177 except Exception as e: 

178 verbose_proxy_logger.error("Error discovering guardrail initializers: %s", e) 

179 

180 return discovered_initializers 

181 

182 

183def get_guardrail_class_from_hooks(): 

184 """ 

185 Get guardrail classes by discovering them from the guardrail_hooks directory structure. 

186 """ 

187 """ 

188 Get guardrail initializers by discovering them from the guardrail_hooks directory structure. 

189 

190 Scans the guardrail_hooks directory for subdirectories containing __init__.py files 

191 with either guardrail_initializer_registry or initialize_guardrail functions. 

192 

193 Returns: 

194 Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions 

195 """ 

196 discovered_classes: Final = {} 

197 

198 try: 

199 # Get the path to the guardrail_hooks directory 

200 current_dir: Final = os.path.dirname(__file__) 

201 hooks_dir: Final = os.path.join(current_dir, "guardrail_hooks") 

202 

203 if not os.path.exists(hooks_dir): 203 ↛ 204line 203 didn't jump to line 204 because the condition on line 203 was never true

204 verbose_proxy_logger.debug("guardrail_hooks directory not found") 

205 return discovered_classes 

206 

207 # Scan each subdirectory in guardrail_hooks 

208 for item in os.listdir(hooks_dir): 

209 item_path = os.path.join(hooks_dir, item) 

210 

211 # Skip files and __pycache__ directories 

212 if not os.path.isdir(item_path) or item.startswith("__"): 

213 continue 

214 

215 # Check if the directory has an __init__.py file 

216 init_file = os.path.join(item_path, "__init__.py") 

217 

218 if not os.path.exists(init_file): 218 ↛ 219line 218 didn't jump to line 219 because the condition on line 218 was never true

219 continue 

220 

221 module_path = f"litellm.proxy.guardrails.guardrail_hooks.{item}" 

222 

223 try: 

224 # Import the module 

225 verbose_proxy_logger.debug("Discovering guardrails in: %s", module_path) 

226 

227 module = importlib.import_module(module_path) 

228 

229 # Check for guardrail_initializer_registry dictionary 

230 if hasattr(module, "guardrail_class_registry"): 

231 registry: Mapping[str, type[CustomGuardrail]] | None = getattr( 

232 module, "guardrail_class_registry", None 

233 ) 

234 if isinstance(registry, dict): 234 ↛ 208line 234 didn't jump to line 208 because the condition on line 234 was always true

235 discovered_classes.update(registry) 

236 

237 except ImportError as e: 

238 verbose_proxy_logger.debug("Could not import %s: %s", module_path, e) 

239 continue 

240 except Exception as e: 

241 verbose_proxy_logger.exception("Error processing %s: %s", module_path, e) 

242 continue 

243 

244 except Exception as e: 

245 verbose_proxy_logger.error("Error discovering guardrail initializers: %s", e) 

246 

247 return discovered_classes 

248 

249 

250guardrail_class_registry.update(get_guardrail_class_from_hooks()) 

251 

252 

253# Merge with dynamically discovered guardrail initializers 

254_discovered_initializers: Final = get_guardrail_initializer_from_hooks() 

255 

256guardrail_initializer_registry.update(_discovered_initializers) 

257 

258 

259class GuardrailRegistry: 

260 """ 

261 Registry for guardrails 

262 

263 Handles adding, removing, and getting guardrails in DB + in memory 

264 """ 

265 

266 def __init__(self): 

267 pass 

268 

269 ########################################################### 

270 ########### In memory management helpers for guardrails ########### 

271 ############################################################ 

272 def get_initialized_guardrail_callback(self, guardrail_name: str) -> CustomGuardrail | None: 

273 """ 

274 Returns the initialized guardrail callback for a given guardrail name 

275 """ 

276 active_guardrails = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomGuardrail) 

277 for active_guardrail in active_guardrails: 277 ↛ 278line 277 didn't jump to line 278 because the loop on line 277 never started

278 if isinstance(active_guardrail, CustomGuardrail): 

279 if active_guardrail.guardrail_name == guardrail_name: 

280 return active_guardrail 

281 return None 

282 

283 ########################################################### 

284 ########### DB management helpers for guardrails ########### 

285 ############################################################ 

286 async def add_guardrail_to_db(self, guardrail: Guardrail, prisma_client: PrismaClient): 

287 """ 

288 Add a guardrail to the database 

289 """ 

290 try: 

291 guardrail_name: Final = guardrail.get("guardrail_name") 

292 # Properly serialize LitellmParams Pydantic model to dict 

293 litellm_params_obj: Final = guardrail.get("litellm_params", {}) 

294 if hasattr(litellm_params_obj, "model_dump"): 294 ↛ 297line 294 didn't jump to line 297 because the condition on line 294 was always true

295 litellm_params_dict = litellm_params_obj.model_dump() 

296 else: 

297 litellm_params_dict = dict(litellm_params_obj) if litellm_params_obj else {} 

298 litellm_params: Final[str] = safe_dumps(litellm_params_dict) 

299 guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {})) 

300 

301 # Create guardrail in DB 

302 created_guardrail: Final[_GuardrailRowLike] = await _guardrail_table(prisma_client).create( 

303 data={ 

304 "guardrail_name": guardrail_name, 

305 "litellm_params": litellm_params, 

306 "guardrail_info": guardrail_info, 

307 "created_at": datetime.now(timezone.utc), 

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

309 } 

310 ) 

311 

312 # Add guardrail_id to the returned guardrail object 

313 guardrail_dict: Final = dict(guardrail) 

314 guardrail_dict["guardrail_id"] = created_guardrail.guardrail_id 

315 

316 return guardrail_dict 

317 except Exception as e: 

318 raise Exception(f"Error adding guardrail to DB: {e}") 

319 

320 async def delete_guardrail_from_db(self, guardrail_id: str, prisma_client: PrismaClient): 

321 """ 

322 Delete a guardrail from the database 

323 """ 

324 try: 

325 # Delete from DB 

326 await _guardrail_table(prisma_client).delete(where={"guardrail_id": guardrail_id}) 

327 

328 return {"message": f"Guardrail {guardrail_id} deleted successfully"} 

329 except Exception as e: 

330 raise Exception(f"Error deleting guardrail from DB: {e}") 

331 

332 async def update_guardrail_in_db(self, guardrail_id: str, guardrail: Guardrail, prisma_client: PrismaClient): 

333 """ 

334 Update a guardrail in the database 

335 """ 

336 try: 

337 guardrail_name: Final = guardrail.get("guardrail_name") 

338 # Properly serialize LitellmParams Pydantic model to dict 

339 litellm_params_obj: Final = guardrail.get("litellm_params", {}) 

340 if hasattr(litellm_params_obj, "model_dump"): 

341 litellm_params_dict = litellm_params_obj.model_dump() 

342 else: 

343 litellm_params_dict = dict(litellm_params_obj) if litellm_params_obj else {} 

344 litellm_params: Final[str] = safe_dumps(litellm_params_dict) 

345 guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {})) 

346 

347 # Update in DB 

348 updated_guardrail: Final[_GuardrailRowLike | None] = await _guardrail_table(prisma_client).update( 

349 where={"guardrail_id": guardrail_id}, 

350 data={ 

351 "guardrail_name": guardrail_name, 

352 "litellm_params": litellm_params, 

353 "guardrail_info": guardrail_info, 

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

355 }, 

356 ) 

357 if updated_guardrail is None: 

358 raise ValueError(f"Guardrail not found, passed guardrail_id={guardrail_id}") 

359 

360 # Convert to dict and return 

361 return dict(updated_guardrail) 

362 except Exception as e: 

363 raise Exception(f"Error updating guardrail in DB: {e}") 

364 

365 @staticmethod 

366 async def get_all_guardrails_from_db( 

367 prisma_client: PrismaClient, 

368 ) -> list[Guardrail]: 

369 """ 

370 Get all active guardrails from the database. 

371 Only rows with status == "active" are returned (pending_review and rejected are excluded). 

372 """ 

373 try: 

374 guardrails_from_db: Final = await _guardrail_table(prisma_client).find_many( 

375 where={"status": "active"}, 

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

377 ) 

378 

379 guardrails: Final[list[Guardrail]] = [] 

380 for guardrail in guardrails_from_db: 380 ↛ 381line 380 didn't jump to line 381 because the loop on line 380 never started

381 guardrails.append(Guardrail(**(dict(guardrail)))) 

382 

383 return guardrails 

384 except Exception as e: 

385 raise Exception(f"Error getting guardrails from DB: {e}") 

386 

387 async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Guardrail | None: 

388 """ 

389 Get a guardrail by its ID from the database 

390 """ 

391 try: 

392 guardrail: Final = await _guardrail_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id}) 

393 

394 if not guardrail: 394 ↛ 397line 394 didn't jump to line 397 because the condition on line 394 was always true

395 return None 

396 

397 return Guardrail(**(dict(guardrail))) 

398 except Exception as e: 

399 raise Exception(f"Error getting guardrail from DB: {e}") 

400 

401 async def get_guardrail_by_name_from_db(self, guardrail_name: str, prisma_client: PrismaClient) -> Guardrail | None: 

402 """ 

403 Get a guardrail by its name from the database 

404 """ 

405 try: 

406 guardrail: Final = await _guardrail_table(prisma_client).find_unique( 

407 where={"guardrail_name": guardrail_name} 

408 ) 

409 

410 if not guardrail: 

411 return None 

412 

413 return Guardrail(**(dict(guardrail))) 

414 except Exception as e: 

415 raise Exception(f"Error getting guardrail from DB: {e}") 

416 

417 

418def _apply_configured_bool_overrides(instance: CustomGuardrail, litellm_params: LitellmParams) -> None: 

419 """Override the parallel/raw-scan flags only when ``litellm_params`` explicitly 

420 sets them, preserving whatever default the guardrail's own constructor chose 

421 otherwise (its constructor default may be True, so blindly copying an 

422 absent/None config value would silently clobber it back to False).""" 

423 if litellm_params.run_in_parallel is not None: 

424 instance.run_in_parallel = bool(litellm_params.run_in_parallel) 

425 if litellm_params.scan_raw_request is not None: 

426 instance.scan_raw_request = bool(litellm_params.scan_raw_request) 

427 

428 

429def _as_callback_tuple( 

430 initialized: CustomGuardrail | Sequence[CustomGuardrail] | None, 

431) -> GuardrailCallbacks: 

432 if initialized is None: 

433 return () 

434 if isinstance(initialized, (list, tuple)): 

435 return tuple(initialized) 

436 return (initialized,) 

437 

438 

439def _configure_callback_scoping( 

440 custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams 

441) -> None: 

442 for scoping_param in ( 

443 "skip_system_message_in_guardrail", 

444 "skip_tool_message_in_guardrail", 

445 "scan_only_tool_results", 

446 ): 

447 setattr(custom_guardrail_callback, scoping_param, getattr(litellm_params, scoping_param, None)) 

448 scan_only_tool_results_enabled: Final = effective_scan_only_tool_results_for_guardrail(custom_guardrail_callback) 

449 if scan_only_tool_results_enabled and not custom_guardrail_callback.supports_scan_only_tool_results(): 

450 raise ValueError( 

451 f"Guardrail {guardrail_name}: scan_only_tool_results is enabled, but this " 

452 "guardrail's role filtering never scans tool results, so no request content would ever " 

453 "be scanned. Remove scan_only_tool_results or the guardrail's role-filtering option." 

454 ) 

455 if scan_only_tool_results_enabled and effective_skip_tool_message_for_guardrail(custom_guardrail_callback): 

456 raise ValueError( 

457 f"Guardrail {guardrail_name}: scan_only_tool_results and " 

458 "skip_tool_message_in_guardrail are enabled together, which excludes every message from " 

459 "scanning, so no request content would ever be scanned. Remove one of the two." 

460 ) 

461 _apply_configured_bool_overrides(custom_guardrail_callback, litellm_params) 

462 

463 

464class InMemoryGuardrailHandler: 

465 """ 

466 Class that handles initializing guardrails and adding them to the CallbackManager 

467 """ 

468 

469 def __init__(self): 

470 self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {} 

471 """ 

472 Guardrail id to Guardrail object mapping 

473 """ 

474 

475 self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {} 

476 """ 

477 Guardrail id to CustomGuardrail object mapping 

478 """ 

479 

480 self.guardrail_id_to_sibling_callbacks: dict[str, GuardrailCallbacks] = {} # mutable-ok: per-id registry 

481 

482 self._sources: dict[str, Literal["db", "config"]] = {} 

483 """ 

484 Guardrail id to provenance marker. "db" entries are reconciled against 

485 the DB on each polling tick; "config" entries are owned by proxy_config.yaml 

486 and never deleted by reconciliation. 

487 """ 

488 

489 def _stable_guardrail_id(self, guardrail_name: str) -> str: 

490 seeds: Final = chain((guardrail_name,), (f"{guardrail_name}:{occurrence}" for occurrence in count(1))) 

491 candidate_ids: Final = (str(uuid.uuid5(CONFIG_GUARDRAIL_ID_NAMESPACE, seed.encode("utf-8"))) for seed in seeds) 

492 return next(candidate_id for candidate_id in candidate_ids if candidate_id not in self.IN_MEMORY_GUARDRAILS) 

493 

494 def initialize_guardrail( 

495 self, 

496 guardrail: Guardrail, 

497 config_file_path: str | None = None, 

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

499 source: Literal["db", "config"] = "config", 

500 ) -> Guardrail | None: 

501 """ 

502 Initialize a guardrail from a dictionary and add it to the litellm callback manager 

503 

504 Returns a Guardrail object if the guardrail is initialized successfully 

505 """ 

506 guardrail_id: Final = guardrail.get("guardrail_id") or self._stable_guardrail_id(guardrail["guardrail_name"]) 

507 guardrail["guardrail_id"] = guardrail_id 

508 if guardrail_id in self.IN_MEMORY_GUARDRAILS: 508 ↛ 509line 508 didn't jump to line 509 because the condition on line 508 was never true

509 verbose_proxy_logger.debug("guardrail_id already exists in IN_MEMORY_GUARDRAILS") 

510 # Honor the caller's source even on the early-return path so a 

511 # racing polling tick or a hot-reload of config can correct an 

512 # entry's provenance. 

513 self._sources[guardrail_id] = source 

514 return self.IN_MEMORY_GUARDRAILS[guardrail_id] 

515 

516 litellm_params_data: Final = guardrail["litellm_params"] 

517 verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) 

518 

519 if isinstance(litellm_params_data, dict): 519 ↛ 520line 519 didn't jump to line 520 because the condition on line 519 was never true

520 litellm_params = LitellmParams(**litellm_params_data) 

521 else: 

522 litellm_params = litellm_params_data 

523 

524 if "category_thresholds" in litellm_params_data and litellm_params_data["category_thresholds"]: 524 ↛ 525line 524 didn't jump to line 525 because the condition on line 524 was never true

525 lakera_category_thresholds: Final = LakeraCategoryThresholds(**litellm_params_data["category_thresholds"]) 

526 litellm_params.category_thresholds = lakera_category_thresholds 

527 

528 if litellm_params.api_key and litellm_params.api_key.startswith("os.environ/"): 528 ↛ 529line 528 didn't jump to line 529 because the condition on line 528 was never true

529 litellm_params.api_key = str(get_secret(litellm_params.api_key)) 

530 

531 if litellm_params.api_base and litellm_params.api_base.startswith("os.environ/"): 531 ↛ 532line 531 didn't jump to line 532 because the condition on line 531 was never true

532 litellm_params.api_base = str(get_secret(litellm_params.api_base)) 

533 

534 guardrail_type: Final = litellm_params.guardrail 

535 

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

537 raise ValueError("guardrail_type is required") 

538 

539 created_callbacks: Final = self._create_callbacks( 

540 guardrail=guardrail, 

541 guardrail_type=guardrail_type, 

542 litellm_params=litellm_params, 

543 config_file_path=config_file_path, 

544 llm_router=llm_router, 

545 ) 

546 for custom_guardrail_callback in created_callbacks: 

547 _configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params) 

548 

549 parsed_guardrail: Final = Guardrail( 

550 guardrail_id=guardrail.get("guardrail_id"), 

551 guardrail_name=guardrail["guardrail_name"], 

552 litellm_params=litellm_params, 

553 guardrail_info=guardrail.get("guardrail_info"), 

554 ) 

555 

556 # store references to the guardrail in memory 

557 self.IN_MEMORY_GUARDRAILS[guardrail_id] = parsed_guardrail 

558 self.guardrail_id_to_custom_guardrail[guardrail_id] = created_callbacks[0] if created_callbacks else None 

559 self.guardrail_id_to_sibling_callbacks[guardrail_id] = created_callbacks[1:] 

560 self._sources[guardrail_id] = source 

561 

562 return parsed_guardrail 

563 

564 def _create_callbacks( 

565 self, 

566 guardrail: Guardrail, 

567 guardrail_type: str, 

568 litellm_params: LitellmParams, 

569 config_file_path: str | None, 

570 llm_router: Optional["Router"], 

571 ) -> GuardrailCallbacks: 

572 initializer: Final = guardrail_initializer_registry.get(guardrail_type) 

573 if initializer: 573 ↛ 574line 573 didn't jump to line 574 because the condition on line 573 was never true

574 import inspect 

575 

576 sig: Final = inspect.signature(initializer) 

577 if "llm_router" in sig.parameters: 

578 return _as_callback_tuple(initializer(litellm_params, guardrail, llm_router)) 

579 return _as_callback_tuple(initializer(litellm_params, guardrail)) 

580 if isinstance(guardrail_type, str) and "." in guardrail_type: 580 ↛ 581line 580 didn't jump to line 581 because the condition on line 580 was never true

581 return _as_callback_tuple( 

582 self.initialize_custom_guardrail( 

583 guardrail=guardrail, 

584 guardrail_type=guardrail_type, 

585 litellm_params=litellm_params, 

586 config_file_path=config_file_path, 

587 ) 

588 ) 

589 raise ValueError(f"Unsupported guardrail: {guardrail_type}") 

590 

591 def _tracked_callbacks(self, guardrail_id: str) -> GuardrailCallbacks: 

592 primary: Final = self.guardrail_id_to_custom_guardrail.get(guardrail_id) 

593 siblings: Final = self.guardrail_id_to_sibling_callbacks.get(guardrail_id, ()) 

594 return (() if primary is None else (primary,)) + siblings 

595 

596 def initialize_custom_guardrail( 

597 self, 

598 guardrail: Guardrail, 

599 guardrail_type: str, 

600 litellm_params: LitellmParams, 

601 config_file_path: str | None = None, 

602 ) -> CustomGuardrail | None: 

603 """ 

604 Initialize a Custom Guardrail from a python file or module path 

605 

606 This initializes it by adding it to the litellm callback manager 

607 """ 

608 if not config_file_path: 

609 raise Exception("GuardrailsAIException - Please pass the config_file_path to initialize_guardrails_v2") 

610 

611 verbose_proxy_logger.debug( 

612 "Initializing custom guardrail: %s", 

613 guardrail_type, 

614 ) 

615 

616 _guardrail_class: Final[Callable[..., CustomGuardrail]] = get_instance_fn( 

617 guardrail_type, config_file_path=config_file_path 

618 ) 

619 

620 mode: Final = litellm_params.mode 

621 if mode is None: 

622 raise ValueError( 

623 f"mode is required for guardrail {guardrail_type} please set mode to one of the following: {', '.join(GuardrailEventHooks)}" 

624 ) 

625 

626 default_on: Final = litellm_params.default_on 

627 

628 # Extract additional params from litellm_params to pass to custom guardrail 

629 # This matches the behavior of other guardrail initializers (e.g., initialize_lakera) 

630 # and aligns with the documented behavior for custom guardrails 

631 if hasattr(litellm_params, "model_dump"): 

632 extra_params = litellm_params.model_dump(exclude_none=True) 

633 else: 

634 extra_params = dict(litellm_params) if litellm_params else {} 

635 

636 # Remove params that are handled explicitly or are internal 

637 for key in ["guardrail", "mode", "default_on"]: 

638 extra_params.pop(key, None) 

639 

640 _guardrail_callback: Final = _guardrail_class( 

641 guardrail_name=guardrail["guardrail_name"], 

642 event_hook=mode, 

643 default_on=default_on, 

644 **extra_params, 

645 ) 

646 litellm.logging_callback_manager.add_litellm_callback(_guardrail_callback) 

647 

648 return _guardrail_callback 

649 

650 def update_in_memory_guardrail( 

651 self, 

652 guardrail_id: str, 

653 guardrail: Guardrail, 

654 source: Literal["db", "config"] = "db", 

655 ) -> None: 

656 """ 

657 Update a guardrail in memory: a changed name or litellm_params rebuilds the 

658 live callback from the new row (fail-closed: an invalid row keeps the 

659 previous instance and raises), anything else only refreshes the stored row 

660 """ 

661 updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id}) 

662 if self._has_guardrail_params_changed(guardrail_id, updated_guardrail): 

663 self.reinitialize_guardrail(guardrail=updated_guardrail, source=source) 

664 return 

665 self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail 

666 self._sources[guardrail_id] = source 

667 

668 def delete_in_memory_guardrail(self, guardrail_id: str) -> None: 

669 """ 

670 Delete a guardrail in memory and remove from litellm callbacks. 

671 

672 The callback is purged from every callback list, not just 

673 litellm.callbacks: request handling promotes guardrail callbacks into the 

674 success/failure/async lists, so removing it from only litellm.callbacks 

675 leaves the old instance stranded in those lists on every re-initialization. 

676 """ 

677 # Remove from in-memory storage 

678 self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None) 

679 self._sources.pop(guardrail_id, None) 

680 

681 tracked_callbacks: Final = self._tracked_callbacks(guardrail_id) 

682 self.guardrail_id_to_custom_guardrail.pop(guardrail_id, None) 

683 self.guardrail_id_to_sibling_callbacks.pop(guardrail_id, None) 

684 for custom_guardrail_callback in tracked_callbacks: 

685 litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback) 

686 

687 def list_in_memory_guardrails(self) -> list[Guardrail]: 

688 """ 

689 List all guardrails in memory 

690 """ 

691 return list(self.IN_MEMORY_GUARDRAILS.values()) 

692 

693 def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None: 

694 """ 

695 Get a guardrail by its ID from memory 

696 """ 

697 return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) 

698 

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

700 """ 

701 Return the provenance of an in-memory guardrail. 

702 """ 

703 return self._sources.get(guardrail_id) 

704 

705 def list_config_guardrails(self) -> list[Guardrail]: 

706 """ 

707 List in-memory guardrails owned by config.yaml. 

708 

709 DB-sourced entries are excluded: a read surface that also queries the DB 

710 would double-count live ones, and a DB-sourced entry that's missing from 

711 the DB is stale (deleted on another pod, awaiting reconciliation here). 

712 """ 

713 return [g for gid, g in self.IN_MEMORY_GUARDRAILS.items() if self._sources.get(gid) == "config"] 

714 

715 def get_config_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None: 

716 """ 

717 Get a config-owned in-memory guardrail by its ID, or None. 

718 

719 Mirrors the fallback in get_guardrail_info: a DB-sourced in-memory entry 

720 that missed the DB lookup is stale and must not be surfaced. 

721 """ 

722 if self._sources.get(guardrail_id) != "config": 722 ↛ 724line 722 didn't jump to line 724 because the condition on line 722 was always true

723 return None 

724 return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) 

725 

726 def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]: 

727 """ 

728 Drop in-memory entries that originated from the DB but are no longer 

729 present in db_guardrail_ids. Config-loaded guardrails are never touched. 

730 

731 Called by the periodic DB polling tick so that a guardrail deleted 

732 on another pod is eventually purged from this pod's memory + callbacks. 

733 """ 

734 stale_ids: Final = [ 

735 guardrail_id 

736 for guardrail_id, source in self._sources.items() 

737 if source == "db" and guardrail_id not in db_guardrail_ids 

738 ] 

739 for guardrail_id in stale_ids: 739 ↛ 740line 739 didn't jump to line 740 because the loop on line 739 never started

740 verbose_proxy_logger.info( 

741 "Reconcile: removing stale DB-backed guardrail '%s' from memory (deleted in DB by another pod)", 

742 guardrail_id, 

743 ) 

744 self.delete_in_memory_guardrail(guardrail_id) 

745 return stale_ids 

746 

747 @staticmethod 

748 def _normalize_litellm_params_for_comparison( 

749 params: LitellmParams | Mapping[str, object] | None, 

750 ) -> Mapping[str, object] | None: 

751 """ 

752 Render litellm_params to a canonical dict so an in-memory LitellmParams and 

753 the raw dict loaded from the DB compare equal when they describe the same 

754 config. The in-memory side is a LitellmParams whose model_dump() carries 

755 every field default and coerces enums, while the DB side is the raw stored 

756 dict holding only the keys originally provided. Comparing those two shapes 

757 directly never matches, so each DB poll would re-initialize the guardrail 

758 forever; normalizing both through LitellmParams keeps the diff meaningful. 

759 """ 

760 if params is None: 

761 return None 

762 if isinstance(params, LitellmParams): 

763 return params.model_dump() 

764 if isinstance(params, dict): 

765 try: 

766 return LitellmParams(**params).model_dump() 

767 except ValidationError as e: 

768 verbose_proxy_logger.warning( 

769 "Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s", 

770 e, 

771 ) 

772 return params 

773 return params 

774 

775 def _has_guardrail_params_changed(self, guardrail_id: str, new_guardrail: Guardrail) -> bool: 

776 """ 

777 Check if guardrail params or name have changed compared to in-memory version. 

778 Returns True if params/name changed or guardrail doesn't exist in memory. 

779 """ 

780 existing: Final = self.IN_MEMORY_GUARDRAILS.get(guardrail_id) 

781 if existing is None: 

782 return True 

783 

784 # Compare guardrail_name 

785 if existing.get("guardrail_name") != new_guardrail.get("guardrail_name"): 

786 return True 

787 

788 # Compare litellm_params 

789 existing_dict: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params")) 

790 new_dict: Final = self._normalize_litellm_params_for_comparison(new_guardrail.get("litellm_params")) 

791 

792 # Compare and identify specific differences 

793 changed_fields = {} 

794 if existing_dict is not None and new_dict is not None: 

795 all_keys: Final = set(existing_dict.keys()) | set(new_dict.keys()) 

796 for key in all_keys: 

797 old_val = existing_dict.get(key) 

798 new_val = new_dict.get(key) 

799 if old_val != new_val: 

800 changed_fields[key] = {"old": old_val, "new": new_val} 

801 elif existing_dict != new_dict: 

802 changed_fields = {"litellm_params": {"old": existing_dict, "new": new_dict}} 

803 

804 # Log differences if any found 

805 if changed_fields: 

806 verbose_proxy_logger.debug("Guardrail params changed. Differences: %s", changed_fields) 

807 

808 # Return True if any fields changed 

809 return len(changed_fields) > 0 

810 

811 def reinitialize_guardrail( 

812 self, 

813 guardrail: Guardrail, 

814 config_file_path: str | None = None, 

815 source: Literal["db", "config"] = "config", 

816 ) -> Guardrail | None: 

817 """ 

818 Force re-initialization of a guardrail even if it exists in memory. 

819 Removes old callback from litellm.callbacks and creates fresh instance. 

820 

821 If the new config fails to initialize (e.g. an invalid on_flagged 

822 combination or an invalid regex), the previous instance is restored 

823 rather than left deleted, and the failure is re-raised as ValueError so 

824 every init failure reaches callers as one exception type: a caller 

825 reaching this point after already deleting the old instance would 

826 otherwise leave the guardrail providing no protection at all, not 

827 merely "still enforcing the old config." 

828 """ 

829 guardrail_id: Final = guardrail.get("guardrail_id") 

830 if not guardrail_id: 

831 verbose_proxy_logger.error("Cannot reinitialize guardrail without guardrail_id") 

832 return None 

833 

834 previous_guardrail: Final = self.IN_MEMORY_GUARDRAILS.get(guardrail_id) 

835 previous_source: Final = self._sources.get(guardrail_id, source) 

836 

837 if guardrail_id in self.IN_MEMORY_GUARDRAILS: 

838 self.delete_in_memory_guardrail(guardrail_id) 

839 

840 # Initialize fresh (will add new callback to litellm.callbacks). If the new 

841 # params are invalid (a raising guardrail __init__), restore the previous 

842 # instance instead of leaving the guardrail silently removed: a guardrail 

843 # that was enforcing must never fail open because an update was bad. 

844 try: 

845 return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source) 

846 except Exception as init_error: 

847 if previous_guardrail is not None: 

848 verbose_proxy_logger.exception( 

849 "Reinitializing guardrail %s with updated params failed; restoring the previous configuration", 

850 guardrail_id, 

851 ) 

852 try: 

853 self.initialize_guardrail( 

854 guardrail=previous_guardrail, config_file_path=config_file_path, source=previous_source 

855 ) 

856 except Exception: # noqa: BLE001 # the original failure must propagate even if the restore breaks 

857 verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id) 

858 raise ValueError(f"Guardrail initialization failed: {init_error}") from init_error 

859 

860 def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None: 

861 """ 

862 Sync a guardrail from DB - initializes if new, re-initializes if changed. 

863 This is the method to call during DB polling. 

864 """ 

865 guardrail_id: Final = guardrail.get("guardrail_id") 

866 if not guardrail_id: 

867 verbose_proxy_logger.error("Cannot sync guardrail without guardrail_id") 

868 return None 

869 

870 if self._has_guardrail_params_changed(guardrail_id, guardrail): 

871 guardrail_name: Final = guardrail.get("guardrail_name", "Unknown") 

872 verbose_proxy_logger.info( 

873 "Guardrail '%s' (ID: %s) params changed, re-initializing...", guardrail_name, guardrail_id 

874 ) 

875 return self.reinitialize_guardrail( 

876 guardrail=guardrail, 

877 config_file_path=config_file_path, 

878 source="db", 

879 ) 

880 

881 # Params unchanged but the entry is still DB-backed; make sure the 

882 # source marker reflects that even if it was previously set differently 

883 # (e.g. a config entry whose UUID later collided with a DB row). 

884 self._sources[guardrail_id] = "db" 

885 return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) 

886 

887 

888######################################################## 

889# In Memory Guardrail Handler for LiteLLM Proxy 

890######################################################## 

891IN_MEMORY_GUARDRAIL_HANDLER: Final = InMemoryGuardrailHandler() 

892 

893GUARDRAIL_RECONCILE_LOCK: Final = asyncio.Lock() 

894########################################################