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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1# litellm/proxy/guardrails/guardrail_registry.py
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
11from pydantic import ValidationError
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)
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)
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
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
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
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}
91CONFIG_GUARDRAIL_ID_NAMESPACE: Final = uuid.UUID("625f63f4-935a-50e5-98b5-fbe77babc74a")
93GuardrailCallbacks: TypeAlias = tuple[CustomGuardrail, ...]
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}
105def get_guardrail_initializer_from_hooks():
106 """
107 Get guardrail initializers by discovering them from the guardrail_hooks directory structure.
109 Scans the guardrail_hooks directory for subdirectories containing __init__.py files
110 with either guardrail_initializer_registry or initialize_guardrail functions.
112 Returns:
113 Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions
114 """
115 discovered_initializers: Final = {}
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")
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
126 # Scan each subdirectory in guardrail_hooks
127 for item in os.listdir(hooks_dir):
128 item_path = os.path.join(hooks_dir, item)
130 # Skip files and __pycache__ directories
131 if not os.path.isdir(item_path) or item.startswith("__"):
132 continue
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
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)
144 module = importlib.import_module(module_path)
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 )
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)
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
171 verbose_proxy_logger.debug(
172 "Discovered %s guardrail initializers: %s",
173 len(discovered_initializers),
174 list(discovered_initializers.keys()),
175 )
177 except Exception as e:
178 verbose_proxy_logger.error("Error discovering guardrail initializers: %s", e)
180 return discovered_initializers
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.
190 Scans the guardrail_hooks directory for subdirectories containing __init__.py files
191 with either guardrail_initializer_registry or initialize_guardrail functions.
193 Returns:
194 Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions
195 """
196 discovered_classes: Final = {}
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")
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
207 # Scan each subdirectory in guardrail_hooks
208 for item in os.listdir(hooks_dir):
209 item_path = os.path.join(hooks_dir, item)
211 # Skip files and __pycache__ directories
212 if not os.path.isdir(item_path) or item.startswith("__"):
213 continue
215 # Check if the directory has an __init__.py file
216 init_file = os.path.join(item_path, "__init__.py")
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
221 module_path = f"litellm.proxy.guardrails.guardrail_hooks.{item}"
223 try:
224 # Import the module
225 verbose_proxy_logger.debug("Discovering guardrails in: %s", module_path)
227 module = importlib.import_module(module_path)
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)
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
244 except Exception as e:
245 verbose_proxy_logger.error("Error discovering guardrail initializers: %s", e)
247 return discovered_classes
250guardrail_class_registry.update(get_guardrail_class_from_hooks())
253# Merge with dynamically discovered guardrail initializers
254_discovered_initializers: Final = get_guardrail_initializer_from_hooks()
256guardrail_initializer_registry.update(_discovered_initializers)
259class GuardrailRegistry:
260 """
261 Registry for guardrails
263 Handles adding, removing, and getting guardrails in DB + in memory
264 """
266 def __init__(self):
267 pass
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
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", {}))
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 )
312 # Add guardrail_id to the returned guardrail object
313 guardrail_dict: Final = dict(guardrail)
314 guardrail_dict["guardrail_id"] = created_guardrail.guardrail_id
316 return guardrail_dict
317 except Exception as e:
318 raise Exception(f"Error adding guardrail to DB: {e}")
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})
328 return {"message": f"Guardrail {guardrail_id} deleted successfully"}
329 except Exception as e:
330 raise Exception(f"Error deleting guardrail from DB: {e}")
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", {}))
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}")
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}")
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 )
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))))
383 return guardrails
384 except Exception as e:
385 raise Exception(f"Error getting guardrails from DB: {e}")
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})
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
397 return Guardrail(**(dict(guardrail)))
398 except Exception as e:
399 raise Exception(f"Error getting guardrail from DB: {e}")
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 )
410 if not guardrail:
411 return None
413 return Guardrail(**(dict(guardrail)))
414 except Exception as e:
415 raise Exception(f"Error getting guardrail from DB: {e}")
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)
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,)
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)
464class InMemoryGuardrailHandler:
465 """
466 Class that handles initializing guardrails and adding them to the CallbackManager
467 """
469 def __init__(self):
470 self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {}
471 """
472 Guardrail id to Guardrail object mapping
473 """
475 self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {}
476 """
477 Guardrail id to CustomGuardrail object mapping
478 """
480 self.guardrail_id_to_sibling_callbacks: dict[str, GuardrailCallbacks] = {} # mutable-ok: per-id registry
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 """
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)
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
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]
516 litellm_params_data: Final = guardrail["litellm_params"]
517 verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
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
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
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))
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))
534 guardrail_type: Final = litellm_params.guardrail
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")
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)
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 )
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
562 return parsed_guardrail
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
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}")
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
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
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")
611 verbose_proxy_logger.debug(
612 "Initializing custom guardrail: %s",
613 guardrail_type,
614 )
616 _guardrail_class: Final[Callable[..., CustomGuardrail]] = get_instance_fn(
617 guardrail_type, config_file_path=config_file_path
618 )
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 )
626 default_on: Final = litellm_params.default_on
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 {}
636 # Remove params that are handled explicitly or are internal
637 for key in ["guardrail", "mode", "default_on"]:
638 extra_params.pop(key, None)
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)
648 return _guardrail_callback
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
668 def delete_in_memory_guardrail(self, guardrail_id: str) -> None:
669 """
670 Delete a guardrail in memory and remove from litellm callbacks.
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)
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)
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())
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)
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)
705 def list_config_guardrails(self) -> list[Guardrail]:
706 """
707 List in-memory guardrails owned by config.yaml.
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"]
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.
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)
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.
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
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
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
784 # Compare guardrail_name
785 if existing.get("guardrail_name") != new_guardrail.get("guardrail_name"):
786 return True
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"))
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}}
804 # Log differences if any found
805 if changed_fields:
806 verbose_proxy_logger.debug("Guardrail params changed. Differences: %s", changed_fields)
808 # Return True if any fields changed
809 return len(changed_fields) > 0
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.
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
834 previous_guardrail: Final = self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
835 previous_source: Final = self._sources.get(guardrail_id, source)
837 if guardrail_id in self.IN_MEMORY_GUARDRAILS:
838 self.delete_in_memory_guardrail(guardrail_id)
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
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
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 )
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)
888########################################################
889# In Memory Guardrail Handler for LiteLLM Proxy
890########################################################
891IN_MEMORY_GUARDRAIL_HANDLER: Final = InMemoryGuardrailHandler()
893GUARDRAIL_RECONCILE_LOCK: Final = asyncio.Lock()
894########################################################