From a9f040dd430ae0c077469a062464e6e9261c19e9 Mon Sep 17 00:00:00 2001 From: Shin Date: Sat, 7 Feb 2026 02:20:37 +0000 Subject: [PATCH] fix(test): resolve A2A and MCP guardrail test regressions from #20619 Update tests to match new implementation: 1. test_invoke_agent_a2a_adds_litellm_data: - Mock ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic instead of add_litellm_data_to_request (which is no longer used) 2. test_process_input_messages_* (MCP guardrail handler): - Handler now processes mcp_tool_name/mcp_arguments instead of messages - Passes GenericGuardrailAPIInputs(tools=[...]) to guardrails - Updated tests to verify new tool-based guardrail behavior - Renamed test_process_input_messages_skips_when_no_messages to test_process_input_messages_skips_when_no_tool_name --- .../test_mcp_guardrail_handler.py | 69 ++++++++++++++----- .../agent_endpoints/test_a2a_endpoints.py | 65 +++++++++++------ 2 files changed, 95 insertions(+), 39 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py b/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py index 0e150e064c7..ed2f2c24ef8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py @@ -16,44 +16,61 @@ class MockGuardrail(CustomGuardrail): self.return_texts = return_texts self.call_count = 0 self.last_inputs = None + self.last_request_data = None async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): self.call_count += 1 self.last_inputs = inputs + self.last_request_data = request_data if self.return_texts is not None: return {"texts": self.return_texts} - texts = inputs.get("texts", []) - return {"texts": [f"{text} [SAFE]" for text in texts]} + # Return original inputs (no modification for tool-based guardrails) + return inputs @pytest.mark.asyncio -async def test_process_input_messages_updates_content(): - """Handler should update the synthetic message content when guardrail modifies text.""" +async def test_process_input_messages_calls_guardrail_with_tool(): + """Handler should call guardrail with tool definition when mcp_tool_name is present.""" handler = MCPGuardrailTranslationHandler() guardrail = MockGuardrail() - original_content = "Tool: weather\nArguments: {'city': 'tokyo'}" data = { - "messages": [{"role": "user", "content": original_content}], "mcp_tool_name": "weather", + "mcp_arguments": {"city": "tokyo"}, + "mcp_tool_description": "Get weather for a city", } result = await handler.process_input_messages(data, guardrail) - assert result["messages"][0]["content"].endswith("[SAFE]") - assert guardrail.last_inputs == {"texts": [original_content]} + # Guardrail should be called once assert guardrail.call_count == 1 + # Guardrail should receive tool definition in inputs + assert "tools" in guardrail.last_inputs + assert len(guardrail.last_inputs["tools"]) == 1 + + tool = guardrail.last_inputs["tools"][0] + assert tool["type"] == "function" + assert tool["function"]["name"] == "weather" + assert tool["function"]["description"] == "Get weather for a city" + + # Request data should be passed through + assert guardrail.last_request_data == data + + # Result should be the original data (unchanged) + assert result == data + @pytest.mark.asyncio -async def test_process_input_messages_skips_when_no_messages(): - """Handler should skip guardrail invocation if messages array is missing or empty.""" +async def test_process_input_messages_skips_when_no_tool_name(): + """Handler should skip guardrail invocation if mcp_tool_name is missing.""" handler = MCPGuardrailTranslationHandler() guardrail = MockGuardrail() - data = {"mcp_tool_name": "noop"} + # No mcp_tool_name in data - guardrail should not be called + data = {"some_other_field": "value"} result = await handler.process_input_messages(data, guardrail) assert result == data @@ -61,18 +78,34 @@ async def test_process_input_messages_skips_when_no_messages(): @pytest.mark.asyncio -async def test_process_input_messages_handles_empty_guardrail_result(): - """Handler should leave content untouched when guardrail returns no text updates.""" +async def test_process_input_messages_handles_name_alias(): + """Handler should accept 'name' as an alias for 'mcp_tool_name'.""" handler = MCPGuardrailTranslationHandler() - guardrail = MockGuardrail(return_texts=[]) + guardrail = MockGuardrail() - original_content = "Tool: calendar\nArguments: {'date': '2024-12-25'}" data = { - "messages": [{"role": "user", "content": original_content}], - "mcp_tool_name": "calendar", + "name": "calendar", + "arguments": {"date": "2024-12-25"}, } result = await handler.process_input_messages(data, guardrail) - assert result["messages"][0]["content"] == original_content assert guardrail.call_count == 1 + assert guardrail.last_inputs["tools"][0]["function"]["name"] == "calendar" + + +@pytest.mark.asyncio +async def test_process_input_messages_handles_missing_arguments(): + """Handler should handle missing mcp_arguments gracefully.""" + handler = MCPGuardrailTranslationHandler() + guardrail = MockGuardrail() + + data = { + "mcp_tool_name": "simple_tool", + # No mcp_arguments provided + } + + result = await handler.process_input_messages(data, guardrail) + + assert guardrail.call_count == 1 + assert guardrail.last_inputs["tools"][0]["function"]["name"] == "simple_tool" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 9588c3b55c3..203c9ebe2a7 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -1,7 +1,8 @@ """ Mock tests for A2A endpoints. -Tests that invoke_agent_a2a properly integrates with add_litellm_data_to_request. +Tests that invoke_agent_a2a properly integrates with ProxyBaseLLMRequestProcessing +for adding litellm data to requests. """ import sys @@ -13,16 +14,26 @@ import pytest @pytest.mark.asyncio async def test_invoke_agent_a2a_adds_litellm_data(): """ - Test that invoke_agent_a2a calls add_litellm_data_to_request + Test that invoke_agent_a2a calls common_processing_pre_call_logic and the resulting data includes proxy_server_request. """ from litellm.proxy._types import UserAPIKeyAuth - # Track the data passed to add_litellm_data_to_request + # Track the data passed to common_processing_pre_call_logic captured_data = {} - async def mock_add_litellm_data(data, **kwargs): - # Simulate what add_litellm_data_to_request does + async def mock_common_processing( + request, + general_settings, + user_api_key_dict, + proxy_logging_obj, + proxy_config, + route_type, + version, + ): + # Get the data from the processor instance via closure + data = mock_processor_instance.data + # Simulate what common_processing_pre_call_logic does data["proxy_server_request"] = { "url": "http://localhost:4000/a2a/test-agent", "method": "POST", @@ -30,7 +41,7 @@ async def test_invoke_agent_a2a_adds_litellm_data(): "body": dict(data), } captured_data.update(data) - return data + return data, MagicMock() # Returns (data, logging_obj) # Mock response from asend_message mock_response = MagicMock() @@ -46,6 +57,8 @@ async def test_invoke_agent_a2a_adds_litellm_data(): "url": "http://backend-agent:10001", "name": "Test Agent", } + mock_agent.litellm_params = {} + mock_agent.agent_id = "test-agent-id" # Mock request mock_request = MagicMock() @@ -71,32 +84,22 @@ async def test_invoke_agent_a2a_adds_litellm_data(): ) # Try to use real a2a.types if available, otherwise create realistic mocks - # This test focuses on LiteLLM integration, not A2A protocol correctness, - # but we want mocks that behave like the real types to catch usage issues try: from a2a.types import ( MessageSendParams, SendMessageRequest, SendStreamingMessageRequest, ) - - # Real types available - use them - pass except ImportError: - # Real types not available - create realistic mocks - pass - def make_mock_pydantic_class(name): """Create a mock class that behaves like a Pydantic model.""" class MockPydanticClass: def __init__(self, **kwargs): self.__dict__.update(kwargs) - # Store kwargs for model_dump() if needed self._kwargs = kwargs def model_dump(self, mode="json", exclude_none=False): - """Mock model_dump method.""" result = dict(self._kwargs) if exclude_none: result = {k: v for k, v in result.items() if v is not None} @@ -117,14 +120,28 @@ async def test_invoke_agent_a2a_adds_litellm_data(): mock_a2a_types.SendMessageRequest = SendMessageRequest mock_a2a_types.SendStreamingMessageRequest = SendStreamingMessageRequest + # Create mock processor instance to capture data + mock_processor_instance = MagicMock() + mock_processor_instance.common_processing_pre_call_logic = AsyncMock( + side_effect=mock_common_processing + ) + + def mock_processor_init(data): + mock_processor_instance.data = data + return mock_processor_instance + # Patch at the source modules with patch( "litellm.proxy.agent_endpoints.a2a_endpoints._get_agent", return_value=mock_agent, ), patch( - "litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request", - side_effect=mock_add_litellm_data, - ) as mock_add_data, patch( + "litellm.proxy.agent_endpoints.a2a_endpoints.AgentRequestHandler.is_agent_allowed", + new_callable=AsyncMock, + return_value=True, + ), patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + side_effect=mock_processor_init, + ) as mock_processor_class, patch( "litellm.a2a_protocol.create_a2a_client", new_callable=AsyncMock, ), patch( @@ -137,6 +154,9 @@ async def test_invoke_agent_a2a_adds_litellm_data(): ), patch( "litellm.proxy.proxy_server.proxy_config", MagicMock(), + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock(), ), patch( "litellm.proxy.proxy_server.version", "1.0.0", @@ -158,8 +178,11 @@ async def test_invoke_agent_a2a_adds_litellm_data(): user_api_key_dict=mock_user_api_key_dict, ) - # Verify add_litellm_data_to_request was called - mock_add_data.assert_called_once() + # Verify ProxyBaseLLMRequestProcessing was instantiated + mock_processor_class.assert_called_once() + + # Verify common_processing_pre_call_logic was called + mock_processor_instance.common_processing_pre_call_logic.assert_called_once() # Verify model and custom_llm_provider were set assert captured_data.get("model") == "a2a_agent/Test Agent"