mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): apply generic guardrail API compressed structured_messages to the request
The generic guardrail API computed structured_messages rewrites (e.g. Headroom prompt compression) but GenericGuardrailAPIResponse.from_dict dropped them and _build_guardrail_return_inputs never returned them, so the original uncompressed messages still reached the model. Parse and surface structured_messages so the handler writes them back to the outbound request Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
56dafd230f
commit
156656d60c
3 changed files with 93 additions and 0 deletions
|
|
@ -334,6 +334,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
return_inputs["tools"] = guardrail_response.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
if guardrail_response.structured_messages is not None:
|
||||
return_inputs["structured_messages"] = guardrail_response.structured_messages
|
||||
if guardrail_response.stream_holdback_chars is not None:
|
||||
return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars
|
||||
return return_inputs
|
||||
|
|
|
|||
|
|
@ -164,6 +164,7 @@ class GenericGuardrailAPIResponse:
|
|||
texts: Optional[List[str]]
|
||||
images: Optional[List[str]]
|
||||
tools: Optional[List[GuardrailToolParam]]
|
||||
structured_messages: Optional[List[AllMessageValues]]
|
||||
action: str
|
||||
blocked_reason: Optional[str]
|
||||
stream_holdback_chars: Optional[List[int]]
|
||||
|
|
@ -175,6 +176,7 @@ class GenericGuardrailAPIResponse:
|
|||
blocked_reason: Optional[str] = None,
|
||||
images: Optional[List[str]] = None,
|
||||
tools: Optional[List[GuardrailToolParam]] = None,
|
||||
structured_messages: Optional[List[AllMessageValues]] = None,
|
||||
stream_holdback_chars: Optional[List[int]] = None,
|
||||
):
|
||||
self.action = action
|
||||
|
|
@ -182,6 +184,7 @@ class GenericGuardrailAPIResponse:
|
|||
self.texts = texts
|
||||
self.images = images
|
||||
self.tools = tools
|
||||
self.structured_messages = structured_messages
|
||||
# Number of trailing chars, indexed the same as ``texts``, that the
|
||||
# framework must withhold from streaming emission until the next
|
||||
# processing round (word-boundary safety for text transformations).
|
||||
|
|
@ -199,5 +202,6 @@ class GenericGuardrailAPIResponse:
|
|||
texts=data.get("texts"),
|
||||
images=data.get("images"),
|
||||
tools=data.get("tools"),
|
||||
structured_messages=data.get("structured_messages"),
|
||||
stream_holdback_chars=stream_holdback_chars,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1402,6 +1402,93 @@ class TestGenericGuardrailAPIResponseParsing:
|
|||
assert result["texts"] == ["Alice went to Berlin"]
|
||||
assert result["stream_holdback_chars"] == [5]
|
||||
|
||||
def test_from_dict_parses_structured_messages(self):
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIResponse,
|
||||
)
|
||||
|
||||
compressed = [{"role": "user", "content": "compressed"}]
|
||||
response = GenericGuardrailAPIResponse.from_dict(
|
||||
{"action": "GUARDRAIL_INTERVENED", "structured_messages": compressed}
|
||||
)
|
||||
|
||||
assert response.structured_messages == compressed
|
||||
|
||||
def test_from_dict_structured_messages_absent_is_none(self):
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIResponse,
|
||||
)
|
||||
|
||||
response = GenericGuardrailAPIResponse.from_dict({"action": "NONE", "texts": ["hi"]})
|
||||
|
||||
assert response.structured_messages is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_flows_structured_messages_back_to_inputs(self, generic_guardrail):
|
||||
"""A guardrail that rewrites the whole conversation (e.g. prompt compression)
|
||||
returns new structured_messages; they must be surfaced on the returned inputs
|
||||
so the framework can send the compressed conversation to the LLM instead of the
|
||||
original. Regression: previously the response's structured_messages were dropped."""
|
||||
original = [
|
||||
{"role": "user", "content": "a very long log " * 500},
|
||||
{"role": "user", "content": "summarize"},
|
||||
]
|
||||
compressed = [
|
||||
{"role": "user", "content": "a very long log [compressed]"},
|
||||
{"role": "user", "content": "summarize"},
|
||||
]
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"structured_messages": compressed,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
result = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": [m["content"] for m in original], "structured_messages": original},
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == compressed
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compressed_structured_messages_are_written_back_to_request(self, generic_guardrail):
|
||||
"""End-to-end pre-call write-back: when the guardrail returns compressed
|
||||
structured_messages, the chat-completions handler must replace data["messages"]
|
||||
with them so the compressed conversation is what reaches the LLM."""
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import (
|
||||
OpenAIChatCompletionsHandler,
|
||||
)
|
||||
|
||||
compressed = [
|
||||
{"role": "user", "content": "log summary [400 matches compressed to 5]"},
|
||||
{"role": "user", "content": "summarize"},
|
||||
]
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"structured_messages": compressed,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [
|
||||
{"role": "user", "content": "a very long log " * 500},
|
||||
{"role": "user", "content": "summarize"},
|
||||
],
|
||||
}
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
result = await OpenAIChatCompletionsHandler().process_input_messages(
|
||||
data=data,
|
||||
guardrail_to_apply=generic_guardrail,
|
||||
)
|
||||
|
||||
assert result["messages"] == compressed
|
||||
|
||||
|
||||
class TestGenericGuardrailAPIStreamingViaUnified:
|
||||
"""Streaming output checks routed through UnifiedLLMGuardrails."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue