fix(guardrails): track and tear down presidio sibling callbacks

initialize_presidio registers up to three callbacks per guardrail but the
registry only kept the first, so deleting or re-syncing the guardrail left
the post_call siblings serving the old config. The initializer now returns
every callback it registered, the registry tracks primary and siblings per
guardrail id, delete purges all of them from every callback list, and
update pushes the new params into each while siblings keep their stage.
This commit is contained in:
mateo-berri 2026-09-01 22:28:38 -07:00
parent 92d453373a
commit c18511be7d
4 changed files with 267 additions and 87 deletions

View file

@ -2,6 +2,7 @@
from typing import Any, Final
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import CommonProxyErrors
from litellm.types.guardrails import *
@ -85,7 +86,7 @@ def initialize_lakera_v2(litellm_params: LitellmParams, guardrail: Guardrail):
return _lakera_v2_callback
def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail):
def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> tuple[CustomGuardrail, ...]:
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
_OPTIONAL_PresidioPIIMasking,
)
@ -94,7 +95,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail):
run_input: Final = filter_scope in ("input", "both")
run_output: Final = filter_scope in ("output", "both")
def _make_presidio_callback(**overrides):
def _make_presidio_callback(**overrides) -> CustomGuardrail:
params: Final = dict(
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
@ -120,27 +121,27 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail):
litellm.logging_callback_manager.add_litellm_callback(callback)
return callback
primary_callback = None
if run_input:
primary_callback = _make_presidio_callback()
if litellm_params.output_parse_pii:
_make_presidio_callback(
output_parse_pii=True,
event_hook=GuardrailEventHooks.post_call.value,
)
if run_output:
output_callback: Final = _make_presidio_callback(
input_callback: Final = _make_presidio_callback() if run_input else None
unmask_output_callback: Final = (
_make_presidio_callback(
output_parse_pii=True,
event_hook=GuardrailEventHooks.post_call.value,
)
if run_input and litellm_params.output_parse_pii
else None
)
mask_output_callback: Final = (
_make_presidio_callback(
apply_to_output=True,
event_hook=GuardrailEventHooks.post_call.value,
output_parse_pii=False,
)
if primary_callback is None:
primary_callback = output_callback
return primary_callback
if run_output
else None
)
return tuple(
callback for callback in (input_callback, unmask_output_callback, mask_output_callback) if callback is not None
)
def initialize_hide_secrets(litellm_params: LitellmParams, guardrail: Guardrail):

View file

@ -3,10 +3,10 @@
import asyncio
import importlib
import os
from collections.abc import Callable, Iterator, Mapping
from collections.abc import Callable, Iterator, Mapping, Sequence
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, TypeAlias, cast
from pydantic import ValidationError
@ -90,6 +90,8 @@ guardrail_initializer_registry: Final = {
CONFIG_GUARDRAIL_ID_NAMESPACE: Final = uuid.UUID("625f63f4-935a-50e5-98b5-fbe77babc74a")
GuardrailCallbacks: TypeAlias = tuple[CustomGuardrail, ...]
guardrail_class_registry: Final[dict[str, type[CustomGuardrail]]] = {
SupportedGuardrailIntegrations.BEDROCK.value: BedrockGuardrail,
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail,
@ -424,6 +426,41 @@ def _apply_configured_bool_overrides(instance: CustomGuardrail, litellm_params:
instance.scan_raw_request = bool(litellm_params.scan_raw_request)
def _as_callback_tuple(
initialized: CustomGuardrail | Sequence[CustomGuardrail] | None,
) -> GuardrailCallbacks:
if initialized is None:
return ()
if isinstance(initialized, (list, tuple)):
return tuple(initialized)
return (initialized,)
def _configure_callback_scoping(
custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams
) -> None:
for scoping_param in (
"skip_system_message_in_guardrail",
"skip_tool_message_in_guardrail",
"scan_only_tool_results",
):
setattr(custom_guardrail_callback, scoping_param, getattr(litellm_params, scoping_param, None))
scan_only_tool_results_enabled: Final = effective_scan_only_tool_results_for_guardrail(custom_guardrail_callback)
if scan_only_tool_results_enabled and not custom_guardrail_callback.supports_scan_only_tool_results():
raise ValueError(
f"Guardrail {guardrail_name}: scan_only_tool_results is enabled, but this "
"guardrail's role filtering never scans tool results, so no request content would ever "
"be scanned. Remove scan_only_tool_results or the guardrail's role-filtering option."
)
if scan_only_tool_results_enabled and effective_skip_tool_message_for_guardrail(custom_guardrail_callback):
raise ValueError(
f"Guardrail {guardrail_name}: scan_only_tool_results and "
"skip_tool_message_in_guardrail are enabled together, which excludes every message from "
"scanning, so no request content would ever be scanned. Remove one of the two."
)
_apply_configured_bool_overrides(custom_guardrail_callback, litellm_params)
class InMemoryGuardrailHandler:
"""
Class that handles initializing guardrails and adding them to the CallbackManager
@ -440,6 +477,8 @@ class InMemoryGuardrailHandler:
Guardrail id to CustomGuardrail object mapping
"""
self.guardrail_id_to_sibling_callbacks: dict[str, GuardrailCallbacks] = {} # mutable-ok: per-id registry
self._sources: dict[str, Literal["db", "config"]] = {}
"""
Guardrail id to provenance marker. "db" entries are reconciled against
@ -474,7 +513,6 @@ class InMemoryGuardrailHandler:
self._sources[guardrail_id] = source
return self.IN_MEMORY_GUARDRAILS[guardrail_id]
custom_guardrail_callback: CustomGuardrail | None = None
litellm_params_data: Final = guardrail["litellm_params"]
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
@ -498,54 +536,15 @@ class InMemoryGuardrailHandler:
if guardrail_type is None:
raise ValueError("guardrail_type is required")
initializer: Final = guardrail_initializer_registry.get(guardrail_type)
if initializer:
# Try to call with llm_router first, fall back to without if it fails
import inspect
sig: Final = inspect.signature(initializer)
if "llm_router" in sig.parameters:
custom_guardrail_callback = initializer(
litellm_params,
guardrail,
llm_router,
)
else:
custom_guardrail_callback = initializer(litellm_params, guardrail)
elif isinstance(guardrail_type, str) and "." in guardrail_type:
custom_guardrail_callback = self.initialize_custom_guardrail(
guardrail=guardrail,
guardrail_type=guardrail_type,
litellm_params=litellm_params,
config_file_path=config_file_path,
)
else:
raise ValueError(f"Unsupported guardrail: {guardrail_type}")
if custom_guardrail_callback is not None:
for scoping_param in (
"skip_system_message_in_guardrail",
"skip_tool_message_in_guardrail",
"scan_only_tool_results",
):
setattr(custom_guardrail_callback, scoping_param, getattr(litellm_params, scoping_param, None))
scan_only_tool_results_enabled: Final = effective_scan_only_tool_results_for_guardrail(
custom_guardrail_callback
)
if scan_only_tool_results_enabled and not custom_guardrail_callback.supports_scan_only_tool_results():
raise ValueError(
f"Guardrail {guardrail['guardrail_name']}: scan_only_tool_results is enabled, but this "
"guardrail's role filtering never scans tool results, so no request content would ever "
"be scanned. Remove scan_only_tool_results or the guardrail's role-filtering option."
)
if scan_only_tool_results_enabled and effective_skip_tool_message_for_guardrail(custom_guardrail_callback):
raise ValueError(
f"Guardrail {guardrail['guardrail_name']}: scan_only_tool_results and "
"skip_tool_message_in_guardrail are enabled together, which excludes every message from "
"scanning, so no request content would ever be scanned. Remove one of the two."
)
_apply_configured_bool_overrides(custom_guardrail_callback, litellm_params)
created_callbacks: Final = self._create_callbacks(
guardrail=guardrail,
guardrail_type=guardrail_type,
litellm_params=litellm_params,
config_file_path=config_file_path,
llm_router=llm_router,
)
for custom_guardrail_callback in created_callbacks:
_configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params)
parsed_guardrail: Final = Guardrail(
guardrail_id=guardrail.get("guardrail_id"),
@ -556,11 +555,44 @@ class InMemoryGuardrailHandler:
# store references to the guardrail in memory
self.IN_MEMORY_GUARDRAILS[guardrail_id] = parsed_guardrail
self.guardrail_id_to_custom_guardrail[guardrail_id] = custom_guardrail_callback
self.guardrail_id_to_custom_guardrail[guardrail_id] = created_callbacks[0] if created_callbacks else None
self.guardrail_id_to_sibling_callbacks[guardrail_id] = created_callbacks[1:]
self._sources[guardrail_id] = source
return parsed_guardrail
def _create_callbacks(
self,
guardrail: Guardrail,
guardrail_type: str,
litellm_params: LitellmParams,
config_file_path: str | None,
llm_router: Optional["Router"],
) -> GuardrailCallbacks:
initializer: Final = guardrail_initializer_registry.get(guardrail_type)
if initializer:
import inspect
sig: Final = inspect.signature(initializer)
if "llm_router" in sig.parameters:
return _as_callback_tuple(initializer(litellm_params, guardrail, llm_router))
return _as_callback_tuple(initializer(litellm_params, guardrail))
if isinstance(guardrail_type, str) and "." in guardrail_type:
return _as_callback_tuple(
self.initialize_custom_guardrail(
guardrail=guardrail,
guardrail_type=guardrail_type,
litellm_params=litellm_params,
config_file_path=config_file_path,
)
)
raise ValueError(f"Unsupported guardrail: {guardrail_type}")
def _tracked_callbacks(self, guardrail_id: str) -> GuardrailCallbacks:
primary: Final = self.guardrail_id_to_custom_guardrail.get(guardrail_id)
siblings: Final = self.guardrail_id_to_sibling_callbacks.get(guardrail_id, ())
return (() if primary is None else (primary,)) + siblings
def initialize_custom_guardrail(
self,
guardrail: Guardrail,
@ -630,10 +662,15 @@ class InMemoryGuardrailHandler:
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)
tracked_callbacks: Final = self._tracked_callbacks(guardrail_id)
if not tracked_callbacks:
return
updated_litellm_params: Final = cast(LitellmParams, guardrail.get("litellm_params", {}))
tracked_callbacks[0].update_in_memory_litellm_params(litellm_params=updated_litellm_params)
for sibling_callback in tracked_callbacks[1:]:
sibling_stage = sibling_callback.event_hook
sibling_callback.update_in_memory_litellm_params(litellm_params=updated_litellm_params)
sibling_callback.event_hook = sibling_stage
def delete_in_memory_guardrail(self, guardrail_id: str) -> None:
"""
@ -648,11 +685,11 @@ class InMemoryGuardrailHandler:
self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None)
self._sources.pop(guardrail_id, None)
custom_guardrail_callback: Final = self.guardrail_id_to_custom_guardrail.pop(guardrail_id, None)
if custom_guardrail_callback is None:
return
litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback)
tracked_callbacks: Final = self._tracked_callbacks(guardrail_id)
self.guardrail_id_to_custom_guardrail.pop(guardrail_id, None)
self.guardrail_id_to_sibling_callbacks.pop(guardrail_id, None)
for custom_guardrail_callback in tracked_callbacks:
litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback)
def list_in_memory_guardrails(self) -> list[Guardrail]:
"""

View file

@ -842,24 +842,37 @@ async def test_presidio_filter_scope_initializer(monkeypatch):
params_input = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="input")
guardrail_dict = {"guardrail_name": "g1"}
cb = initialize_presidio(params_input, guardrail_dict)
assert cb is created[0]
callbacks = initialize_presidio(params_input, guardrail_dict)
assert callbacks == (created[0],)
assert created[0].apply_to_output is False
# output-only
created.clear()
params_output = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="output")
cb = initialize_presidio(params_output, guardrail_dict)
callbacks = initialize_presidio(params_output, guardrail_dict)
assert len(created) == 1
assert callbacks == (created[0],)
assert created[0].apply_to_output is True
# both -> expect two callbacks (input + output)
# both -> expect two callbacks (input + output), both returned, input first
created.clear()
params_both = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="both")
cb = initialize_presidio(params_both, guardrail_dict)
callbacks = initialize_presidio(params_both, guardrail_dict)
assert len(created) == 2
assert any(not c.apply_to_output for c in created)
assert any(c.apply_to_output for c in created)
assert callbacks == tuple(created)
assert callbacks[0].apply_to_output is False
assert callbacks[1].apply_to_output is True
# both + output_parse_pii -> three callbacks, all returned, input first
created.clear()
params_all = LitellmParams(
guardrail="presidio", mode="pre_call", presidio_filter_scope="both", output_parse_pii=True
)
callbacks = initialize_presidio(params_all, guardrail_dict)
assert len(created) == 3
assert callbacks == tuple(created)
assert callbacks[0].apply_to_output is False
assert mgr.added[-3:] == list(created)
@pytest.mark.asyncio

View file

@ -491,6 +491,135 @@ def test_repeated_db_sync_does_not_accumulate_runner_instances():
cb_list[:] = snapshot
PRESIDIO_SIBLINGS_GID = "55555555-5555-5555-5555-555555555555"
PRESIDIO_SIBLINGS_NAME = "presidio-siblings"
def _presidio_db_guardrail(pii_entities_config: dict) -> Guardrail:
return Guardrail(
guardrail_id=PRESIDIO_SIBLINGS_GID,
guardrail_name=PRESIDIO_SIBLINGS_NAME,
litellm_params={
"guardrail": "presidio",
"mode": "pre_call",
"default_on": True,
"output_parse_pii": True,
"presidio_filter_scope": "both",
"presidio_analyzer_api_base": "https://fakelink.com/v1/presidio/analyze",
"presidio_anonymizer_api_base": "https://fakelink.com/v1/presidio/anonymize",
"pii_entities_config": pii_entities_config,
},
)
def _presidio_callbacks_in(cb_list) -> list:
return [
callback
for callback in cb_list
if isinstance(callback, CustomGuardrail) and getattr(callback, "guardrail_name", None) == PRESIDIO_SIBLINGS_NAME
]
def test_presidio_siblings_are_tracked_and_deleted_together():
"""
A presidio guardrail scoped to both stages registers the pre_call primary plus
the post_call unmask and mask-output siblings. Deleting the guardrail must remove
all three from every callback list, not just the primary.
"""
import litellm
handler = InMemoryGuardrailHandler()
lists = _all_callback_lists()
snapshots = [list(cb_list) for cb_list in lists]
try:
handler.initialize_guardrail(_presidio_db_guardrail({"EMAIL_ADDRESS": "MASK"}))
registered = _presidio_callbacks_in(litellm.callbacks)
assert len(registered) == 3
primary = handler.guardrail_id_to_custom_guardrail[PRESIDIO_SIBLINGS_GID]
siblings = handler.guardrail_id_to_sibling_callbacks[PRESIDIO_SIBLINGS_GID]
assert primary is registered[0]
assert siblings == tuple(registered[1:])
assert [sibling.event_hook for sibling in siblings] == [GuardrailEventHooks.post_call] * 2
for cb_list in lists[1:]:
cb_list.extend(registered)
handler.delete_in_memory_guardrail(PRESIDIO_SIBLINGS_GID)
for cb_list in lists:
assert _presidio_callbacks_in(cb_list) == []
assert PRESIDIO_SIBLINGS_GID not in handler.guardrail_id_to_custom_guardrail
assert PRESIDIO_SIBLINGS_GID not in handler.guardrail_id_to_sibling_callbacks
finally:
for cb_list, snapshot in zip(lists, snapshots):
cb_list[:] = snapshot
def test_update_in_memory_guardrail_reaches_presidio_siblings_and_keeps_their_stage():
import litellm
handler = InMemoryGuardrailHandler()
lists = _all_callback_lists()
snapshots = [list(cb_list) for cb_list in lists]
try:
handler.initialize_guardrail(_presidio_db_guardrail({"EMAIL_ADDRESS": "MASK", "IP_ADDRESS": "MASK"}))
tracked = _presidio_callbacks_in(litellm.callbacks)
roles_before = [(callback.apply_to_output, callback.event_hook) for callback in tracked]
updated = Guardrail(
guardrail_id=PRESIDIO_SIBLINGS_GID,
guardrail_name=PRESIDIO_SIBLINGS_NAME,
litellm_params=LitellmParams(
guardrail="presidio",
mode="pre_call",
default_on=True,
output_parse_pii=True,
presidio_filter_scope="both",
presidio_analyzer_api_base="https://fakelink.com/v1/presidio/analyze",
presidio_anonymizer_api_base="https://fakelink.com/v1/presidio/anonymize",
pii_entities_config={"EMAIL_ADDRESS": "MASK"},
),
)
handler.update_in_memory_guardrail(guardrail_id=PRESIDIO_SIBLINGS_GID, guardrail=updated)
assert [callback.pii_entities_config for callback in tracked] == [{"EMAIL_ADDRESS": "MASK"}] * 3
assert [(callback.apply_to_output, callback.event_hook) for callback in tracked] == roles_before
assert _presidio_callbacks_in(litellm.callbacks) == tracked
finally:
for cb_list, snapshot in zip(lists, snapshots):
cb_list[:] = snapshot
def test_repeated_db_sync_replaces_presidio_siblings_instead_of_leaking_stale_ones():
"""
The callback manager dedupes custom loggers by their scalar attributes, so a
leaked post_call sibling blocks the re-initialized sibling from registering and
keeps serving the previous entity config. After every DB re-sync, each callback
list must hold exactly the three current instances, all on the latest config.
"""
import litellm
handler = InMemoryGuardrailHandler()
lists = _all_callback_lists()
snapshots = [list(cb_list) for cb_list in lists]
try:
entity_configs = [{"EMAIL_ADDRESS": "MASK"}, {"EMAIL_ADDRESS": "MASK", "IP_ADDRESS": "MASK"}]
for cycle in range(4):
latest = entity_configs[cycle % 2]
handler.sync_guardrail_from_db(_presidio_db_guardrail(latest))
for cb_list in lists[1:]:
cb_list.extend(_presidio_callbacks_in(litellm.callbacks))
for cb_list in lists:
current = _presidio_callbacks_in(cb_list)
assert len({id(callback) for callback in current}) == 3
assert all(callback.pii_entities_config == latest for callback in current)
finally:
for cb_list, snapshot in zip(lists, snapshots):
cb_list[:] = snapshot
def _judge_guardrail(guardrail_id: str) -> Guardrail:
return Guardrail(
guardrail_id=guardrail_id,