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
This commit is contained in:
Shin 2026-02-07 02:20:37 +00:00
parent 271877ffb5
commit a9f040dd43
2 changed files with 95 additions and 39 deletions

View file

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

View file

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