mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): match multi-mode guardrail_mode without false-COMPLIANT (#32832)
* fix(proxy): match list/dict guardrail_mode in compliance mode checks * test(compliance): cover ComplianceChecker guardrail_mode shapes (str/list/dict/None) * fix(proxy): trust only Mode.default in compliance mode matching (ignore tag overrides) * fix(proxy): match dict guardrail_mode only when every branch runs in mode (no false-compliant) * fix(proxy): treat multi-mode guardrail_mode as unresolved (no false-compliant) The list branch previously counted a guardrail configured with mode: [pre_call, post_call] under every listed mode. But when the writer cannot infer the concrete hook that fired (apply_guardrail invocations), the raw list is logged, and an image-only request that only reaches the post-call path still records both modes. That let a pre_call compliance check pass on a request that only ran post_call. Match the tightened dict semantics: a list now counts for mode only when every listed mode equals mode. Same trade-off (under-report instead of false-COMPLIANT). Speculative set support is dropped (spend logs are JSON-serialized, sets do not cross the wire). Tests updated to reflect the tightened list semantics, deduplicated (single TestModeMatching class), and shortened. The invariant is now expressed as a computed check: True implies every branch runs in the matched mode. --------- Co-authored-by: Marton Schneider <marton@schneider.co.nl>
This commit is contained in:
parent
8f24f2f767
commit
51f6e1b7a5
2 changed files with 255 additions and 12 deletions
|
|
@ -35,15 +35,49 @@ class ComplianceChecker:
|
|||
If a guardrail doesn't have a mode specified, it's treated as pre-call
|
||||
(the most common case).
|
||||
"""
|
||||
result = []
|
||||
for g in self.guardrails:
|
||||
g_mode = g.get("guardrail_mode")
|
||||
# If no mode specified, default to pre_call
|
||||
if g_mode is None and mode == "pre_call":
|
||||
result.append(g)
|
||||
elif g_mode == mode:
|
||||
result.append(g)
|
||||
return result
|
||||
return [g for g in self.guardrails if self._mode_matches(g.get("guardrail_mode"), mode)]
|
||||
|
||||
@staticmethod
|
||||
def _mode_matches(g_mode: object, mode: str) -> bool:
|
||||
"""
|
||||
Return True only when a guardrail with logged ``guardrail_mode`` of
|
||||
``g_mode`` is guaranteed to have run in ``mode`` for the audited request.
|
||||
|
||||
``guardrail_mode`` in a spend log can take several shapes because
|
||||
``LitellmParams.mode`` is typed ``Union[str, List[str], Mode]``, and
|
||||
when the event type cannot be inferred at write time the raw config is
|
||||
logged verbatim. The spend log records the configured mode(s), not the
|
||||
concrete hook that fired for a given request; a match reports a mode
|
||||
satisfied only when every configured branch runs in that mode, so True
|
||||
never claims a hook the guardrail may not have actually executed.
|
||||
|
||||
Fails safe: if the guarantee cannot be established (missing default,
|
||||
divergent per-tag override, or a list that runs in more than one mode),
|
||||
the guardrail counts for no mode. The precise fix is to log the
|
||||
resolved event mode and match on it; this is the safe interim.
|
||||
"""
|
||||
if g_mode is None:
|
||||
return mode == "pre_call"
|
||||
if isinstance(g_mode, str):
|
||||
return g_mode == mode
|
||||
if isinstance(g_mode, (list, tuple)):
|
||||
return bool(g_mode) and all(m == mode for m in g_mode)
|
||||
if isinstance(g_mode, dict):
|
||||
default = g_mode.get("default")
|
||||
if default is None:
|
||||
return False
|
||||
tags = g_mode.get("tags")
|
||||
tag_branches = list(tags.values()) if isinstance(tags, dict) else []
|
||||
|
||||
def _branch_runs_in_mode(branch: object) -> bool:
|
||||
if isinstance(branch, str):
|
||||
return branch == mode
|
||||
if isinstance(branch, (list, tuple)):
|
||||
return bool(branch) and all(m == mode for m in branch)
|
||||
return False
|
||||
|
||||
return all(_branch_runs_in_mode(branch) for branch in [default, *tag_branches])
|
||||
return False
|
||||
|
||||
def _has_guardrail_intervention(self, guardrails: List[Dict]) -> bool:
|
||||
"""Check if any guardrail intervened (blocked/masked content)."""
|
||||
|
|
|
|||
|
|
@ -7,9 +7,7 @@ import sys
|
|||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.proxy.compliance_checks import ComplianceChecker
|
||||
from litellm.types.proxy.compliance_endpoints import ComplianceCheckRequest
|
||||
|
|
@ -385,3 +383,214 @@ class TestGdprCompliant:
|
|||
)
|
||||
checks = ComplianceChecker(data).check_gdpr()
|
||||
assert all(c.passed for c in checks)
|
||||
|
||||
|
||||
class TestModeMatching:
|
||||
"""Direct coverage of ComplianceChecker._mode_matches for every shape.
|
||||
|
||||
LitellmParams.mode is Union[str, List[str], Mode], so a spend-log
|
||||
guardrail_mode can be None / str / list / tuple / dict. A prior
|
||||
implementation compared `g_mode == mode`, which silently failed for the
|
||||
non-str shapes and reported NON-COMPLIANT for every multi-mode guardrail.
|
||||
A match now reports a mode satisfied only when every configured branch
|
||||
runs in that mode: fails safe, no false-COMPLIANT.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"g_mode, mode, expected",
|
||||
[
|
||||
(None, "pre_call", True),
|
||||
(None, "post_call", False),
|
||||
(None, "during_call", False),
|
||||
("pre_call", "pre_call", True),
|
||||
("post_call", "pre_call", False),
|
||||
("during_call", "during_call", True),
|
||||
# list/tuple: only guaranteed when every listed mode equals `mode`
|
||||
(["pre_call"], "pre_call", True),
|
||||
(["pre_call", "pre_call"], "pre_call", True),
|
||||
(["pre_call", "post_call"], "pre_call", False),
|
||||
(["pre_call", "post_call"], "post_call", False),
|
||||
([], "pre_call", False),
|
||||
(("during_call",), "during_call", True),
|
||||
(("pre_call", "post_call"), "pre_call", False),
|
||||
# dict: default only
|
||||
({"default": "pre_call"}, "pre_call", True),
|
||||
({"default": "pre_call"}, "post_call", False),
|
||||
({"default": ["pre_call", "post_call"]}, "post_call", False),
|
||||
({"default": ["pre_call"]}, "pre_call", True),
|
||||
# dict with tags: every branch must run in mode
|
||||
({"default": "pre_call", "tags": {"a": "pre_call"}}, "pre_call", True),
|
||||
({"default": "pre_call", "tags": {"a": ["pre_call"]}}, "pre_call", True),
|
||||
({"default": "pre_call", "tags": {"a": ["pre_call", "post_call"]}}, "pre_call", False),
|
||||
({"default": "pre_call", "tags": {"eu": "post_call"}}, "pre_call", False),
|
||||
({"default": "pre_call", "tags": {"eu": "post_call"}}, "post_call", False),
|
||||
({"default": "pre_call", "tags": {"eu": ["during_call"]}}, "during_call", False),
|
||||
({"default": ["pre_call", "post_call"], "tags": {"a": "pre_call"}}, "post_call", False),
|
||||
# Missing default: untagged routing is unknown, nothing guaranteed
|
||||
({"tags": {"x": "post_call"}}, "pre_call", False),
|
||||
({"tags": {"x": "post_call"}}, "post_call", False),
|
||||
({}, "pre_call", False),
|
||||
({}, "post_call", False),
|
||||
({"default": 123}, "pre_call", False),
|
||||
# Unknown top-level shapes never match
|
||||
(5, "pre_call", False),
|
||||
(object(), "pre_call", False),
|
||||
],
|
||||
)
|
||||
def test_mode_matches(self, g_mode, mode, expected):
|
||||
assert ComplianceChecker._mode_matches(g_mode, mode) is expected
|
||||
|
||||
def test_list_mode_guardrail_not_misclassified(self):
|
||||
"""A guardrail configured with mode ["pre_call", "post_call"] is logged
|
||||
with the raw list when the writer cannot infer the concrete hook that
|
||||
ran (e.g. apply_guardrail invocations). The spend log records "this
|
||||
guardrail could have run at either hook", not "which hook fired this
|
||||
request". Counting it for both would let a request that only fired
|
||||
post_call pass a pre_call compliance check. It counts for neither."""
|
||||
data = ComplianceCheckRequest(
|
||||
request_id="req-mode-1",
|
||||
user_id="user-1",
|
||||
model="gpt-4",
|
||||
timestamp="2026-02-17T00:00:00Z",
|
||||
guardrail_information=[
|
||||
{
|
||||
"guardrail_name": "pii_masking",
|
||||
"guardrail_mode": ["pre_call", "post_call"],
|
||||
"guardrail_status": "success",
|
||||
}
|
||||
],
|
||||
)
|
||||
checker = ComplianceChecker(data)
|
||||
assert len(checker._get_guardrails_by_mode("pre_call")) == 0
|
||||
assert len(checker._get_guardrails_by_mode("post_call")) == 0
|
||||
results = {c.check_name: c.passed for c in checker.check_eu_ai_act()}
|
||||
assert results["Content screened before LLM"] is False
|
||||
|
||||
def test_list_mode_single_value_counts(self):
|
||||
"""A single-entry list ["pre_call"] runs pre_call unconditionally, so it
|
||||
counts for pre_call and no other mode."""
|
||||
data = ComplianceCheckRequest(
|
||||
request_id="req-mode-1b",
|
||||
user_id="user-1",
|
||||
model="gpt-4",
|
||||
timestamp="2026-02-17T00:00:00Z",
|
||||
guardrail_information=[
|
||||
{
|
||||
"guardrail_name": "pii_masking",
|
||||
"guardrail_mode": ["pre_call"],
|
||||
"guardrail_status": "success",
|
||||
}
|
||||
],
|
||||
)
|
||||
checker = ComplianceChecker(data)
|
||||
assert len(checker._get_guardrails_by_mode("pre_call")) == 1
|
||||
assert len(checker._get_guardrails_by_mode("post_call")) == 0
|
||||
|
||||
def test_dict_tag_routed_guardrail_not_misclassified(self):
|
||||
"""A tag-routed guardrail (default=pre_call, a post_call tag) is not
|
||||
guaranteed to run in either mode, so it counts for neither."""
|
||||
data = ComplianceCheckRequest(
|
||||
request_id="req-mode-2",
|
||||
user_id="user-1",
|
||||
model="gpt-4",
|
||||
timestamp="2026-02-17T12:00:00Z",
|
||||
guardrail_information=[
|
||||
{
|
||||
"guardrail_name": "pii_masking",
|
||||
"guardrail_mode": {"default": "pre_call", "tags": {"eu": "post_call"}},
|
||||
"guardrail_status": "success",
|
||||
}
|
||||
],
|
||||
)
|
||||
checker = ComplianceChecker(data)
|
||||
assert len(checker._get_guardrails_by_mode("pre_call")) == 0
|
||||
assert len(checker._get_guardrails_by_mode("post_call")) == 0
|
||||
results = {c.check_name: c.passed for c in checker.check_eu_ai_act()}
|
||||
assert results["Content screened before LLM"] is False
|
||||
|
||||
def test_dict_all_branches_pre_call_counts(self):
|
||||
"""When default and every tag override all run pre_call, the guardrail is
|
||||
guaranteed pre_call regardless of routing, so it counts for pre_call."""
|
||||
data = ComplianceCheckRequest(
|
||||
request_id="req-mode-4",
|
||||
user_id="user-1",
|
||||
model="gpt-4",
|
||||
timestamp="2026-02-17T12:00:00Z",
|
||||
guardrail_information=[
|
||||
{
|
||||
"guardrail_name": "pii_masking",
|
||||
"guardrail_mode": {"default": "pre_call", "tags": {"eu": "pre_call"}},
|
||||
"guardrail_status": "success",
|
||||
}
|
||||
],
|
||||
)
|
||||
checker = ComplianceChecker(data)
|
||||
assert len(checker._get_guardrails_by_mode("pre_call")) == 1
|
||||
|
||||
def test_none_mode_defaults_to_pre_call(self):
|
||||
"""A guardrail logged without a mode counts as pre_call only."""
|
||||
data = ComplianceCheckRequest(
|
||||
request_id="req-mode-3",
|
||||
user_id="user-1",
|
||||
model="gpt-4",
|
||||
timestamp="2026-02-17T12:00:00Z",
|
||||
guardrail_information=[{"guardrail_name": "pii_masking", "guardrail_status": "success"}],
|
||||
)
|
||||
checker = ComplianceChecker(data)
|
||||
assert len(checker._get_guardrails_by_mode("pre_call")) == 1
|
||||
assert len(checker._get_guardrails_by_mode("post_call")) == 0
|
||||
|
||||
def test_never_reports_false_compliant(self):
|
||||
"""The core invariant: a match reports `mode` satisfied only when every
|
||||
configured branch runs in that mode. So True can never claim a hook the
|
||||
guardrail may not have actually executed. The only allowed error
|
||||
direction is under-reporting."""
|
||||
|
||||
def _branch_modes(value):
|
||||
if isinstance(value, str):
|
||||
return {value}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return {v for v in value if isinstance(v, str)}
|
||||
return set()
|
||||
|
||||
def _guaranteed_modes(g_mode):
|
||||
"""Modes every branch of ``g_mode`` runs in."""
|
||||
if isinstance(g_mode, str):
|
||||
return {g_mode}
|
||||
if isinstance(g_mode, (list, tuple)):
|
||||
sets = [_branch_modes(m) for m in g_mode]
|
||||
return set.intersection(*sets) if sets else set()
|
||||
if isinstance(g_mode, dict):
|
||||
default = g_mode.get("default")
|
||||
if default is None:
|
||||
return set()
|
||||
branches = [default, *(g_mode.get("tags") or {}).values()]
|
||||
sets = [_branch_modes(b) for b in branches]
|
||||
return set.intersection(*sets) if sets else set()
|
||||
return set()
|
||||
|
||||
shapes = [
|
||||
None,
|
||||
"pre_call",
|
||||
"post_call",
|
||||
["pre_call"],
|
||||
["pre_call", "post_call"],
|
||||
[],
|
||||
{"default": "pre_call"},
|
||||
{"default": ["pre_call", "post_call"]},
|
||||
{"default": "pre_call", "tags": {"a": "pre_call"}},
|
||||
{"default": "pre_call", "tags": {"a": "post_call"}},
|
||||
{"default": ["pre_call", "post_call"], "tags": {"a": "pre_call"}},
|
||||
{"tags": {"a": "post_call"}},
|
||||
{},
|
||||
{"default": 123},
|
||||
5,
|
||||
]
|
||||
for g_mode in shapes:
|
||||
for mode in ("pre_call", "post_call", "during_call"):
|
||||
matched = ComplianceChecker._mode_matches(g_mode, mode)
|
||||
if g_mode is None:
|
||||
assert matched is (mode == "pre_call"), (g_mode, mode)
|
||||
continue
|
||||
if matched:
|
||||
assert mode in _guaranteed_modes(g_mode), (g_mode, mode)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue