mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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.
This commit is contained in:
parent
1b81bb451d
commit
0f9ca3fd77
2 changed files with 60 additions and 19 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue