From 8c72342ad58ed4731022c1f2d6401b620426a6ee Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 1 Sep 2026 18:25:34 -0700 Subject: [PATCH] fix(guardrails): resync event_hook and accept raw dicts in in-memory guardrail updates --- basedpyright-code-budget.json | 18 ++--- litellm/integrations/custom_guardrail.py | 73 ++++++++++++++----- .../guardrail_hooks/azure/prompt_shield.py | 40 ++++------ .../guardrail_hooks/bedrock_guardrails.py | 5 +- .../guardrail_hooks/lakera_ai_v2.py | 25 +++---- .../model_armor/model_armor.py | 2 +- .../guardrails/guardrail_hooks/presidio.py | 47 ++++++------ .../guardrail_hooks/qualifire/qualifire.py | 9 ++- .../guardrail_hooks/tool_permission.py | 18 ++--- .../zscaler_ai_guard/zscaler_ai_guard.py | 7 +- .../proxy/guardrails/guardrail_registry.py | 16 ++-- ruff-strict-budget.json | 2 +- .../integrations/test_custom_guardrail.py | 64 ++++++++++++++++ .../guardrail_hooks/test_presidio.py | 37 +++++++++- .../guardrails/test_guardrail_registry.py | 59 +++++++++++++++ type-discipline-budget.json | 8 +- 16 files changed, 308 insertions(+), 122 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index df52069e71f..4e82e0752ca 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 14076 + "limit": 14075 }, "reportArgumentType": { "limit": 2216 @@ -9,7 +9,7 @@ "limit": 319 }, "reportAttributeAccessIssue": { - "limit": 480 + "limit": 479 }, "reportCallIssue": { "limit": 112 @@ -57,7 +57,7 @@ "limit": 5601 }, "reportMissingTypeArgument": { - "limit": 15306 + "limit": 15303 }, "reportMissingTypeStubs": { "limit": 40 @@ -99,22 +99,22 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44364 + "limit": 44360 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38350 + "limit": 38346 }, "reportUnknownParameterType": { - "limit": 19626 + "limit": 19623 }, "reportUnknownVariableType": { - "limit": 29890 + "limit": 29881 }, "reportUnnecessaryCast": { - "limit": 111 + "limit": 110 }, "reportUnnecessaryComparison": { "limit": 695 @@ -123,7 +123,7 @@ "limit": 5 }, "reportUnnecessaryIsInstance": { - "limit": 826 + "limit": 825 }, "reportUntypedBaseClass": { "limit": 0 diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index e87ac9521ae..8754d116537 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -2,10 +2,12 @@ import contextvars import hashlib import os import secrets -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from datetime import datetime from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args +from pydantic import TypeAdapter + from litellm._logging import verbose_logger from litellm.caching import DualCache from litellm.integrations.custom_logger import CustomLogger @@ -121,6 +123,18 @@ def _strict_guardrail_modes_enabled() -> bool: return True if parsed is None else parsed +def updated_litellm_param(litellm_params: "LitellmParams | Mapping[str, object]", key: str) -> object: + if isinstance(litellm_params, Mapping): + return litellm_params.get(key) + value: Final[object] = getattr(litellm_params, key, None) + return value + + +GUARDRAIL_MODE_ADAPTER: Final[TypeAdapter[GuardrailEventHooks | list[GuardrailEventHooks] | Mode]] = TypeAdapter( + GuardrailEventHooks | list[GuardrailEventHooks] | Mode +) + + def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None: """Extract session_id from request data (litellm_session_id or metadata).""" session_id = request_data.get("litellm_session_id") @@ -214,18 +228,7 @@ class CustomGuardrail(CustomLogger): self.only_scan_new_messages: bool = only_scan_new_messages if supported_event_hooks: - ## validate event_hook is in supported_event_hooks - try: - self._validate_event_hook(event_hook, supported_event_hooks) - except ValueError as validation_error: - if _strict_guardrail_modes_enabled(): - raise - verbose_logger.warning( - "%s. LITELLM_STRICT_GUARDRAIL_MODES=false; continuing " - "with unsupported event_hook. Set the env var to true " - "(default) to enforce validation and fail at startup.", - validation_error, - ) + self._validate_or_warn_event_hook(event_hook, supported_event_hooks) super().__init__(**kwargs) def render_violation_message(self, default: str, context: Mapping[str, object] | None = None) -> str: @@ -588,12 +591,12 @@ class CustomGuardrail(CustomLogger): def _validate_event_hook( self, - event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None, - supported_event_hooks: list[GuardrailEventHooks], + event_hook: GuardrailEventHooks | Sequence[GuardrailEventHooks] | Mode | None, + supported_event_hooks: Sequence[GuardrailEventHooks], ) -> None: def _validate_event_hook_list_is_in_supported_event_hooks( - event_hook: list[GuardrailEventHooks] | list[str], - supported_event_hooks: list[GuardrailEventHooks], + event_hook: Sequence[GuardrailEventHooks] | Sequence[str], + supported_event_hooks: Sequence[GuardrailEventHooks], ) -> None: for hook in event_hook: if isinstance(hook, str): @@ -622,6 +625,23 @@ class CustomGuardrail(CustomLogger): if event_hook not in supported_event_hooks: raise ValueError(f"Event hook {event_hook} is not in the supported event hooks {supported_event_hooks}") + def _validate_or_warn_event_hook( + self, + event_hook: GuardrailEventHooks | Sequence[GuardrailEventHooks] | Mode | None, + supported_event_hooks: Sequence[GuardrailEventHooks], + ) -> None: + try: + self._validate_event_hook(event_hook, supported_event_hooks) + except ValueError as validation_error: + if _strict_guardrail_modes_enabled(): + raise + verbose_logger.warning( + "%s. LITELLM_STRICT_GUARDRAIL_MODES=false; continuing " + "with unsupported event_hook. Set the env var to true " + "(default) to enforce validation and fail at startup.", + validation_error, + ) + @staticmethod def _get_admin_metadata(data: dict) -> dict: """Return merged admin-configured key and team metadata from the request data. @@ -1271,12 +1291,25 @@ class CustomGuardrail(CustomLogger): # Mask the content return content_string[:start_index] + mask_string + content_string[end_index:] - def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: """ - Update the guardrails litellm params in memory + Update the guardrails litellm params in memory, accepting either a + LitellmParams object or the raw params mapping stored in the DB, and + resync ``event_hook`` when the update carries a new ``mode``. The new + mode is validated against ``supported_event_hooks`` before any state + is mutated, so a rejected update leaves the guardrail untouched. """ - for key, value in vars(litellm_params).items(): + updated_params: Final[Mapping[str, object]] = ( + litellm_params if isinstance(litellm_params, Mapping) else vars(litellm_params) + ) + raw_mode: Final = updated_params.get("mode") + new_event_hook: Final = None if raw_mode is None else GUARDRAIL_MODE_ADAPTER.validate_python(raw_mode) + if new_event_hook is not None and self.supported_event_hooks: + self._validate_or_warn_event_hook(new_event_hook, self.supported_event_hooks) + for key, value in updated_params.items(): setattr(self, key, value) + if new_event_hook is not None: + self.event_hook = new_event_hook def get_guardrails_messages_for_call_type( self, call_type: CallTypes, data: dict | None = None diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 6e29d44662e..58bebfbdb6e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -14,6 +14,7 @@ from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, + updated_litellm_param, ) from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( AZURE_PROMPT_SHIELD_TEXT_RECORD_UNIT, @@ -61,15 +62,6 @@ def _resolved_secret_value(value: object) -> object: return value -def _updated_param(litellm_params: "LitellmParams | dict", key: str) -> object: # mutable-ok: DB dict - """Read one param from a Mapping or a pydantic object, including pydantic - extras (cost_tier / price_per_1000_text_records live there), which the base - class ``vars()`` loop never sees.""" - if isinstance(litellm_params, Mapping): - return litellm_params.get(key) - return getattr(litellm_params, key, None) - - def _resolved_cost_tier(raw: object) -> str | None: """Normalize the configured cost_tier to 'free' / 'paid' / None.""" value: Final = _resolved_secret_value(raw) @@ -270,29 +262,27 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai verbose_proxy_logger.warning("Azure Prompt Shield: No user prompt found") return None - def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | dict") -> None: # mutable-ok: DB dict + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: """Apply updated params in place, re-resolving billing and credentials. - Pricing is read via ``_updated_param`` (the values are pydantic extras, and - the immediate PUT sync hands this method the raw DB dict). Pricing and any - ``os.environ/`` credential references are validated and resolved BEFORE any - state is mutated, so an invalid update leaves the running guardrail - untouched and a raw reference never overwrites a resolved credential. + Pricing is read via ``updated_litellm_param`` (the values are pydantic + extras, and the immediate PUT sync hands this method the raw DB dict). + Pricing and any ``os.environ/`` credential references are validated and + resolved BEFORE any state is mutated, so an invalid update leaves the + running guardrail untouched and a raw reference never overwrites a + resolved credential. Both input shapes flow through the base update so + the event_hook resync applies to each. """ - cost_tier: Final = _resolved_cost_tier(_updated_param(litellm_params, "cost_tier")) - price: Final = _resolved_price(_updated_param(litellm_params, "price_per_1000_text_records"), cost_tier) + cost_tier: Final = _resolved_cost_tier(updated_litellm_param(litellm_params, "cost_tier")) + price: Final = _resolved_price(updated_litellm_param(litellm_params, "price_per_1000_text_records"), cost_tier) resolved_credentials: dict[str, object] = {} # mutable-ok: staged before mutation for cred_key in ("api_key", "api_base"): - cred_value = _updated_param(litellm_params, cred_key) + cred_value = updated_litellm_param(litellm_params, cred_key) if isinstance(cred_value, str) and cred_value.startswith("os.environ/"): resolved_credentials[cred_key] = _resolved_secret_value(cred_value) - if isinstance(litellm_params, Mapping): - for key, value in litellm_params.items(): - setattr(self, key, resolved_credentials.get(key, value)) - else: - super().update_in_memory_litellm_params(litellm_params) - for cred_key, cred_value in resolved_credentials.items(): - setattr(self, cred_key, cred_value) + super().update_in_memory_litellm_params(litellm_params) + for cred_key, cred_value in resolved_credentials.items(): + setattr(self, cred_key, cred_value) self.cost_tier = cost_tier self.price_per_1000_text_records = price diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 30526d30dc5..17ca36ea6a4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -317,9 +317,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): self.streaming_sampling_rate = streaming_params.streaming_sampling_rate self.streaming_end_of_stream_only = streaming_params.streaming_end_of_stream_only - def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: super().update_in_memory_litellm_params(litellm_params) - self._set_streaming_params(BedrockGuardrailStreamingParams.from_extras(litellm_params.model_extra)) + extras: Final = litellm_params if isinstance(litellm_params, Mapping) else litellm_params.model_extra + self._set_streaming_params(BedrockGuardrailStreamingParams.from_extras(extras)) def _streams_incrementally(self) -> bool: return not self.streaming_buffer_until_moderated and not self.mask_response_content diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index 2f98a9afbd8..9e1382feb21 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -13,6 +13,7 @@ from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( DEFAULT_ADVISORY_MESSAGE, CustomGuardrail, + updated_litellm_param, ) from litellm.llms.base_llm.guardrail_translation.utils import ( effective_skip_system_message_for_guardrail, @@ -304,7 +305,7 @@ class LakeraAIGuardrail(CustomGuardrail): breakdown=self.breakdown, ) - def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: """ The base implementation blindly ``setattr``s every field on ``litellm_params`` (including ``on_flagged``/``advisory_system_message``/``payload``/``breakdown``) @@ -313,24 +314,18 @@ class LakeraAIGuardrail(CustomGuardrail): on_flagged combinations __init__ rejects. Validate the prospective post-update state *before* mutating, so a rejected update leaves the live instance untouched instead of raising after it's already been corrupted. - - The base setattr also writes ``litellm_params.mode`` onto a new ``self.mode`` - attribute rather than the ``self.event_hook`` dispatch actually reads - (LitellmParams has no field literally named ``event_hook``), so without the - explicit sync below a hot reload that changes mode would pass validation but - keep dispatching on the stale event_hook. """ - new_event_hook: Final = litellm_params.mode or self.event_hook - prospective_payload: Final = litellm_params.payload - prospective_breakdown: Final = litellm_params.breakdown + raw_on_flagged: Final = updated_litellm_param(litellm_params, "on_flagged") + raw_advisory: Final = updated_litellm_param(litellm_params, "advisory_system_message") + raw_payload: Final = updated_litellm_param(litellm_params, "payload") + raw_breakdown: Final = updated_litellm_param(litellm_params, "breakdown") self._validate_advisory_config( - on_flagged=litellm_params.on_flagged or self.on_flagged, - advisory_system_message=litellm_params.advisory_system_message, - payload=self.payload if prospective_payload is None else prospective_payload, - breakdown=self.breakdown if prospective_breakdown is None else prospective_breakdown, + on_flagged=raw_on_flagged if isinstance(raw_on_flagged, str) and raw_on_flagged else self.on_flagged, + advisory_system_message=raw_advisory if isinstance(raw_advisory, str) else None, + payload=raw_payload if isinstance(raw_payload, bool) else self.payload, + breakdown=raw_breakdown if isinstance(raw_breakdown, bool) else self.breakdown, ) super().update_in_memory_litellm_params(litellm_params=litellm_params) - self.event_hook = new_event_hook def _validate_advisory_config( self, diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index d187b5b12e9..c96334fb2f9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -185,7 +185,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if self.optional_params.get("fail_on_error", True): raise e from None - def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: super().update_in_memory_litellm_params(litellm_params) self.sanitize_error_detail = self.sanitize_error_detail is not False diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index da51a905ae3..da2776eda7d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -11,7 +11,7 @@ import asyncio import json import threading -from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Sequence +from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Mapping, Sequence from contextlib import asynccontextmanager from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast @@ -35,6 +35,7 @@ if TYPE_CHECKING: from litellm.caching.caching import DualCache from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException from litellm.integrations.custom_guardrail import ( + GUARDRAIL_MODE_ADAPTER, CustomGuardrail, log_guardrail_information, ) @@ -530,17 +531,17 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return created @staticmethod - def _coerce_analyze_chunk_size(value: int | None) -> int: + def _coerce_analyze_chunk_size(value: object) -> int: """ Validate a configured chunk size, falling back to the default. - Non-positive values would either bypass chunking entirely or degenerate - it into per-character splits (silently disabling detection), so they are - replaced by the default; values below 4 bytes are floored to 4 and the - splitter always emits at least one character per chunk, so the chunked - path can never re-enter itself. + Non-positive or non-integer values would either bypass chunking entirely + or degenerate it into per-character splits (silently disabling + detection), so they are replaced by the default; values below 4 bytes + are floored to 4 and the splitter always emits at least one character + per chunk, so the chunked path can never re-enter itself. """ - if not value or value <= 0: + if not isinstance(value, int) or value <= 0: return DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES return max(value, 4) @@ -1628,20 +1629,24 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): inputs["texts"] = new_texts return inputs - def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: """ Update the guardrails litellm params in memory """ super().update_in_memory_litellm_params(litellm_params) - if litellm_params.pii_entities_config: - self.pii_entities_config = litellm_params.pii_entities_config - if litellm_params.presidio_score_thresholds: - self.presidio_score_thresholds = litellm_params.presidio_score_thresholds - if litellm_params.presidio_entities_deny_list: - self.presidio_entities_deny_list = litellm_params.presidio_entities_deny_list - if litellm_params.presidio_analyze_chunk_size_bytes is not None: - # Same validation as __init__: a non-positive value from a guardrail - # update must not silently disable detection via degenerate chunking. - self.presidio_analyze_chunk_size_bytes = self._coerce_analyze_chunk_size( - litellm_params.presidio_analyze_chunk_size_bytes - ) + self.presidio_analyze_chunk_size_bytes = self._coerce_analyze_chunk_size(self.presidio_analyze_chunk_size_bytes) + self._resync_output_stage_event_hook() + + def _resync_output_stage_event_hook(self) -> None: + if self.event_hook == GuardrailEventHooks.logging_only: + return + if self.apply_to_output: + self.event_hook = GuardrailEventHooks.post_call + return + if not self.output_parse_pii: + return + current_hook: Final = self.event_hook + if isinstance(current_hook, str) and current_hook != "post_call": + self.event_hook = GUARDRAIL_MODE_ADAPTER.validate_python((current_hook, GuardrailEventHooks.post_call)) + elif isinstance(current_hook, list) and "post_call" not in current_hook: + self.event_hook = GUARDRAIL_MODE_ADAPTER.validate_python((*current_hook, GuardrailEventHooks.post_call)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py index f834426d619..cc1234da4d9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py @@ -7,6 +7,7 @@ import json import os +from collections.abc import Mapping from typing import Any, Final, Literal from fastapi import HTTPException @@ -15,6 +16,7 @@ from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, + updated_litellm_param, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.custom_httpx.http_handler import ( @@ -111,7 +113,7 @@ class QualifireGuardrail(CustomGuardrail): "only 'block' and 'monitor' are supported." ) - def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: """ The base implementation blindly ``setattr``s every field on ``litellm_params`` (including ``on_flagged``) onto this live instance with no revalidation, so an @@ -121,7 +123,10 @@ class QualifireGuardrail(CustomGuardrail): the live instance untouched instead of raising after it's already been corrupted. Mirrors LakeraAIGuardrail's own override of this same method. """ - prospective_on_flagged: Final = litellm_params.on_flagged or self.on_flagged + raw_on_flagged: Final = updated_litellm_param(litellm_params, "on_flagged") + prospective_on_flagged: Final = ( + raw_on_flagged if isinstance(raw_on_flagged, str) and raw_on_flagged else self.on_flagged + ) self._validate_on_flagged(prospective_on_flagged) super().update_in_memory_litellm_params(litellm_params=litellm_params) diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index a8b33109900..88e24db207a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -159,7 +159,7 @@ class ToolPermissionGuardrail(CustomGuardrail): self._compiled_rule_targets = compiled_targets self._compiled_rule_patterns = compiled_patterns - def update_in_memory_litellm_params(self, litellm_params: LitellmParams | dict) -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: """Apply updated params in place, rebuilding the compiled rule state. The base implementation only ``setattr``s raw fields, which would leave @@ -169,17 +169,11 @@ class ToolPermissionGuardrail(CustomGuardrail): immediate in-memory sync take effect, mirroring the PresidioGuardrail override of this method. """ - # ``litellm_params`` may arrive as the raw DB dict (the proxy ``cast()``s - # it to ``LitellmParams`` without converting), so handle both shapes. The - # base ``setattr`` loop is model-only, so apply the dict case here. previous_rules: Final = self.rules - if isinstance(litellm_params, dict): - params = litellm_params - for key, value in params.items(): - setattr(self, key, value) - else: - super().update_in_memory_litellm_params(litellm_params) - params = vars(litellm_params) + params: Final[Mapping[str, object]] = ( + litellm_params if isinstance(litellm_params, Mapping) else vars(litellm_params) + ) + super().update_in_memory_litellm_params(litellm_params) # The generic update above sets ``self.rules`` from the incoming value # (None on a partial update that omits rules), but never rebuilds the @@ -187,7 +181,7 @@ class ToolPermissionGuardrail(CustomGuardrail): # the previous ruleset so a partial update doesn't silently wipe it. An # explicit empty list still clears the rules. rules: Final = params.get("rules") - if rules is not None: + if isinstance(rules, list): try: self._load_rules(rules) except Exception: diff --git a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py index 1aefa38ecf8..2928ea5d068 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/zscaler_ai_guard/zscaler_ai_guard.py @@ -4,6 +4,7 @@ # # +-------------------------------------------------------------+ import os +from collections.abc import Mapping from typing import TYPE_CHECKING, Final, Literal, Optional from fastapi import HTTPException @@ -12,6 +13,7 @@ from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, + updated_litellm_param, ) from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -102,9 +104,10 @@ class ZscalerAIGuard(CustomGuardrail): return timeout - def update_in_memory_litellm_params(self, litellm_params: "LitellmParams") -> None: + def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | Mapping[str, object]") -> None: super().update_in_memory_litellm_params(litellm_params) - self.timeout = self._resolve_timeout(litellm_params.timeout) + raw_timeout: Final = updated_litellm_param(litellm_params, "timeout") + self.timeout = self._resolve_timeout(raw_timeout if isinstance(raw_timeout, (int, float)) else None) @staticmethod def _resolve_metadata_value(request_data: dict | None, key: str) -> str | None: diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index dc13c09dd38..3abd952da8d 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -6,7 +6,7 @@ import os from collections.abc import Callable, Iterator, Mapping from datetime import datetime, timezone from itertools import chain, count -from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, cast +from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol from pydantic import ValidationError @@ -624,17 +624,19 @@ class InMemoryGuardrailHandler: """ Update a guardrail in memory - - updates the guardrail in memory - updates the guardrail params in litellm.callback_manager + - stores the guardrail in memory only after the callback update + succeeds, so a failed update stays visible as a diff to the + per-worker DB poller and gets retried instead of going stale """ + custom_guardrail_callback: Final = self.guardrail_id_to_custom_guardrail.get(guardrail_id) + updated_litellm_params: Final = guardrail.get("litellm_params") + if custom_guardrail_callback and updated_litellm_params: + custom_guardrail_callback.update_in_memory_litellm_params(litellm_params=updated_litellm_params) + self.IN_MEMORY_GUARDRAILS[guardrail_id] = guardrail self._sources[guardrail_id] = source - custom_guardrail_callback: Final = self.guardrail_id_to_custom_guardrail.get(guardrail_id) - if custom_guardrail_callback: - updated_litellm_params: Final = cast(LitellmParams, guardrail.get("litellm_params", {})) - custom_guardrail_callback.update_in_memory_litellm_params(litellm_params=updated_litellm_params) - def delete_in_memory_guardrail(self, guardrail_id: str) -> None: """ Delete a guardrail in memory and remove from litellm callbacks. diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 9b1cc977a64..f43e8c6e93e 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1073 + "limit": 1072 }, "TRY002": { "limit": 524 diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index d978eb48c12..c84fcd5ac11 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -9,6 +9,7 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.proxy._types import CallTypes, UserAPIKeyAuth +from litellm.types.guardrails import GuardrailEventHooks, LitellmParams from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailTracingDetail @@ -2237,3 +2238,66 @@ class TestRecordsOwnGuardrailInformation: ) assert _guardrail_entries(request_data) == [] + + +class TestUpdateInMemoryLitellmParams: + """A PUT /guardrails update reaches the live callback through + update_in_memory_litellm_params: it must accept both a LitellmParams object + and the raw DB dict, and resync self.event_hook (which dispatch reads) from + the incoming mode instead of only writing a dead self.mode attribute (LIT-6591).""" + + def _guardrail(self) -> CustomGuardrail: + return CustomGuardrail( + guardrail_name="update-test", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + supported_event_hooks=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call], + ) + + def test_mode_change_resyncs_event_hook_dispatch(self): + guardrail = self._guardrail() + + guardrail.update_in_memory_litellm_params( + LitellmParams(guardrail="update-test", mode="post_call", default_on=True) + ) + + assert guardrail.event_hook is GuardrailEventHooks.post_call + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is False + + def test_raw_db_dict_copies_params_and_resyncs_event_hook(self): + guardrail = self._guardrail() + + guardrail.update_in_memory_litellm_params( + { + "guardrail": "update-test", + "mode": "post_call", + "api_base": "https://guardrail.example.com", + "default_on": True, + } + ) + + assert guardrail.event_hook is GuardrailEventHooks.post_call + assert getattr(guardrail, "api_base", None) == "https://guardrail.example.com" + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_call) is True + + def test_strict_mode_rejects_unsupported_mode_without_mutating(self, monkeypatch): + monkeypatch.delenv("LITELLM_STRICT_GUARDRAIL_MODES", raising=False) + guardrail = self._guardrail() + + with pytest.raises(ValueError, match="not in the supported event hooks"): + guardrail.update_in_memory_litellm_params( + {"mode": "during_call", "api_base": "https://guardrail.example.com"} + ) + + assert guardrail.event_hook is GuardrailEventHooks.pre_call + assert getattr(guardrail, "api_base", None) is None + + def test_non_strict_mode_warns_and_applies_unsupported_mode(self, monkeypatch): + monkeypatch.setenv("LITELLM_STRICT_GUARDRAIL_MODES", "false") + guardrail = self._guardrail() + + guardrail.update_in_memory_litellm_params({"mode": "during_call", "api_base": "https://guardrail.example.com"}) + + assert guardrail.event_hook is GuardrailEventHooks.during_call + assert getattr(guardrail, "api_base", None) == "https://guardrail.example.com" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 4ee6741ee02..519df980d22 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -17,7 +17,7 @@ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( _OPTIONAL_PresidioPIIMasking, ) from litellm.exceptions import GuardrailRaisedException -from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType +from litellm.types.guardrails import GuardrailEventHooks, LitellmParams, PiiAction, PiiEntityType from litellm.types.utils import Choices, Message, ModelResponse from litellm.exceptions import BlockedPiiEntityError @@ -3167,6 +3167,41 @@ def test_update_in_memory_coerces_invalid_chunk_size(): assert guardrail.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES +def test_update_in_memory_output_callback_keeps_forced_post_call(): + """The registry-tracked callback for filter_scope='output' is initialized with a + forced post_call hook regardless of the configured mode; a mode-changing update + must not move it off the response stage (LIT-6591).""" + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + event_hook=GuardrailEventHooks.post_call.value, + ) + + guardrail.update_in_memory_litellm_params({"guardrail": "presidio", "mode": "pre_call", "default_on": True}) + + assert guardrail.event_hook is GuardrailEventHooks.post_call + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is False + + +def test_update_in_memory_output_parse_pii_keeps_post_call_expansion(): + """A guardrail with output_parse_pii must keep running on post_call to unmask the + response after a mode-changing update, mirroring the constructor's expansion.""" + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + event_hook="pre_call", + ) + + guardrail.update_in_memory_litellm_params( + LitellmParams(guardrail="presidio", mode="during_call", output_parse_pii=True, default_on=True) + ) + + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.during_call) is True + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is False + + def test_split_text_handles_chunk_size_below_char_width(): chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis( text="\U0001f642\U0001f642", chunk_size_bytes=3, overlap_chars=8 diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 2c0735970d3..015d530257b 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -179,6 +179,65 @@ def test_update_in_memory_guardrail(): assert handler.guardrail_id_to_custom_guardrail["123"].event_hook is GuardrailEventHooks.pre_call +def test_update_in_memory_guardrail_raw_db_dict_resyncs_event_hook(): + """PUT /guardrails hands this method the raw DB row, whose litellm_params is a + plain dict; the update must still apply and move dispatch to the new mode + instead of raising inside vars() and leaving the worker stale (LIT-6591).""" + handler = InMemoryGuardrailHandler() + handler.guardrail_id_to_custom_guardrail["123"] = CustomGuardrail( + guardrail_name="test-guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + supported_event_hooks=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call], + ) + + updated_row = { + "guardrail_id": "123", + "guardrail_name": "test-guardrail", + "litellm_params": {"guardrail": "test-guardrail", "mode": "post_call", "default_on": True}, + } + handler.update_in_memory_guardrail("123", updated_row) + + callback = handler.guardrail_id_to_custom_guardrail["123"] + assert callback.event_hook is GuardrailEventHooks.post_call + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_call) is True + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is False + assert handler.IN_MEMORY_GUARDRAILS["123"] == updated_row + + +def test_update_in_memory_guardrail_failed_callback_update_stays_visible_to_poller(monkeypatch): + """When the callback update raises, IN_MEMORY_GUARDRAILS must keep the old row: + storing the new row first would make the per-worker DB poller see no diff and + never re-initialize, leaving the PUT-serving worker stale until restart.""" + monkeypatch.delenv("LITELLM_STRICT_GUARDRAIL_MODES", raising=False) + handler = InMemoryGuardrailHandler() + stale_row = Guardrail( + guardrail_id="123", + guardrail_name="test-guardrail", + litellm_params=LitellmParams(guardrail="test-guardrail", mode="pre_call", default_on=True), + ) + handler.IN_MEMORY_GUARDRAILS["123"] = stale_row + handler.guardrail_id_to_custom_guardrail["123"] = CustomGuardrail( + guardrail_name="test-guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + supported_event_hooks=[GuardrailEventHooks.pre_call], + ) + + with pytest.raises(ValueError, match="not in the supported event hooks"): + handler.update_in_memory_guardrail( + "123", + { + "guardrail_id": "123", + "guardrail_name": "test-guardrail", + "litellm_params": {"guardrail": "test-guardrail", "mode": "post_call", "default_on": True}, + }, + ) + + assert handler.IN_MEMORY_GUARDRAILS["123"] == stale_row + assert handler.guardrail_id_to_custom_guardrail["123"].event_hook is GuardrailEventHooks.pre_call + + def _make_guardrail(guardrail_id: str, name: str = "g") -> Guardrail: return Guardrail( guardrail_id=guardrail_id, diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 3d2e97d55a5..73ea794f8a8 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22367 + "limit": 22362 }, "LIT002": { - "limit": 26777 + "limit": 26776 }, "LIT003": { "limit": 269 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1039 + "limit": 1038 }, "LIT007": { "limit": 0 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16507 + "limit": 16505 }, "LIT011": { "limit": 5535