fix(guardrails): resync event_hook and accept raw dicts in in-memory guardrail updates

This commit is contained in:
mateo-berri 2026-09-01 18:25:34 -07:00
parent 81277252e1
commit 8c72342ad5
16 changed files with 308 additions and 122 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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,

View file

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

View file

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

View file

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

View file

@ -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:

View file

@ -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:

View file

@ -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.

View file

@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1073
"limit": 1072
},
"TRY002": {
"limit": 524

View file

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

View file

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

View file

@ -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,

View file

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