mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #39271 from BerriAI/litellm_fix_presidio_sibling_callback_leak
fix(guardrails): track and tear down presidio sibling callbacks on delete and update
This commit is contained in:
commit
86ca146ea2
5 changed files with 291 additions and 87 deletions
|
|
@ -1633,6 +1633,8 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
Update the guardrails litellm params in memory
|
||||
"""
|
||||
super().update_in_memory_litellm_params(litellm_params)
|
||||
if self.apply_to_output:
|
||||
self.output_parse_pii = False
|
||||
if litellm_params.pii_entities_config:
|
||||
self.pii_entities_config = litellm_params.pii_entities_config
|
||||
if litellm_params.presidio_score_thresholds:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -3116,6 +3129,18 @@ def test_update_in_memory_applies_analyze_chunk_size():
|
|||
assert guardrail.presidio_analyze_chunk_size_bytes == 99_000
|
||||
|
||||
|
||||
def test_update_in_memory_keeps_output_masker_from_unmasking():
|
||||
masker = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True, output_parse_pii=False)
|
||||
unmasker = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True)
|
||||
params = LitellmParams(guardrail="presidio", mode="pre_call", output_parse_pii=True)
|
||||
|
||||
masker.update_in_memory_litellm_params(params)
|
||||
unmasker.update_in_memory_litellm_params(params)
|
||||
|
||||
assert (masker.apply_to_output, masker.output_parse_pii) == (True, False)
|
||||
assert (unmasker.apply_to_output, unmasker.output_parse_pii) == (False, True)
|
||||
|
||||
|
||||
def test_merge_drops_truncated_same_type_fragment_from_overlap():
|
||||
"""A boundary entity seen truncated by chunk 1 and whole by chunk 2 must
|
||||
merge to the single full span; keeping both overlapping spans corrupts the
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Iterable
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -491,6 +492,144 @@ 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[str, str]) -> 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: Iterable[object]) -> list[CustomGuardrail]:
|
||||
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.output_parse_pii, callback.event_hook) for callback in tracked
|
||||
]
|
||||
assert roles_before == [
|
||||
(False, True, [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]),
|
||||
(False, True, GuardrailEventHooks.post_call),
|
||||
(True, False, GuardrailEventHooks.post_call),
|
||||
]
|
||||
|
||||
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.output_parse_pii, 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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue