diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 6b1c9a8b2b3..2186edb1b8d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -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 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 9e198e13902..f98f26a9548 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -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, ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 5be0d43c250..e7517be9013 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -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."""