From c23c04170eb8a851c4fe02d0cc811b9eb348c862 Mon Sep 17 00:00:00 2001 From: Shin Bot Date: Sat, 7 Feb 2026 02:42:17 +0000 Subject: [PATCH] fix(test): address Greptile review comments - Use dict() conversion for TypedDict access (safer for future changes) - Use .get() for safer dict key access - Improve A2A test to validate call arguments more explicitly - Add comments explaining TypedDict vs Pydantic model distinction --- .../test_mcp_guardrail_handler.py | 23 +++++-- .../agent_endpoints/test_a2a_endpoints.py | 60 ++++++++++++++----- 2 files changed, 63 insertions(+), 20 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 ed2f2c24ef8..4e2e1fd0f35 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 @@ -51,10 +51,15 @@ async def test_process_input_messages_calls_guardrail_with_tool(): 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" + # ChatCompletionToolParam is a TypedDict (dict), so dict access works. + # Convert to dict explicitly to ensure compatibility with any future changes. + tool = dict(guardrail.last_inputs["tools"][0]) + assert tool.get("type") == "function" + + # The function is also a TypedDict (ChatCompletionToolParamFunctionChunk) + function = dict(tool.get("function", {})) + assert function.get("name") == "weather" + assert function.get("description") == "Get weather for a city" # Request data should be passed through assert guardrail.last_request_data == data @@ -91,7 +96,10 @@ async def test_process_input_messages_handles_name_alias(): result = await handler.process_input_messages(data, guardrail) assert guardrail.call_count == 1 - assert guardrail.last_inputs["tools"][0]["function"]["name"] == "calendar" + # Convert to dict for safe access + tool = dict(guardrail.last_inputs["tools"][0]) + function = dict(tool.get("function", {})) + assert function.get("name") == "calendar" @pytest.mark.asyncio @@ -108,4 +116,7 @@ async def test_process_input_messages_handles_missing_arguments(): result = await handler.process_input_messages(data, guardrail) assert guardrail.call_count == 1 - assert guardrail.last_inputs["tools"][0]["function"]["name"] == "simple_tool" + # Convert to dict for safe access + tool = dict(guardrail.last_inputs["tools"][0]) + function = dict(tool.get("function", {})) + assert function.get("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 203c9ebe2a7..b6f8be47220 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -6,7 +6,7 @@ for adding litellm data to requests. """ import sys -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch, call import pytest @@ -19,8 +19,9 @@ async def test_invoke_agent_a2a_adds_litellm_data(): """ from litellm.proxy._types import UserAPIKeyAuth - # Track the data passed to common_processing_pre_call_logic - captured_data = {} + # Track calls to common_processing_pre_call_logic + processing_call_args = {} + returned_data = {} async def mock_common_processing( request, @@ -31,8 +32,16 @@ async def test_invoke_agent_a2a_adds_litellm_data(): route_type, version, ): - # Get the data from the processor instance via closure + # Capture the actual arguments passed to common_processing_pre_call_logic + processing_call_args["request"] = request + processing_call_args["general_settings"] = general_settings + processing_call_args["user_api_key_dict"] = user_api_key_dict + processing_call_args["route_type"] = route_type + processing_call_args["version"] = version + + # Get the data from the processor instance data = mock_processor_instance.data + # Simulate what common_processing_pre_call_logic does data["proxy_server_request"] = { "url": "http://localhost:4000/a2a/test-agent", @@ -40,8 +49,11 @@ async def test_invoke_agent_a2a_adds_litellm_data(): "headers": {}, "body": dict(data), } - captured_data.update(data) - return data, MagicMock() # Returns (data, logging_obj) + + # Store the returned data to verify endpoint uses it + returned_data.update(data) + mock_logging_obj = MagicMock() + return data, mock_logging_obj # Mock response from asend_message mock_response = MagicMock() @@ -51,6 +63,14 @@ async def test_invoke_agent_a2a_adds_litellm_data(): "result": {"status": "success"}, } + # Track what gets passed to asend_message + asend_message_call_args = {} + + async def mock_asend_message(*args, **kwargs): + asend_message_call_args["args"] = args + asend_message_call_args["kwargs"] = kwargs + return mock_response + # Mock agent mock_agent = MagicMock() mock_agent.agent_card_params = { @@ -147,7 +167,7 @@ async def test_invoke_agent_a2a_adds_litellm_data(): ), patch( "litellm.a2a_protocol.asend_message", new_callable=AsyncMock, - return_value=mock_response, + side_effect=mock_asend_message, ), patch( "litellm.proxy.proxy_server.general_settings", {}, @@ -178,16 +198,28 @@ async def test_invoke_agent_a2a_adds_litellm_data(): user_api_key_dict=mock_user_api_key_dict, ) - # Verify ProxyBaseLLMRequestProcessing was instantiated + # Verify ProxyBaseLLMRequestProcessing was instantiated with data dict mock_processor_class.assert_called_once() + init_call_args = mock_processor_class.call_args + assert isinstance(init_call_args[0][0], dict), "Processor should be initialized with a dict" # Verify common_processing_pre_call_logic was called mock_processor_instance.common_processing_pre_call_logic.assert_called_once() + + # Verify the call included correct route_type and version + assert processing_call_args.get("route_type") == "a2a_request" + assert processing_call_args.get("version") == "1.0.0" - # Verify model and custom_llm_provider were set - assert captured_data.get("model") == "a2a_agent/Test Agent" - assert captured_data.get("custom_llm_provider") == "a2a_agent" + # Verify model and custom_llm_provider were set in the data + assert returned_data.get("model") == "a2a_agent/Test Agent" + assert returned_data.get("custom_llm_provider") == "a2a_agent" - # Verify proxy_server_request was added - assert "proxy_server_request" in captured_data - assert captured_data["proxy_server_request"]["method"] == "POST" + # Verify proxy_server_request was added by common_processing_pre_call_logic + assert "proxy_server_request" in returned_data + assert returned_data["proxy_server_request"]["method"] == "POST" + + # Verify the data with proxy_server_request is what gets passed downstream + # (The endpoint should use the returned data from common_processing_pre_call_logic) + assert "metadata" in asend_message_call_args.get("kwargs", {}) or \ + any("proxy_server_request" in str(arg) for arg in asend_message_call_args.get("args", [])), \ + "Data from common_processing_pre_call_logic should be passed to asend_message"