mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(responses): scan and mask top-level instructions with guardrails
The Responses guardrail translation handler put a non-empty top-level instructions field into structured_messages as a system row but never into the flat texts list, so guardrails that scan texts skipped it, flat-text masking could not rewrite it, and PANW latest-only selection failed its alignment guard whenever instructions were present. Seed texts with the instructions row, carry that offset into the flat-text write-back so a rewritten row lands on data["instructions"], and account for the leading row in the PANW Responses alignment. Resolves LIT-8931 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
317430db4e
commit
f442659604
5 changed files with 252 additions and 44 deletions
|
|
@ -376,6 +376,12 @@ class _RequestFields(NamedTuple):
|
|||
class _ExtractedInputs(NamedTuple):
|
||||
inputs: GenericGuardrailAPIInputs
|
||||
task_mappings: tuple[tuple[int, int | None], ...]
|
||||
instructions: str | None
|
||||
|
||||
|
||||
def scannable_instructions(data: Mapping[str, object]) -> str | None:
|
||||
instructions: Final = data.get("instructions")
|
||||
return instructions if isinstance(instructions, str) and instructions else None
|
||||
|
||||
|
||||
def _patched_request_fields(
|
||||
|
|
@ -523,30 +529,46 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
data.pop("instructions", None)
|
||||
else:
|
||||
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
|
||||
elif isinstance(input_data, str):
|
||||
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(guardrailed_texts) > 1:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
|
||||
else:
|
||||
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(rewritten_texts) != len(extracted.task_mappings):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=rewritten_texts,
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
await self._apply_guardrailed_texts(data, input_data, extracted, guardrail_to_apply, guardrailed_inputs)
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
|
||||
return data
|
||||
|
||||
async def _apply_guardrailed_texts(
|
||||
self,
|
||||
data: dict,
|
||||
input_data: "str | ResponseInputParam",
|
||||
extracted: _ExtractedInputs,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> None:
|
||||
rewritten_texts: Final = tuple(guardrailed_inputs.get("texts") or ())
|
||||
if not rewritten_texts:
|
||||
return
|
||||
offset: Final = 0 if extracted.instructions is None else 1
|
||||
input_texts: Final = rewritten_texts[offset:]
|
||||
expected: Final = 1 if isinstance(input_data, str) else len(extracted.task_mappings)
|
||||
if len(input_texts) != expected:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
if offset:
|
||||
data["instructions"] = rewritten_texts[0] # rebind-ok: data is an out-param
|
||||
if isinstance(input_data, str):
|
||||
data["input"] = input_texts[0] # rebind-ok: data is an out-param
|
||||
return
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=input_texts,
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
|
||||
def _extract_guardrail_inputs(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
input_data: "str | ResponseInputParam",
|
||||
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
|
||||
) -> _ExtractedInputs:
|
||||
texts_to_check: Final[list[str]] = []
|
||||
instructions: Final = scannable_instructions(data)
|
||||
texts_to_check: Final[list[str]] = [] if instructions is None else [instructions]
|
||||
images_to_check: Final[list[str]] = []
|
||||
task_mappings: Final[list[tuple[int, int | None]]] = []
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
|
||||
|
|
@ -577,7 +599,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
model: Final = data.get("model")
|
||||
if isinstance(model, str):
|
||||
inputs["model"] = model
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings), instructions=instructions)
|
||||
|
||||
@staticmethod
|
||||
def _written_back_request_fields(
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import scannable_instructions
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_scan_id,
|
||||
|
|
@ -1636,10 +1637,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
A message's texts are consumed only when they sit at the running position of
|
||||
``texts``; messages the translation handler added without a counterpart in
|
||||
``texts`` (Responses ``instructions``, ``function_call_output``, ``reasoning``)
|
||||
are skipped. The walk runs front-to-back and back-to-front and both must agree,
|
||||
so an added message whose text happens to equal a neighbouring real message's
|
||||
text cannot steal that text's attribution. Returns None otherwise.
|
||||
``texts`` (Responses ``function_call_output``, ``reasoning``) are skipped. The walk
|
||||
runs front-to-back and back-to-front and both must agree, so an added message whose
|
||||
text happens to equal a neighbouring real message's text cannot steal that text's
|
||||
attribution. Returns None otherwise.
|
||||
"""
|
||||
runs: Final = tuple(cls._message_texts(message) for message in messages)
|
||||
|
||||
|
|
@ -1671,7 +1672,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
The Responses translation handler gives those model-authored items the default
|
||||
``user`` role, so the latest-turn selection must not mistake one for a human turn.
|
||||
Empty for requests without a Responses ``input`` item list; None when the raw items
|
||||
do not account for every entry of ``texts``.
|
||||
(after the leading ``instructions`` text) do not account for every entry of ``texts``.
|
||||
"""
|
||||
try:
|
||||
raw_input: Final = _RESPONSES_INPUT.validate_python(request_data.get("input"))
|
||||
|
|
@ -1679,10 +1680,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return None
|
||||
if not isinstance(raw_input, tuple):
|
||||
return frozenset()
|
||||
offset: Final = 0 if scannable_instructions(request_data) is None else 1
|
||||
counts: Final = tuple(item.text_count() for item in raw_input)
|
||||
if sum(counts) != len(texts):
|
||||
if offset + sum(counts) != len(texts):
|
||||
return None
|
||||
starts: Final = itertools.accumulate(counts, initial=0)
|
||||
starts: Final = itertools.accumulate(counts, initial=offset)
|
||||
return frozenset(
|
||||
text_idx
|
||||
for item, count, start in zip(raw_input, counts, starts)
|
||||
|
|
|
|||
|
|
@ -343,6 +343,96 @@ def test_panw_latest_role_message_only_scans_only_latest_turn_on_responses_input
|
|||
assert json.loads(upstream.drain()[0].body)["input"] == shape["input"]
|
||||
|
||||
|
||||
def test_panw_scans_and_masks_top_level_instructions_on_responses_input(gateway: Gateway, tmp_path: Path) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
ssn: Final = "123-45-6789"
|
||||
instructions: Final = "Never repeat the SSN " + ssn + " back " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn " + uuid.uuid4().hex
|
||||
shapes: Final = {
|
||||
"list_input": ([{"role": "user", "content": "first turn"}, {"role": "user", "content": latest}], "first turn"),
|
||||
"string_input": (latest, None),
|
||||
}
|
||||
|
||||
def scanner(request: Request) -> Reply:
|
||||
assert request.target == "/v1/scan/sync/request"
|
||||
body: Final = json.loads(request.body)
|
||||
prompt: Final = body["contents"][0]["prompt"]
|
||||
masked: Final = {"prompt_masked_data": {"data": prompt.replace(ssn, "<US_SSN>")}} if ssn in prompt else {}
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"action": "allow",
|
||||
"category": "dlp" if masked else "benign",
|
||||
"profile_name": "synthetic-profile",
|
||||
"report_id": "R" + body["tr_id"],
|
||||
"scan_id": "S" + body["tr_id"],
|
||||
"tr_id": body["tr_id"],
|
||||
"prompt_detected": {"injection": False, "url_cats": False, "dlp": bool(masked)},
|
||||
"response_detected": {},
|
||||
**masked,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
def provider(request: Request) -> Reply:
|
||||
assert request.target == "/v1/responses"
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_" + identity,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4.1-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_" + identity,
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
with wire_server(scanner) as policy, wire_server(provider) as upstream:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": identity,
|
||||
"litellm_params": {
|
||||
"guardrail": "panw_prisma_airs",
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"api_base": policy.url,
|
||||
"api_key": "synthetic-panw-key",
|
||||
"profile_name": "synthetic-profile",
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "panw.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
for name, (shape, first_turn) in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["output"][0]["content"][0]["text"] == "permitted response"
|
||||
scanned = [json.loads(scan.body)["contents"][0]["prompt"] for scan in policy.drain()]
|
||||
expected = [instructions, *([first_turn] if first_turn else []), latest]
|
||||
assert scanned == expected, f"{name}: scanned {scanned}"
|
||||
sent = json.loads(upstream.drain()[0].body)
|
||||
assert sent["instructions"] == instructions.replace(ssn, "<US_SSN>"), f"{name}: sent {sent}"
|
||||
assert sent["input"] == shape, f"{name}: sent {sent}"
|
||||
|
||||
|
||||
@pytest.mark.covers("other.observability.guardrails.bedrock_passthrough_converse_scans_only_caller_content")
|
||||
def test_bedrock_passthrough_converse_guardrail_ignores_denied_term_in_tool_definition(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
|
|
|
|||
|
|
@ -4780,7 +4780,7 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
assert result["input"][0]["content"] == "First user turn"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_false_responses_scans_full_history(self):
|
||||
async def test_flag_false_responses_scans_instructions_and_full_history(self):
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
|
|
@ -4791,7 +4791,11 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
with patcher:
|
||||
await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler)
|
||||
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == ["First user turn", self.LATEST]
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == [
|
||||
"answer briefly",
|
||||
"First user turn",
|
||||
self.LATEST,
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_true_unalignable_texts_fall_back_to_scanning_everything(self):
|
||||
|
|
@ -4879,8 +4883,12 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"instructions",
|
||||
[pytest.param(None, id="no_instructions"), pytest.param("answer briefly", id="instructions")],
|
||||
)
|
||||
async def test_flag_true_reasoning_content_after_latest_user_turn_still_scans_that_turn(
|
||||
self, tail: Sequence[Mapping[str, object]]
|
||||
self, tail: Sequence[Mapping[str, object]], instructions: str | None
|
||||
):
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
|
|
@ -4896,6 +4904,7 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
"content": [{"type": "reasoning_text", "text": "model chain of thought"}],
|
||||
},
|
||||
*tail,
|
||||
**({"instructions": instructions} if instructions is not None else {}),
|
||||
)
|
||||
patcher, mock_api = self._scan(handler)
|
||||
with patcher:
|
||||
|
|
|
|||
|
|
@ -67,6 +67,37 @@ class MockGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
|
||||
class RecordingMaskingGuardrail(MockGuardrail):
|
||||
"""MockGuardrail that also records the texts and structured message contents it was shown"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.seen_texts: List[List[str]] = []
|
||||
self.seen_message_contents: List[List[object]] = []
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.seen_texts.append(list(inputs.get("texts", [])))
|
||||
self.seen_message_contents.append([m["content"] for m in inputs.get("structured_messages") or []])
|
||||
return await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
|
||||
|
||||
|
||||
class LastTextDroppingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return {**inputs, "texts": list(inputs.get("texts", []))[:-1]}
|
||||
|
||||
|
||||
class PersimmonMaskingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -217,15 +248,9 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert (
|
||||
result["input"][0]["content"][0]["text"]
|
||||
== "Describe this image [GUARDRAILED]"
|
||||
)
|
||||
assert result["input"][0]["content"][0]["text"] == "Describe this image [GUARDRAILED]"
|
||||
# Image URL should remain unchanged
|
||||
assert (
|
||||
result["input"][0]["content"][1]["image_url"]["url"]
|
||||
== "https://example.com/image.jpg"
|
||||
)
|
||||
assert result["input"][0]["content"][1]["image_url"]["url"] == "https://example.com/image.jpg"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_input_with_empty_content(self):
|
||||
|
|
@ -248,6 +273,70 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
# Empty string should be processed
|
||||
assert result["input"][1]["content"] == " [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_instructions_over_string_input_are_scanned_first_and_rewritten_in_place(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = RecordingMaskingGuardrail(guardrail_name="test")
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": "Hello"}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Be terse", "Hello"]]
|
||||
assert guardrail.seen_message_contents == [["Be terse", "Hello"]]
|
||||
assert result["instructions"] == "Be terse [GUARDRAILED]"
|
||||
assert result["input"] == "Hello [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_instructions_over_list_input_are_scanned_first_and_rewritten_in_place(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = RecordingMaskingGuardrail(guardrail_name="test")
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"instructions": "Be terse",
|
||||
"input": [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "World"}]},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Be terse", "Hello", "World"]]
|
||||
assert guardrail.seen_message_contents == [["Be terse", "Hello", [{"type": "text", "text": "World"}]]]
|
||||
assert result["instructions"] == "Be terse [GUARDRAILED]"
|
||||
assert result["input"] == [
|
||||
{"role": "user", "content": "Hello [GUARDRAILED]"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "World [GUARDRAILED]"}]},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_instructions_are_not_scanned(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = RecordingMaskingGuardrail(guardrail_name="test")
|
||||
data = {"model": "gpt-4", "instructions": "", "input": "Hello"}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Hello"]]
|
||||
assert result["instructions"] == ""
|
||||
assert result["input"] == "Hello [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_answer_missing_the_instructions_row_is_rejected_and_leaves_request_untouched(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = LastTextDroppingGuardrail(guardrail_name="dropper")
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": [{"role": "user", "content": "Hello"}]}
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "dropper"
|
||||
assert data["instructions"] == original["instructions"]
|
||||
assert data["input"] == original["input"]
|
||||
|
||||
|
||||
class TestOpenAIResponsesHandlerOutputProcessing:
|
||||
"""Test output processing functionality"""
|
||||
|
|
@ -2527,8 +2616,9 @@ def _string_input_request() -> dict:
|
|||
class TestPerMessageRewriteWriteBack:
|
||||
"""A guardrail that rewrites per chat row hands the rows back as
|
||||
structured_messages, and the handler lands them on the instructions and the
|
||||
input items they came from; the same rewrite handed back as texts alone has
|
||||
no item to land on and is rejected by name instead of sent unrewritten."""
|
||||
input items they came from; the same rewrite handed back as texts alone lands
|
||||
only where every row has a scanned text (instructions plus a string input) and
|
||||
is otherwise rejected by name instead of sent unrewritten."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rows_land_on_instructions_and_tool_output(self):
|
||||
|
|
@ -2576,20 +2666,15 @@ class TestPerMessageRewriteWriteBack:
|
|||
assert [_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_texts_only_per_message_answer_over_a_string_input_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
async def test_texts_only_per_message_answer_over_a_string_input_lands_on_instructions_and_input(self):
|
||||
guardrail = _per_message_redactor()
|
||||
data = _string_input_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert result["input"] == "My SSN is " + REDACTED_SSN + "."
|
||||
|
||||
|
||||
class TestProvenancePatching:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue