mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(responses): keep skipped rows through full-coverage rewrites and align latest-only with skip_system
Trust a guardrail's structured_messages_cover_full_request claim only when it returns as many rows as the full normalized request, otherwise merge the scoped rows back so skipped instructions and system items survive the write-back. Make PANW's Responses reasoning alignment skip-aware so latest-only still picks the latest user turn when system content is excluded from texts. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
09055232d5
commit
21990819ea
4 changed files with 158 additions and 17 deletions
|
|
@ -645,9 +645,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
guardrailed: Final = guardrailed_inputs.get("structured_messages")
|
||||
if guardrailed is None or guardrailed is scoped_structured_messages:
|
||||
return None
|
||||
covers_full_request: Final = (
|
||||
len(scoped_indices) == len(structured_messages)
|
||||
or guardrail_to_apply.structured_messages_cover_full_request()
|
||||
covers_full_request: Final = len(scoped_indices) == len(structured_messages) or (
|
||||
guardrail_to_apply.structured_messages_cover_full_request() and len(guardrailed) == len(structured_messages)
|
||||
)
|
||||
merged: Final = (
|
||||
guardrailed
|
||||
|
|
|
|||
|
|
@ -28,6 +28,8 @@ from litellm.integrations.custom_guardrail import (
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
role_out_of_guardrail_scope,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
|
|
@ -106,9 +108,14 @@ class _ResponsesInputItem(BaseModel):
|
|||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
type: str | None = None
|
||||
role: str | None = None
|
||||
content: str | tuple[_ResponsesContentPart, ...] | None = None
|
||||
|
||||
def text_count(self) -> int:
|
||||
def text_count(self, *, skip_system: bool) -> int:
|
||||
if role_out_of_guardrail_scope(
|
||||
(self.role or "").lower(), skip_system_message=skip_system, skip_tool_message=False
|
||||
):
|
||||
return 0
|
||||
if isinstance(self.content, str):
|
||||
return 1
|
||||
if self.content is None:
|
||||
|
|
@ -1661,18 +1668,20 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
return forward if len(forward) == len(texts) and forward == backward else None
|
||||
|
||||
@classmethod
|
||||
@staticmethod
|
||||
def _reasoning_item_text_indices(
|
||||
cls,
|
||||
texts: Sequence[str],
|
||||
request_data: Mapping[str, object],
|
||||
*,
|
||||
skip_system: bool,
|
||||
) -> frozenset[int] | None:
|
||||
"""Return the ``texts`` indices flattened from Responses ``reasoning`` input items.
|
||||
|
||||
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
|
||||
(after the leading ``instructions`` text) do not account for every entry of ``texts``.
|
||||
(after the leading ``instructions`` text, both minus whatever ``skip_system`` drops)
|
||||
do not account for every entry of ``texts``.
|
||||
"""
|
||||
try:
|
||||
raw_input: Final = _RESPONSES_INPUT.validate_python(request_data.get("input"))
|
||||
|
|
@ -1680,8 +1689,8 @@ 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)
|
||||
offset: Final = 0 if scannable_instructions(request_data, skip_system=skip_system) is None else 1
|
||||
counts: Final = tuple(item.text_count(skip_system=skip_system) for item in raw_input)
|
||||
if offset + sum(counts) != len(texts):
|
||||
return None
|
||||
starts: Final = itertools.accumulate(counts, initial=offset)
|
||||
|
|
@ -1692,9 +1701,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
for text_idx in range(start, start + count)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_latest_user_text_indices(
|
||||
cls,
|
||||
self,
|
||||
texts: Sequence[str],
|
||||
messages: Sequence[AllMessageValues],
|
||||
request_data: Mapping[str, object],
|
||||
|
|
@ -1708,10 +1716,12 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
user/developer message exists, or the latest one carries text that never reached
|
||||
``texts`` (safety fallback to the role-filter scan).
|
||||
"""
|
||||
sources: Final = cls._text_source_message_indices(texts, messages)
|
||||
sources: Final = self._text_source_message_indices(texts, messages)
|
||||
if sources is None:
|
||||
return None
|
||||
reasoning: Final = cls._reasoning_item_text_indices(texts, request_data)
|
||||
reasoning: Final = self._reasoning_item_text_indices(
|
||||
texts, request_data, skip_system=effective_skip_system_message_for_guardrail(self)
|
||||
)
|
||||
if reasoning is None:
|
||||
return None
|
||||
reasoning_messages: Final = frozenset(sources[text_idx] for text_idx in reasoning)
|
||||
|
|
@ -1725,7 +1735,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
if latest_human is None:
|
||||
return None
|
||||
if latest_human not in sources and cls._message_texts(messages[latest_human]):
|
||||
if latest_human not in sources and self._message_texts(messages[latest_human]):
|
||||
return None
|
||||
return frozenset(text_idx for text_idx, source in enumerate(sources) if source == latest_human)
|
||||
|
||||
|
|
|
|||
|
|
@ -4779,6 +4779,38 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
assert result["input"][-1]["content"] == "[MASKED]"
|
||||
assert result["input"][0]["content"] == "First user turn"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"history_tail",
|
||||
[
|
||||
pytest.param((), id="plain"),
|
||||
pytest.param(
|
||||
({"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]},),
|
||||
id="reasoning",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_flag_true_with_skip_system_still_scans_only_the_latest_turn_on_responses(
|
||||
self, history_tail: Sequence[Mapping[str, object]]
|
||||
):
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
|
||||
handler = make_handler(experimental_use_latest_role_message_only=True)
|
||||
handler.skip_system_message_in_guardrail = True
|
||||
request_data = self._responses_request(
|
||||
{"role": "system", "content": "House rules"},
|
||||
*history_tail,
|
||||
{"role": "user", "content": self.LATEST},
|
||||
instructions="answer briefly",
|
||||
)
|
||||
patcher, mock_api = self._scan(handler)
|
||||
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] == [self.LATEST]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_false_responses_scans_instructions_and_full_history(self):
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
|
|
@ -4938,6 +4970,28 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
"thinking",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_true_texts_short_of_the_input_items_fall_back_to_scanning_everything(self):
|
||||
handler = make_handler(experimental_use_latest_role_message_only=True)
|
||||
reasoning = {"type": "reasoning", "id": "rs_1", "content": [{"type": "reasoning_text", "text": "thinking"}]}
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["thinking", self.LATEST],
|
||||
"structured_messages": [{"role": "user", "content": "thinking"}, {"role": "user", "content": self.LATEST}],
|
||||
}
|
||||
request_data: dict[str, object] = {
|
||||
"litellm_call_id": "test-call-id",
|
||||
"input": [
|
||||
{"role": "user", "content": "First user turn"},
|
||||
reasoning,
|
||||
{"role": "user", "content": self.LATEST},
|
||||
],
|
||||
}
|
||||
patcher, mock_api = self._scan(handler)
|
||||
with patcher:
|
||||
await handler.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request")
|
||||
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == ["thinking", self.LATEST]
|
||||
|
||||
|
||||
class TestPanwAirsMcpToolCallWithoutCallId:
|
||||
"""Tests for MCP tool invocations flowing through apply_guardrail without
|
||||
|
|
|
|||
|
|
@ -357,7 +357,9 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
|
||||
@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):
|
||||
async def test_empty_texts_answer_is_rejected_instead_of_forwarding_the_raw_request(
|
||||
self, data_input: str | list[dict[str, str]]
|
||||
) -> None:
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
|
|
@ -374,7 +376,9 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
|
||||
@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):
|
||||
async def test_answer_without_texts_key_leaves_instructions_and_input_untouched(
|
||||
self, data_input: str | list[dict[str, str]]
|
||||
) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = TextsReplacingGuardrail(guardrail_name="silent", texts=None)
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": data_input}
|
||||
|
|
@ -398,7 +402,7 @@ class TestSkipSystemMessageScopesInstructions:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("data_input", ["Hello", [{"role": "user", "content": "Hello"}]])
|
||||
async def test_instructions_are_neither_scanned_nor_rewritten(self, data_input):
|
||||
async def test_instructions_are_neither_scanned_nor_rewritten(self, data_input: str | list[dict[str, str]]) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = _skipping_system(RecordingMaskingGuardrail(guardrail_name="test"))
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": data_input}
|
||||
|
|
@ -472,6 +476,50 @@ class TestSkipSystemMessageScopesInstructions:
|
|||
("user", ["What is the codename?"]),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_coverage_claim_over_only_the_scoped_rows_still_keeps_the_skipped_system_prompt(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Answer from the memo only.",
|
||||
"input": [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "user", "content": "memo " * 400},
|
||||
{"role": "user", "content": "What is the codename?"},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, _skipping_system(ScopedRowsFullCoverageGuardrail()))
|
||||
|
||||
assert result["instructions"] == "Answer from the memo only."
|
||||
assert [(item["role"], _texts(item)) for item in result["input"]] == [
|
||||
("system", ["House rules"]),
|
||||
("user", [COMPRESSED_MARKER]),
|
||||
("user", ["What is the codename?"]),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_coverage_claim_over_the_whole_request_is_installed_without_a_second_merge(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Answer from the memo only.",
|
||||
"input": [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "user", "content": "memo " * 400},
|
||||
{"role": "user", "content": "What is the codename?"},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, _skipping_system(RebuildingFullCoverageGuardrail()))
|
||||
|
||||
assert result["instructions"] == "Answer from the memo only."
|
||||
assert [(item["role"], _texts(item)) for item in result["input"]] == [
|
||||
("system", ["House rules"]),
|
||||
("user", [COMPRESSED_MARKER]),
|
||||
("user", ["What is the codename?"]),
|
||||
]
|
||||
|
||||
|
||||
class TestOpenAIResponsesHandlerOutputProcessing:
|
||||
"""Test output processing functionality"""
|
||||
|
|
@ -2380,6 +2428,36 @@ class StructuredRewriteGuardrail(CustomGuardrail):
|
|||
return {**inputs, "structured_messages": rewritten}
|
||||
|
||||
|
||||
class ScopedRowsFullCoverageGuardrail(StructuredRewriteGuardrail):
|
||||
"""Claims its structured_messages span the whole request but, like CrowdStrike AIDR on a
|
||||
Responses body (no `messages` to rebuild from), only ever returns the scoped rows it was given."""
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
class RebuildingFullCoverageGuardrail(CustomGuardrail):
|
||||
"""Claims full coverage and honours it: rebuilds every conversation row from the raw request,
|
||||
compressing the first user turn, the way CrowdStrike AIDR does on a chat body."""
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return True
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
raw_input = request_data["input"]
|
||||
assert isinstance(raw_input, list)
|
||||
full: list[dict[str, object]] = [{"role": "system", "content": request_data["instructions"]}, *raw_input]
|
||||
first_user = next(i for i, m in enumerate(full) if m.get("role") == "user")
|
||||
rewritten = [{**m, "content": COMPRESSED_MARKER} if i == first_user else m for i, m in enumerate(full)]
|
||||
return {**inputs, "structured_messages": rewritten}
|
||||
|
||||
|
||||
class ToolOutputRewriteGuardrail(CustomGuardrail):
|
||||
"""Guardrail that compresses the first tool-result row, the way Headroom does."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue