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:
yucheng-berri 2026-07-10 16:41:28 -07:00 • committed by GitHub
parent 8f24f2f767
commit 51f6e1b7a5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 255 additions and 12 deletions

View file

@ -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)."""

View file

@ -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)