diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index b14b3f1a0af..f32228a6204 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -277,7 +277,7 @@ class GraySwanGuardrail(CustomGuardrail): role: Final = "assistant" if input_type == "response" else "user" merged_tail: Final = ( _MonitorMessage(role="assistant", content=texts[-1], tool_calls=response_tool_calls) - if texts and response_tool_calls + if len(texts) == 1 and response_tool_calls else None ) messages: Final = ( @@ -286,7 +286,7 @@ class GraySwanGuardrail(CustomGuardrail): *((merged_tail,) if merged_tail else ()), *( (_MonitorMessage(role="assistant", tool_calls=response_tool_calls),) - if response_tool_calls and not texts + if response_tool_calls and not merged_tail else () ), ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index a3d7f9e73ef..4c32c0bbd55 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -772,6 +772,33 @@ async def test_post_call_merges_response_text_and_tool_calls_into_one_message() ] +@pytest.mark.asyncio +async def test_post_call_multi_choice_texts_and_tool_calls_stay_split() -> None: + guardrail = _post_call_guardrail() + client = _CapturingClient() + guardrail.async_handler = client + + tool_call = { + "id": "call_send", + "type": "function", + "function": {"name": "send_email", "arguments": '{"to": "cfo@example.com"}'}, + } + await guardrail.apply_guardrail( + inputs={"texts": ["first answer", "second answer"], "tool_calls": [tool_call]}, + request_data={**_REQUEST_DATA, "litellm_logging_obj": _LoggingObj("acompletion")}, + input_type="response", + logging_obj=_LoggingObj("acompletion"), + ) + + messages = list(client.calls[0]["json"]["messages"]) + assert messages == [ + *_REQUEST_DATA["messages"], + {"role": "assistant", "content": "first answer"}, + {"role": "assistant", "content": "second answer"}, + {"role": "assistant", "tool_calls": (tool_call,)}, + ] + + @pytest.mark.asyncio async def test_post_call_prefers_request_route_over_logging_call_type() -> None: guardrail = _post_call_guardrail()