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
This commit is contained in:
Shin Bot 2026-02-07 02:42:17 +00:00
parent a9f040dd43
commit c23c04170e
2 changed files with 63 additions and 20 deletions

View file

@ -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"

View file

@ -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"