From 51f6e1b7a5a32e0c85e55d35625f30cdec4abee9 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Fri, 10 Jul 2026 16:41:28 -0700 Subject: [PATCH] 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 --- litellm/proxy/compliance_checks.py | 52 ++++- .../test_compliance_endpoints.py | 215 +++++++++++++++++- 2 files changed, 255 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/compliance_checks.py b/litellm/proxy/compliance_checks.py index 8cde91e32b2..445257a1e01 100644 --- a/litellm/proxy/compliance_checks.py +++ b/litellm/proxy/compliance_checks.py @@ -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).""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py index 2c41b16ba7f..33e45ccb22c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_compliance_endpoints.py @@ -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)