From 0f9ca3fd777c9bc1390a644f372e4f65f6f2c776 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 10:39:16 -0500 Subject: [PATCH] fix(guardrails): redact llm_shield_proxy schema enum and const values `enum` and `const` were skipped by the schema walk, so a value holding PII went to the provider in clear. They now go to the caller's vault rather than the non-restorable one: the model emits the stand-in in its tool arguments or structured output, and restoring the reply turns it back into the value the schema allows, so the call still routes. --- .../llm_shield_proxy/llm_shield_proxy.py | 42 +++++++++++-------- .../guardrail_hooks/test_llm_shield_proxy.py | 37 +++++++++++++++- 2 files changed, 60 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index ff5ca24c6e9..bd2796086db 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -128,17 +128,13 @@ _RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( _RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) # JSON Schema keywords whose value has to reach the model or a validator verbatim, so the -# schema walk leaves them alone: types, formats, patterns, references, required-property -# lists, and `enum` / `const`, which the model must reproduce exactly -- a value redacted -# into the non-restorable vault would come back as a stand-in and break the call. -# Everything else is scanned. +# schema walk leaves them alone: types, formats, patterns, references and +# required-property lists. Everything else is scanned. _SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( ( "type", "format", "pattern", - "enum", - "const", "required", "dependentRequired", "propertyOrdering", @@ -161,6 +157,13 @@ _SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( # whatever the keys around it are called. _SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) +# Keywords holding the literal values the model must reproduce. These go to the CALLER's +# vault, not the privileged one: the model emits the stand-in in its tool arguments or +# structured output, and restoring the reply turns it back into the value the schema +# allows, so the call still routes. In the non-restorable vault it would come back as a +# stand-in no validator accepts. +_SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const")) + # Keywords whose value maps names to subschemas. Their keys are property names, not # keywords, so a property called `type` or `enum` is walked like any other subschema. _SCHEMA_MAP_KEYWORDS: Final = frozenset( @@ -333,14 +336,14 @@ def _collect_text_parts(container: MutableRequest, key: str, slots: _SlotSink) - _collect(part, "text", slots) -def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> None: +def _collect_tool_definitions(data: MutableRequest, slots: _SlotSink, privileged: _SlotSink) -> None: """Tool definitions are application-authored free text bound for the provider. A tool's description and the free text in its parameter schema are where callers put examples and customer context, so they carry PII as often as a prompt does. They are collected into the privileged sink, like a system prompt: redacted outbound, and never - restorable from the reply. Names, types, `enum` and `const` values are left as sent, - because the model has to reproduce them exactly for a call to route. + restorable from the reply. `enum` and `const` values are the exception, and go to the + caller's vault -- see `_SCHEMA_LITERAL_KEYWORDS`. Names and types are left as sent. Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. @@ -353,17 +356,19 @@ def _collect_tool_definitions(data: MutableRequest, privileged: _SlotSink) -> No function = tool.get("function") for holder in (tool, function) if isinstance(function, dict) else (tool,): _collect(holder, "description", privileged) - _collect_schema_text(holder.get("parameters"), privileged) - _collect_schema_text(holder.get("input_schema"), privileged) + _collect_schema_text(holder.get("parameters"), slots, privileged) + _collect_schema_text(holder.get("input_schema"), slots, privileged) -def _collect_schema_text(schema: object, privileged: _SlotSink) -> None: - """Collects the free text in a JSON Schema, at any depth. +def _collect_schema_text(schema: object, slots: _SlotSink, privileged: _SlotSink) -> None: + """Collects the text in a JSON Schema, at any depth. Scan by default: every string is collected except under the keywords in `_SCHEMA_STRUCTURAL_KEYWORDS`, whose values must go out verbatim. A list of keywords *to* collect would leak every one it forgot -- draft-07 `dependencies`, a `$comment`, - a vendor `x-` extension -- which is how this walk started out. + a vendor `x-` extension -- which is how this walk started out. Free text goes to the + privileged sink; `enum` / `const` literals go to the caller's, so the model's use of + them is restored. Structure matters in two places. Under `properties` and the other name -> subschema maps, keys are property names rather than keywords, so a property called `type` is a @@ -389,7 +394,10 @@ def _collect_schema_text(schema: object, privileged: _SlotSink) -> None: for keyword, value in tuple(node.items()): if keyword in _SCHEMA_STRUCTURAL_KEYWORDS: continue - if keyword in _SCHEMA_VALUE_KEYWORDS: + if keyword in _SCHEMA_LITERAL_KEYWORDS: + _collect(node, keyword, slots) + _collect_json_leaves(value, slots, strict=True) + elif keyword in _SCHEMA_VALUE_KEYWORDS: _collect(node, keyword, privileged) _collect_json_leaves(value, privileged, strict=True) elif keyword in _SCHEMA_MAP_KEYWORDS and isinstance(value, dict): @@ -421,7 +429,7 @@ def _collect_output_contracts(data: MutableRequest, slots: _SlotSink, privileged ): if isinstance(wrapper, dict): _collect(wrapper, "description", privileged) - _collect_schema_text(wrapper.get("schema"), privileged) + _collect_schema_text(wrapper.get("schema"), slots, privileged) def _collect_end_user_ids(data: MutableRequest, privileged: _SlotSink) -> None: @@ -1013,7 +1021,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): _collect_responses_fields(data, slots, privileged) _collect_prompt(data, slots) _collect_system(data, privileged) - _collect_tool_definitions(data, privileged) + _collect_tool_definitions(data, slots, privileged) _collect_output_contracts(data, slots, privileged) _collect_end_user_ids(data, privileged) return tuple(slots), tuple(privileged) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 475ab06994a..d909181e541 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -635,7 +635,7 @@ class TestRequestCoverage: def test_tool_schemas_give_up_their_free_text_and_nothing_else(self): """Every string is collected except what must reach the model verbatim: names, - types, formats, patterns, required lists, enum and const values.""" + types, formats, patterns and required lists.""" data = { "tools": [ { @@ -673,8 +673,9 @@ class TestRequestCoverage: } ] } - _, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + assert sorted(text for text, _ in caller) == ["a", "a", "b"], "enum and const go to the caller vault" assert sorted(text for text, _ in privileged) == [ "comment", "default", @@ -691,6 +692,38 @@ class TestRequestCoverage: "vendor", ] + @pytest.mark.asyncio + async def test_enum_values_are_redacted_and_restored_in_the_tool_call(self): + """An enum value holding PII is redacted, and the model's use of the stand-in is + restored in its tool arguments, so the call still carries a value the schema allows.""" + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + shield = _FakeShield({"[EMAIL_1]": "ops@example.com"}) + redact_mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + data = { + "messages": [], + "tools": [ + { + "type": "function", + "function": { + "name": "notify", + "parameters": {"properties": {"to": {"type": "string", "enum": ["ops@example.com"]}}}, + }, + } + ], + } + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + assert redact_mock.call_args_list[0].kwargs["json"]["texts"] == ["ops@example.com"] + assert data["tools"][0]["function"]["parameters"]["properties"]["to"]["enum"] == ["[EMAIL_1]"] + + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + call = SimpleNamespace(function=SimpleNamespace(name="notify", arguments='{"to": "[EMAIL_1]"}')) + reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))]) + reply.choices[0].message.tool_calls = [call] + await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert json.loads(call.function.arguments) == {"to": "ops@example.com"} + def test_schema_nesting_past_the_bound_is_refused(self): schema: dict = {"type": "object", "description": "past-the-bound@example.com"} for _ in range(100):