mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): honor scan_raw_request for pipeline-managed guardrails
A scan_raw_request=True guardrail that is itself a pipeline step never saw raw_request_snapshot: PipelineExecutor.execute_steps had no way to receive it, and pipeline-managed guardrails are fully excluded from the normal sequential/parallel loops that implement the flag. Such a guardrail silently evaluated whatever an earlier pass_data step in the same pipeline had already rewritten, defeating the flag for pipeline-managed guardrails. Moves the snapshot helper (renamed independent_snapshot) from proxy/utils.py to litellm_core_utils/core_helpers.py so pipeline_executor.py can use the same independent-copy logic without a circular import, threads raw_request_snapshot through _maybe_execute_pipelines and PipelineExecutor.execute_steps/_run_step, and discards a scan_raw_request step's returned data the same way the sequential/parallel loops already do.
This commit is contained in:
parent
62aa350f57
commit
4a483bc70e
3 changed files with 135 additions and 4 deletions
|
|
@ -454,6 +454,62 @@ def safe_deep_copy(data):
|
|||
return new_data
|
||||
|
||||
|
||||
def independent_snapshot(
|
||||
data: dict, # mutable-ok: caller-defined request-payload shape
|
||||
) -> dict: # mutable-ok: caller-defined request-payload shape
|
||||
"""
|
||||
A copy of ``data`` whose top-level keys are deep-copied independently
|
||||
where possible -- always attempted, regardless of
|
||||
``litellm.safe_memory_mode``. Unlike ``safe_deep_copy``, which can return
|
||||
the *original* object outright under that mode (defeating any isolation
|
||||
guarantee for every key, not just the ones that need it), this never
|
||||
skips copying wholesale.
|
||||
|
||||
Real proxy requests carry ``data["litellm_logging_obj"]`` (a ``Logging``
|
||||
instance nesting a live OTel span with a real lock) by the time
|
||||
``pre_call_hook`` runs, which can never be deep-copied. Any individual
|
||||
key that fails to deep-copy falls back to sharing its original
|
||||
reference, same crash tolerance as ``safe_deep_copy``'s own per-key
|
||||
fallback; callers needing true isolation (e.g. a guardrail's
|
||||
``scan_raw_request`` snapshot) only depend on the keys that are plain,
|
||||
cleanly-copyable structures (``messages``/``input``,
|
||||
``metadata``/``litellm_metadata``).
|
||||
"""
|
||||
sanitized: Final = {
|
||||
key: (
|
||||
{ # mutable-ok: same request-payload shape as data
|
||||
inner_key: ("placeholder" if inner_key == "litellm_parent_otel_span" else inner_value)
|
||||
for inner_key, inner_value in value.items()
|
||||
}
|
||||
if key in ("metadata", "litellm_metadata") and isinstance(value, dict)
|
||||
else value
|
||||
)
|
||||
for key, value in data.items()
|
||||
}
|
||||
|
||||
def _copied_value(key: str, sanitized_value: object) -> object:
|
||||
try:
|
||||
copied_value: Final = copy.deepcopy(sanitized_value)
|
||||
except Exception: # noqa: BLE001 # any unpicklable value falls back to the original reference for this key only
|
||||
return data.get(key)
|
||||
original_value: Final = data.get(key)
|
||||
if (
|
||||
key in ("metadata", "litellm_metadata")
|
||||
and isinstance(copied_value, dict)
|
||||
and isinstance(original_value, dict)
|
||||
and "litellm_parent_otel_span" in original_value
|
||||
):
|
||||
return { # mutable-ok: same request-payload shape as data
|
||||
**copied_value,
|
||||
"litellm_parent_otel_span": original_value["litellm_parent_otel_span"],
|
||||
}
|
||||
return copied_value
|
||||
|
||||
return { # mutable-ok: same request-payload shape as data
|
||||
key: _copied_value(key, value) for key, value in sanitized.items()
|
||||
}
|
||||
|
||||
|
||||
def filter_exceptions_from_params(data: Any, max_depth: int = 20) -> Any:
|
||||
"""
|
||||
Recursively filter out Exception objects and callable objects from dicts/lists.
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.integrations.custom_guardrail import (
|
|||
ModifyResponseException,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import independent_snapshot
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
|
@ -41,6 +42,7 @@ class PipelineExecutor:
|
|||
user_api_key_dict: Any,
|
||||
call_type: str,
|
||||
policy_name: str,
|
||||
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
|
||||
) -> PipelineExecutionResult:
|
||||
"""
|
||||
Execute pipeline steps sequentially with conditional actions.
|
||||
|
|
@ -52,6 +54,11 @@ class PipelineExecutor:
|
|||
user_api_key_dict: User API key auth
|
||||
call_type: Type of call (completion, etc.)
|
||||
policy_name: Name of the owning policy (for logging)
|
||||
raw_request_snapshot: pristine pre-pipeline, pre-guardrail request
|
||||
(taken by the caller before any guardrail or pipeline ran), so a
|
||||
step whose guardrail opted into ``scan_raw_request`` evaluates
|
||||
the original request instead of whatever an earlier
|
||||
``pass_data`` step in this same pipeline already rewrote.
|
||||
|
||||
Returns:
|
||||
PipelineExecutionResult with terminal action and step results
|
||||
|
|
@ -75,6 +82,7 @@ class PipelineExecutor:
|
|||
data=working_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
raw_request_snapshot=raw_request_snapshot,
|
||||
)
|
||||
|
||||
duration = time.perf_counter() - start_time
|
||||
|
|
@ -143,6 +151,7 @@ class PipelineExecutor:
|
|||
data: dict,
|
||||
user_api_key_dict: Any,
|
||||
call_type: str,
|
||||
raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data
|
||||
) -> tuple[
|
||||
Literal["pass", "fail", "error"],
|
||||
dict | None,
|
||||
|
|
@ -172,20 +181,33 @@ class PipelineExecutor:
|
|||
data["metadata"] = {}
|
||||
data["metadata"]["guardrails"] = [step.guardrail]
|
||||
|
||||
# A scan_raw_request step evaluates the pristine pre-pipeline
|
||||
# snapshot instead of `data` (which earlier pass_data steps in
|
||||
# this same pipeline may have already rewritten), same reason
|
||||
# the normal sequential/parallel guardrail loops do this.
|
||||
scans_raw_request: Final = getattr(callback, "scan_raw_request", False)
|
||||
hook_input: Final[dict] = ( # mutable-ok: same request-payload shape as data
|
||||
independent_snapshot(raw_request_snapshot)
|
||||
if scans_raw_request and raw_request_snapshot is not None
|
||||
else data
|
||||
)
|
||||
if hook_input is not data:
|
||||
hook_input.setdefault("metadata", {})["guardrails"] = [step.guardrail]
|
||||
|
||||
# Use unified_guardrail path if callback implements apply_guardrail
|
||||
target: CustomLogger = callback
|
||||
use_unified: Final = (
|
||||
"apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
|
||||
)
|
||||
if use_unified:
|
||||
data["guardrail_to_apply"] = callback
|
||||
hook_input["guardrail_to_apply"] = callback
|
||||
target = UnifiedLLMGuardrails()
|
||||
|
||||
if mode == "pre_call":
|
||||
response = await target.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=None,
|
||||
data=data,
|
||||
data=hook_input,
|
||||
call_type=call_type,
|
||||
)
|
||||
if isinstance(callback, CustomGuardrail):
|
||||
|
|
@ -201,9 +223,13 @@ class PipelineExecutor:
|
|||
else:
|
||||
return ("error", None, f"Unsupported pipeline mode: {mode}", None)
|
||||
|
||||
# Normal return means pass
|
||||
# Normal return means pass. A scan_raw_request step is block-only,
|
||||
# same contract as run_in_parallel/scan_raw_request elsewhere: any
|
||||
# data it returned is discarded, since applying it on top of the
|
||||
# raw snapshot would silently undo whatever an earlier step in
|
||||
# this pipeline already did.
|
||||
modified_data = None
|
||||
if response is not None and isinstance(response, dict):
|
||||
if response is not None and isinstance(response, dict) and not scans_raw_request:
|
||||
modified_data = response
|
||||
return ("pass", modified_data, None, None)
|
||||
|
||||
|
|
|
|||
|
|
@ -468,6 +468,55 @@ async def test_data_forwarding_pii_masking(monkeypatch):
|
|||
assert result.modified_data["messages"][0]["content"] == "Hello [REDACTED]"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_raw_request_step_sees_pre_pipeline_content(monkeypatch):
|
||||
"""
|
||||
veria-ai finding on BerriAI/litellm#34940: a scan_raw_request=True guardrail
|
||||
that is itself a pipeline step never saw raw_request_snapshot at all --
|
||||
execute_steps had no way to receive it, so it evaluated whatever an earlier
|
||||
pass_data step in the same pipeline had already rewritten, defeating the
|
||||
whole point of the flag for pipeline-managed guardrails.
|
||||
|
||||
Pipeline: pii-masker (pass_data: true, on_pass: next) -> content-check
|
||||
(scan_raw_request=True, on_pass: allow). Input: "Hello John Smith".
|
||||
content-check must still see the original, unmasked content.
|
||||
"""
|
||||
pii_guard = PiiMaskingGuardrail(guardrail_name="pii-masker")
|
||||
content_guard = ContentCheckGuardrail(guardrail_name="content-check")
|
||||
content_guard.scan_raw_request = True
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="pii-masker",
|
||||
on_fail="block",
|
||||
on_pass="next",
|
||||
pass_data=True,
|
||||
),
|
||||
PipelineStep(guardrail="content-check", on_fail="block", on_pass="allow"),
|
||||
],
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [pii_guard, content_guard])
|
||||
original_data = {"messages": [{"role": "user", "content": "Hello John Smith"}]}
|
||||
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data=original_data,
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="pii-then-safety",
|
||||
raw_request_snapshot=original_data,
|
||||
)
|
||||
|
||||
assert pii_guard.calls == 1
|
||||
assert content_guard.calls == 1
|
||||
assert content_guard.received_messages[0]["content"] == "Hello John Smith"
|
||||
assert result.terminal_action == "allow"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_not_found_uses_on_fail(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue