mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
271877ffb5
commit
a9f040dd43
2 changed files with 95 additions and 39 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue