mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): resync event_hook and accept raw dicts in in-memory guardrail updates
This commit is contained in:
parent
81277252e1
commit
8c72342ad5
16 changed files with 308 additions and 122 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -231,7 +231,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 1073
|
||||
"limit": 1072
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 524
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue