mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(responses): reject empty guardrail rewrites instead of forwarding raw input
An explicit texts=[] answer from a guardrail now fails the count check and raises UnappliableRequestRewrite like any other misaligned rewrite; only a missing texts key means no rewrite. Types the out-param as dict[str, object] and adds integration coverage for instructions blocking, masking, empty instructions, tool loops, latest-only and concurrent workers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f442659604
commit
1f8adc4c85
4 changed files with 438 additions and 32 deletions
|
|
@ -536,19 +536,20 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
async def _apply_guardrailed_texts(
|
||||
self,
|
||||
data: dict,
|
||||
data: dict[str, object],
|
||||
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:
|
||||
returned_texts: Final = guardrailed_inputs.get("texts")
|
||||
if returned_texts is None:
|
||||
return
|
||||
rewritten_texts: Final = tuple(returned_texts)
|
||||
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:
|
||||
if len(rewritten_texts) != offset + expected:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
if offset:
|
||||
data["instructions"] = rewritten_texts[0] # rebind-ok: data is an out-param
|
||||
|
|
|
|||
|
|
@ -433,6 +433,371 @@ def test_panw_scans_and_masks_top_level_instructions_on_responses_input(gateway:
|
|||
assert sent["input"] == shape, f"{name}: sent {sent}"
|
||||
|
||||
|
||||
_SSN: Final = "123-45-6789"
|
||||
_MASKED_SSN: Final = "<US_SSN>"
|
||||
_DENIED_TERM: Final = "RIGBLOCKME"
|
||||
|
||||
|
||||
def _panw_scanner(request: Request) -> Reply:
|
||||
assert request.target == "/v1/scan/sync/request"
|
||||
body: Final = json.loads(request.body)
|
||||
prompt: Final = body["contents"][0]["prompt"]
|
||||
denied: Final = _DENIED_TERM in prompt
|
||||
masked: Final = {"prompt_masked_data": {"data": prompt.replace(_SSN, _MASKED_SSN)}} if _SSN in prompt else {}
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"action": "block" if denied else "allow",
|
||||
"category": "malicious" if denied else ("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": denied, "url_cats": False, "dlp": bool(masked)},
|
||||
"response_detected": {},
|
||||
**masked,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _responses_provider(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target.endswith("/models"):
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
assert request.target == "/v1/responses", request.target
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_" + uuid.uuid4().hex,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4.1-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_synthetic",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _panw_config(tmp_path: Path, identity: str, policy_url: str, **flags: bool) -> Path:
|
||||
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",
|
||||
**flags,
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "panw.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _scanned_prompts(scans: tuple[Request, ...]) -> list[str]:
|
||||
return [json.loads(scan.body)["contents"][0]["prompt"] for scan in scans]
|
||||
|
||||
|
||||
def _forwarded_bodies(requests: tuple[Request, ...]) -> list[dict[str, object]]:
|
||||
return [json.loads(request.body) for request in requests if request.method == "POST"]
|
||||
|
||||
|
||||
def test_guardrail_denies_responses_request_whose_only_flagged_text_is_in_instructions(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "You are terse and say " + _DENIED_TERM + " " + uuid.uuid4().hex
|
||||
shapes: Final = {"string_input": "say hi", "list_input": [{"role": "user", "content": "say hi"}]}
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) 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 in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 400, f"{name}: {response.text}"
|
||||
assert "Prompt blocked by PANW Prisma AI Security policy" in response.text, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [instructions], name
|
||||
assert _forwarded_bodies(upstream.drain()) == [], (
|
||||
f"{name}: denied instructions must not reach the provider"
|
||||
)
|
||||
|
||||
|
||||
def test_empty_instructions_are_not_scanned_while_input_and_chat_system_masking_are_unchanged(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
secret: Final = "my SSN is " + _SSN + " " + uuid.uuid4().hex
|
||||
masked: Final = secret.replace(_SSN, _MASKED_SSN)
|
||||
|
||||
def chat_provider(request: Request) -> Reply:
|
||||
assert request.target == "/v1/chat/completions"
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_" + identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4.1-mini",
|
||||
"choices": [
|
||||
{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "ok"}}
|
||||
],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
def provider(request: Request) -> Reply:
|
||||
return chat_provider(request) if request.target == "/v1/chat/completions" else _responses_provider(request)
|
||||
|
||||
with wire_server(_panw_scanner) as policy, wire_server(provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url, mask_request_content=True)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) 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 instructions in ("", None):
|
||||
body = {"model": model, "input": secret, **({} if instructions is None else {"instructions": ""})}
|
||||
response = candidate.request("POST", "/v1/responses", body)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [secret], f"instructions={instructions!r}"
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent.get("instructions") == instructions, f"instructions={instructions!r}: sent {sent}"
|
||||
assert sent["input"] == masked, f"instructions={instructions!r}: sent {sent}"
|
||||
|
||||
response = candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "system", "content": secret}, {"role": "user", "content": "hi"}],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [secret, "hi"]
|
||||
(sent_chat,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent_chat["messages"] == [
|
||||
{"role": "system", "content": masked},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
|
||||
|
||||
def test_instructions_masking_lands_next_to_multimodal_and_tool_loop_input_items(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "Never repeat the SSN " + _SSN + " back " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn with " + _SSN + " " + uuid.uuid4().hex
|
||||
image: Final = {"type": "input_image", "image_url": "https://example.test/receipt.png", "detail": "low"}
|
||||
shapes: Final = {
|
||||
"multimodal": [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "first turn"}, image]},
|
||||
{"role": "user", "content": [image, {"type": "input_text", "text": latest}]},
|
||||
],
|
||||
"tool_loop": [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": "tool result with " + _SSN},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
}
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url, mask_request_content=True)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) 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 in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 200, f"{name}: {response.text}"
|
||||
assert _scanned_prompts(policy.drain()) == [instructions, "first turn", latest], name
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions.replace(_SSN, _MASKED_SSN), f"{name}: sent {sent}"
|
||||
expected = json.loads(json.dumps(shape).replace(latest, latest.replace(_SSN, _MASKED_SSN)))
|
||||
assert sent["input"] == expected, f"{name}: sent {sent}"
|
||||
|
||||
|
||||
def test_panw_latest_only_with_instructions_masks_only_the_latest_turn(gateway: Gateway, tmp_path: Path) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "Keep " + _SSN + " confidential " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn with " + _SSN + " " + uuid.uuid4().hex
|
||||
history: Final = ({"role": "user", "content": "first turn"}, {"role": "assistant", "content": "first reply"})
|
||||
shapes: Final = {
|
||||
"plain": [*history, {"role": "user", "content": latest}],
|
||||
"reasoning": [
|
||||
*history,
|
||||
{"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
}
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(
|
||||
tmp_path, identity, policy.url, mask_request_content=True, experimental_use_latest_role_message_only=True
|
||||
)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) 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 in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 200, f"{name}: {response.text}"
|
||||
assert _scanned_prompts(policy.drain()) == [latest], name
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions, f"{name}: latest-only must leave instructions alone"
|
||||
assert sent["input"] == [*shape[:-1], {"role": "user", "content": latest.replace(_SSN, _MASKED_SSN)}], (
|
||||
f"{name}: sent {sent}"
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_latest_only_masks_latest_turn_on_responses_input_with_instructions(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
guardrail_id: Final = "synthetic" + uuid.uuid4().hex[:8]
|
||||
instructions: Final = "Keep " + _SSN + " confidential " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn with " + _SSN + " " + uuid.uuid4().hex
|
||||
|
||||
def guardrail(request: Request) -> Reply:
|
||||
assert request.target == f"/guardrail/{guardrail_id}/version/DRAFT/apply", request.target
|
||||
body: Final = json.loads(request.body)
|
||||
assert body["source"] == "INPUT", body
|
||||
assert body["content"] == [{"text": {"text": latest}}], body
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"outputs": [{"text": latest.replace(_SSN, _MASKED_SSN)}],
|
||||
"assessments": [
|
||||
{
|
||||
"sensitiveInformationPolicy": {
|
||||
"piiEntities": [
|
||||
{"type": "US_SOCIAL_SECURITY_NUMBER", "match": _SSN, "action": "ANONYMIZED"}
|
||||
]
|
||||
}
|
||||
}
|
||||
],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
with wire_server(guardrail) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": identity,
|
||||
"litellm_params": {
|
||||
"guardrail": "bedrock",
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"mask_request_content": True,
|
||||
"experimental_use_latest_role_message_only": True,
|
||||
"guardrailIdentifier": guardrail_id,
|
||||
"guardrailVersion": "DRAFT",
|
||||
"aws_region_name": "us-east-1",
|
||||
"aws_access_key_id": "AKIASYNTHETICGUARDRAIL",
|
||||
"aws_secret_access_key": "synthetic-secret",
|
||||
"aws_bedrock_runtime_endpoint": policy.url,
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "bedrock-instructions.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path, workers=2) 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"
|
||||
)
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"instructions": instructions,
|
||||
"input": [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"role": "assistant", "content": "first reply"},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(policy.drain()) == 1
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions, sent
|
||||
assert sent["input"] == [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"role": "assistant", "content": "first reply"},
|
||||
{"role": "user", "content": latest.replace(_SSN, _MASKED_SSN)},
|
||||
], sent
|
||||
|
||||
|
||||
def test_instructions_masking_holds_under_concurrent_load_across_two_workers(gateway: Gateway, tmp_path: Path) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url, mask_request_content=True)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) 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"
|
||||
)
|
||||
tags: Final = tuple(uuid.uuid4().hex for _ in range(16))
|
||||
|
||||
def send(tag: str) -> httpx.Response:
|
||||
return candidate.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{"model": model, "instructions": "Keep " + _SSN + " private " + tag, "input": "say hi " + tag},
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=8) as pool:
|
||||
responses: Final = tuple(pool.map(send, tags))
|
||||
assert [response.status_code for response in responses] == [200] * len(tags), [
|
||||
response.text for response in responses
|
||||
]
|
||||
sent: Final = {str(body["input"]): body for body in _forwarded_bodies(upstream.drain())}
|
||||
assert sorted(_scanned_prompts(policy.drain())) == sorted(
|
||||
[text for tag in tags for text in ("Keep " + _SSN + " private " + tag, "say hi " + tag)]
|
||||
)
|
||||
assert {tag: sent["say hi " + tag]["instructions"] for tag in tags} == {
|
||||
tag: "Keep " + _MASKED_SSN + " private " + tag for tag in tags
|
||||
}
|
||||
|
||||
|
||||
@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
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Final, cast
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -12,9 +12,9 @@ from pydantic import ValidationError
|
|||
import litellm
|
||||
from litellm.exceptions import Timeout
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import initialize_guardrail
|
||||
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr.crowdstrike_aidr import (
|
||||
CrowdStrikeAIDRGuardrailMissingSecrets,
|
||||
|
|
@ -1805,36 +1805,22 @@ class _MessageShapedGuardrail(CustomGuardrail):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("case", "instructions", "responses_input"),
|
||||
[
|
||||
(
|
||||
"instructions add a system message",
|
||||
"be terse",
|
||||
[{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}],
|
||||
),
|
||||
(
|
||||
"tool items add messages that carry no text",
|
||||
None,
|
||||
[
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]},
|
||||
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "42"},
|
||||
],
|
||||
),
|
||||
],
|
||||
("case", "instructions"),
|
||||
[("tool items add messages that carry no text", None), ("instructions do not rescue the tool desync", "be terse")],
|
||||
)
|
||||
async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
|
||||
case: str,
|
||||
instructions: str | None,
|
||||
responses_input: list[dict[str, object]],
|
||||
) -> None:
|
||||
async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(case: str, instructions: str | None) -> None:
|
||||
"""An unalignable rewrite must fail the request, not forward the raw prompt.
|
||||
|
||||
Skipping the write-back would hand the model the unredacted text, so a
|
||||
guardrail could be bypassed by adding ``instructions`` or a tool call.
|
||||
guardrail could be bypassed by adding a tool call.
|
||||
"""
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
responses_input: list[dict[str, object]] = [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]},
|
||||
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "42"},
|
||||
]
|
||||
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
|
||||
if instructions is not None:
|
||||
data["instructions"] = instructions
|
||||
|
|
@ -1846,21 +1832,27 @@ async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
|
|||
)
|
||||
|
||||
assert "078-05-1120" in str(responses_input), case
|
||||
assert data.get("instructions") == instructions, case
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aligned_rewrite_is_written_back() -> None:
|
||||
"""Matching counts must still redact the input in place."""
|
||||
@pytest.mark.parametrize("instructions", [None, "be terse"])
|
||||
async def test_aligned_rewrite_is_written_back(instructions: str | None) -> None:
|
||||
"""Matching counts must redact the input, and the instructions when present, in place."""
|
||||
responses_input: list[dict[str, object]] = [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}
|
||||
]
|
||||
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
|
||||
if instructions is not None:
|
||||
data["instructions"] = instructions
|
||||
|
||||
await OpenAIResponsesHandler().process_input_messages(
|
||||
data={"model": "gpt-4o", "input": responses_input},
|
||||
data=data,
|
||||
guardrail_to_apply=_MessageShapedGuardrail("my ssn is <US_SSN>"),
|
||||
)
|
||||
|
||||
assert cast(list, responses_input[0]["content"])[0]["text"] == "my ssn is <US_SSN>"
|
||||
assert data.get("instructions") == (None if instructions is None else "my ssn is <US_SSN>")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -98,6 +98,24 @@ class LastTextDroppingGuardrail(CustomGuardrail):
|
|||
return {**inputs, "texts": list(inputs.get("texts", []))[:-1]}
|
||||
|
||||
|
||||
class TextsReplacingGuardrail(CustomGuardrail):
|
||||
"""Answers with the given texts list, or without a texts key at all when given None"""
|
||||
|
||||
def __init__(self, texts: list[str] | None, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.texts = texts
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
answer = {key: value for key, value in inputs.items() if key != "texts"}
|
||||
return answer if self.texts is None else {**answer, "texts": list(self.texts)}
|
||||
|
||||
|
||||
class PersimmonMaskingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -337,6 +355,36 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
assert data["instructions"] == original["instructions"]
|
||||
assert data["input"] == original["input"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("data_input", ["Hello", [{"role": "user", "content": "Hello"}]])
|
||||
async def test_empty_texts_answer_is_rejected_instead_of_forwarding_the_raw_request(self, data_input):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = TextsReplacingGuardrail(guardrail_name="emptier", texts=[])
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": data_input}
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "emptier"
|
||||
assert data["instructions"] == original["instructions"]
|
||||
assert data["input"] == original["input"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("data_input", ["Hello", [{"role": "user", "content": "Hello"}]])
|
||||
async def test_answer_without_texts_key_leaves_instructions_and_input_untouched(self, data_input):
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = TextsReplacingGuardrail(guardrail_name="silent", texts=None)
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": data_input}
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == original["instructions"]
|
||||
assert result["input"] == original["input"]
|
||||
|
||||
|
||||
class TestOpenAIResponsesHandlerOutputProcessing:
|
||||
"""Test output processing functionality"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue