[router][grpc] Add more mcp test cases to responses api (#12749)

This commit is contained in:
Chang Su
2025-11-06 10:20:00 -08:00
committed by GitHub
parent 0c006b8809
commit 2cb42dc184
3 changed files with 266 additions and 223 deletions
@@ -23,7 +23,7 @@ from util import kill_process_tree
class TestGrpcBackend(StateManagementTests, MCPTests): class TestGrpcBackend(StateManagementTests, MCPTests):
"""End to end tests for gRPC backend.""" """End to end tests for gRPC backend (Regular backend with Llama)."""
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -37,7 +37,7 @@ class TestGrpcBackend(StateManagementTests, MCPTests):
num_workers=1, num_workers=1,
tp_size=2, tp_size=2,
policy="round_robin", policy="round_robin",
router_args=["--history-backend", "memory"], router_args=["--history-backend", "memory", "--tool-call-parser", "llama"],
) )
cls.base_url = cls.cluster["base_url"] cls.base_url = cls.cluster["base_url"]
@@ -48,9 +48,6 @@ class TestGrpcBackend(StateManagementTests, MCPTests):
for worker in cls.cluster.get("workers", []): for worker in cls.cluster.get("workers", []):
kill_process_tree(worker.pid) kill_process_tree(worker.pid)
@unittest.skip(
"TODO: transport error, details: [], metadata: MetadataMap { headers: {} }"
)
def test_previous_response_id_chaining(self): def test_previous_response_id_chaining(self):
super().test_previous_response_id_chaining() super().test_previous_response_id_chaining()
@@ -62,18 +59,17 @@ class TestGrpcBackend(StateManagementTests, MCPTests):
def test_mutually_exclusive_parameters(self): def test_mutually_exclusive_parameters(self):
super().test_mutually_exclusive_parameters() super().test_mutually_exclusive_parameters()
@unittest.skip(
"TODO: Pipeline execution failed: Pipeline stage WorkerSelection failed"
)
def test_mcp_basic_tool_call(self):
super().test_mcp_basic_tool_call()
@unittest.skip("TODO: no event fields")
def test_mcp_basic_tool_call_streaming(self): def test_mcp_basic_tool_call_streaming(self):
return super().test_mcp_basic_tool_call_streaming() return super().test_mcp_basic_tool_call_streaming()
# Inherited from MCPTests:
# - test_mcp_basic_tool_call
# - test_mcp_basic_tool_call_streaming
# - test_mixed_mcp_and_function_tools (requires external MCP server)
# - test_mixed_mcp_and_function_tools_streaming (requires external MCP server)
class TestHarmonyBackend(StateManagementTests, MCPTests, FunctionCallingBaseTest):
class TestGrpcHarmonyBackend(StateManagementTests, MCPTests, FunctionCallingBaseTest):
"""End to end tests for Harmony backend.""" """End to end tests for Harmony backend."""
@classmethod @classmethod
@@ -108,168 +104,16 @@ class TestHarmonyBackend(StateManagementTests, MCPTests, FunctionCallingBaseTest
def test_mutually_exclusive_parameters(self): def test_mutually_exclusive_parameters(self):
super().test_mutually_exclusive_parameters() super().test_mutually_exclusive_parameters()
def test_mcp_basic_tool_call(self):
"""Test basic MCP tool call (non-streaming)."""
tools = [
{
"type": "mcp",
"server_label": "deepwiki",
"server_url": "https://mcp.deepwiki.com/mcp",
"require_approval": "never",
}
]
resp = self.create_response(
"What transport protocols does the 2025-03-26 version of the MCP spec (modelcontextprotocol/modelcontextprotocol) support?",
tools=tools,
stream=False,
)
# Should successfully make the request
self.assertEqual(resp.status_code, 200)
data = resp.json()
# Basic response structure
self.assertIn("id", data)
self.assertIn("status", data)
self.assertEqual(data["status"], "completed")
self.assertIn("output", data)
self.assertIn("model", data)
# Verify output array is not empty
output = data["output"]
self.assertIsInstance(output, list)
self.assertGreater(len(output), 0)
# Check for MCP-specific output types
output_types = [item.get("type") for item in output]
# Should have mcp_list_tools - tools are listed before calling
self.assertIn(
"mcp_list_tools", output_types, "Response should contain mcp_list_tools"
)
# Should have at least one mcp_call
mcp_calls = [item for item in output if item.get("type") == "mcp_call"]
self.assertGreater(
len(mcp_calls), 0, "Response should contain at least one mcp_call"
)
# Verify mcp_call structure
for mcp_call in mcp_calls:
self.assertIn("id", mcp_call)
self.assertIn("status", mcp_call)
self.assertEqual(mcp_call["status"], "completed")
self.assertIn("server_label", mcp_call)
self.assertEqual(mcp_call["server_label"], "deepwiki")
self.assertIn("name", mcp_call)
self.assertIn("arguments", mcp_call)
self.assertIn("output", mcp_call)
def test_mcp_basic_tool_call_streaming(self):
"""Test basic MCP tool call (streaming)."""
tools = [
{
"type": "mcp",
"server_label": "deepwiki",
"server_url": "https://mcp.deepwiki.com/mcp",
"require_approval": "never",
}
]
resp = self.create_response(
"What transport protocols does the 2025-03-26 version of the MCP spec (modelcontextprotocol/modelcontextprotocol) support?",
tools=tools,
stream=True,
)
# Should successfully make the request
self.assertEqual(resp.status_code, 200)
events = self.parse_sse_events(resp)
self.assertGreater(len(events), 0)
event_types = [e.get("event") for e in events]
# Check for lifecycle events
self.assertIn(
"response.created", event_types, "Should have response.created event"
)
self.assertIn(
"response.completed", event_types, "Should have response.completed event"
)
# Check for MCP list tools events
self.assertIn(
"response.output_item.added",
event_types,
"Should have output_item.added events",
)
self.assertIn(
"response.mcp_list_tools.in_progress",
event_types,
"Should have mcp_list_tools.in_progress event",
)
self.assertIn(
"response.mcp_list_tools.completed",
event_types,
"Should have mcp_list_tools.completed event",
)
# Check for MCP call events
self.assertIn(
"response.mcp_call.in_progress",
event_types,
"Should have mcp_call.in_progress event",
)
self.assertIn(
"response.mcp_call_arguments.delta",
event_types,
"Should have mcp_call_arguments.delta event",
)
self.assertIn(
"response.mcp_call_arguments.done",
event_types,
"Should have mcp_call_arguments.done event",
)
self.assertIn(
"response.mcp_call.completed",
event_types,
"Should have mcp_call.completed event",
)
# Verify final completed event has full response
completed_events = [e for e in events if e.get("event") == "response.completed"]
self.assertEqual(len(completed_events), 1)
final_response = completed_events[0].get("data", {}).get("response", {})
self.assertIn("id", final_response)
self.assertEqual(final_response.get("status"), "completed")
self.assertIn("output", final_response)
# Verify final output contains expected items
final_output = final_response.get("output", [])
final_output_types = [item.get("type") for item in final_output]
self.assertIn("mcp_list_tools", final_output_types)
self.assertIn("mcp_call", final_output_types)
# Verify mcp_call items in final output
mcp_calls = [item for item in final_output if item.get("type") == "mcp_call"]
self.assertGreater(len(mcp_calls), 0)
for mcp_call in mcp_calls:
self.assertEqual(mcp_call.get("status"), "completed")
self.assertEqual(mcp_call.get("server_label"), "deepwiki")
self.assertIn("name", mcp_call)
self.assertIn("arguments", mcp_call)
self.assertIn("output", mcp_call)
@unittest.skip("TODO: 501 Not Implemented") @unittest.skip("TODO: 501 Not Implemented")
def test_conversation_with_multiple_turns(self): def test_conversation_with_multiple_turns(self):
super().test_conversation_with_multiple_turns() super().test_conversation_with_multiple_turns()
# Inherited from MCPTests:
# - test_mcp_basic_tool_call
# - test_mcp_basic_tool_call_streaming
# - test_mixed_mcp_and_function_tools (requires external MCP server)
# - test_mixed_mcp_and_function_tools_streaming (requires external MCP server)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -36,6 +36,7 @@ class TestOpenaiBackend(
"""End to end tests for OpenAI backend.""" """End to end tests for OpenAI backend."""
api_key = os.environ.get("OPENAI_API_KEY") api_key = os.environ.get("OPENAI_API_KEY")
mcp_validation_mode = "strict" # Enable strict validation for HTTP backend
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -54,6 +55,24 @@ class TestOpenaiBackend(
def tearDownClass(cls): def tearDownClass(cls):
kill_process_tree(cls.cluster["router"].pid) kill_process_tree(cls.cluster["router"].pid)
# Inherited from MCPTests:
# - test_mcp_basic_tool_call (with strict validation)
# - test_mcp_basic_tool_call_streaming (with strict validation)
# - test_mixed_mcp_and_function_tools (requires external MCP server)
# - test_mixed_mcp_and_function_tools_streaming (requires external MCP server)
@unittest.skip(
"Requires external MCP server (deepwiki) - may not be accessible in CI"
)
def test_mixed_mcp_and_function_tools(self):
super().test_mixed_mcp_and_function_tools()
@unittest.skip(
"Requires external MCP server (deepwiki) - may not be accessible in CI"
)
def test_mixed_mcp_and_function_tools_streaming(self):
super().test_mixed_mcp_and_function_tools_streaming()
class TestXaiBackend(StateManagementTests): class TestXaiBackend(StateManagementTests):
"""End to end tests for XAI backend.""" """End to end tests for XAI backend."""
+232 -52
View File
@@ -5,14 +5,24 @@ Tests MCP tool calling in both streaming and non-streaming modes.
These tests should work across all backends that support MCP (OpenAI, XAI). These tests should work across all backends that support MCP (OpenAI, XAI).
""" """
import json
from basic_crud import ResponseAPIBaseTest from basic_crud import ResponseAPIBaseTest
class MCPTests(ResponseAPIBaseTest): class MCPTests(ResponseAPIBaseTest):
"""Tests for MCP tool calling in both streaming and non-streaming modes.""" """Tests for MCP tool calling in both streaming and non-streaming modes."""
# Class attribute to control validation strictness
# Subclasses can override this to enable strict validation
mcp_validation_mode = "relaxed"
def test_mcp_basic_tool_call(self): def test_mcp_basic_tool_call(self):
"""Test basic MCP tool call (non-streaming).""" """Test basic MCP tool call (non-streaming).
Validation strictness is controlled by the class attribute `mcp_validation_mode`.
Set to "strict" in subclasses for additional HTTP-specific validation.
"""
tools = [ tools = [
{ {
"type": "mcp", "type": "mcp",
@@ -70,26 +80,31 @@ class MCPTests(ResponseAPIBaseTest):
self.assertIn("arguments", mcp_call) self.assertIn("arguments", mcp_call)
self.assertIn("output", mcp_call) self.assertIn("output", mcp_call)
# Should have final message output # Strict mode: additional validation for HTTP backends
messages = [item for item in output if item.get("type") == "message"] if self.mcp_validation_mode == "strict":
self.assertGreater( # Should have final message output
len(messages), 0, "Response should contain at least one message" messages = [item for item in output if item.get("type") == "message"]
) self.assertGreater(
len(messages), 0, "Response should contain at least one message"
)
# Verify message structure
for msg in messages:
self.assertIn("content", msg)
self.assertIsInstance(msg["content"], list)
# Verify message structure # Check content has text
for msg in messages: for content_item in msg["content"]:
self.assertIn("content", msg) if content_item.get("type") == "output_text":
self.assertIsInstance(msg["content"], list) self.assertIn("text", content_item)
self.assertIsInstance(content_item["text"], str)
# Check content has text self.assertGreater(len(content_item["text"]), 0)
for content_item in msg["content"]:
if content_item.get("type") == "output_text":
self.assertIn("text", content_item)
self.assertIsInstance(content_item["text"], str)
self.assertGreater(len(content_item["text"]), 0)
def test_mcp_basic_tool_call_streaming(self): def test_mcp_basic_tool_call_streaming(self):
"""Test basic MCP tool call (streaming).""" """Test basic MCP tool call (streaming).
Validation strictness is controlled by the class attribute `mcp_validation_mode`.
Set to "strict" in subclasses for additional HTTP-specific validation.
"""
tools = [ tools = [
{ {
"type": "mcp", "type": "mcp",
@@ -160,28 +175,6 @@ class MCPTests(ResponseAPIBaseTest):
"Should have mcp_call.completed event", "Should have mcp_call.completed event",
) )
# Check for text output events
self.assertIn(
"response.content_part.added",
event_types,
"Should have content_part.added event",
)
self.assertIn(
"response.output_text.delta",
event_types,
"Should have output_text.delta events",
)
self.assertIn(
"response.output_text.done",
event_types,
"Should have output_text.done event",
)
self.assertIn(
"response.content_part.done",
event_types,
"Should have content_part.done event",
)
# Verify final completed event has full response # Verify final completed event has full response
completed_events = [e for e in events if e.get("event") == "response.completed"] completed_events = [e for e in events if e.get("event") == "response.completed"]
self.assertEqual(len(completed_events), 1) self.assertEqual(len(completed_events), 1)
@@ -197,7 +190,6 @@ class MCPTests(ResponseAPIBaseTest):
self.assertIn("mcp_list_tools", final_output_types) self.assertIn("mcp_list_tools", final_output_types)
self.assertIn("mcp_call", final_output_types) self.assertIn("mcp_call", final_output_types)
self.assertIn("message", final_output_types)
# Verify mcp_call items in final output # Verify mcp_call items in final output
mcp_calls = [item for item in final_output if item.get("type") == "mcp_call"] mcp_calls = [item for item in final_output if item.get("type") == "mcp_call"]
@@ -210,19 +202,207 @@ class MCPTests(ResponseAPIBaseTest):
self.assertIn("arguments", mcp_call) self.assertIn("arguments", mcp_call)
self.assertIn("output", mcp_call) self.assertIn("output", mcp_call)
# Verify text deltas combine to final message # Strict mode: additional validation for HTTP backends
text_deltas = [ if self.mcp_validation_mode == "strict":
e.get("data", {}).get("delta", "") # Check for text output events
self.assertIn(
"response.content_part.added",
event_types,
"Should have content_part.added event",
)
self.assertIn(
"response.output_text.delta",
event_types,
"Should have output_text.delta events",
)
self.assertIn(
"response.output_text.done",
event_types,
"Should have output_text.done event",
)
self.assertIn(
"response.content_part.done",
event_types,
"Should have content_part.done event",
)
self.assertIn("message", final_output_types)
# Verify text deltas combine to final message
text_deltas = [
e.get("data", {}).get("delta", "")
for e in events
if e.get("event") == "response.output_text.delta"
]
self.assertGreater(len(text_deltas), 0, "Should have text deltas")
# Get final text from output_text.done event
text_done_events = [
e for e in events if e.get("event") == "response.output_text.done"
]
self.assertGreater(len(text_done_events), 0)
final_text = text_done_events[0].get("data", {}).get("text", "")
self.assertGreater(len(final_text), 0, "Final text should not be empty")
def test_mixed_mcp_and_function_tools(self):
"""Test mixed MCP and function tools (non-streaming)."""
tools = [
{
"type": "mcp",
"server_url": "https://mcp.deepwiki.com/mcp",
"server_label": "deepwiki",
"require_approval": "never",
},
{
"type": "function",
"name": "get_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
]
resp = self.create_response(
"What is the weather in seattle now?",
tools=tools,
stream=False,
tool_choice="auto",
)
# Should successfully make the request
self.assertEqual(resp.status_code, 200)
data = resp.json()
# Basic response structure
self.assertIn("id", data)
self.assertIn("status", data)
self.assertIn("output", data)
# Verify output array is not empty
output = data["output"]
self.assertIsInstance(output, list)
self.assertGreater(len(output), 0)
# Check for function_call (not mcp_call for get_weather)
function_calls = [
item for item in output if item.get("type") == "function_call"
]
self.assertGreater(
len(function_calls), 0, "Response should contain at least one function_call"
)
# Verify function_call structure for get_weather
weather_call = function_calls[0]
self.assertIn("name", weather_call)
self.assertEqual(weather_call["name"], "get_weather")
self.assertIn("call_id", weather_call)
self.assertIn("arguments", weather_call)
self.assertIn("status", weather_call)
# Parse and verify arguments
args = json.loads(weather_call["arguments"])
self.assertIn("location", args)
self.assertIn("seattle", args["location"].lower())
def test_mixed_mcp_and_function_tools_streaming(self):
"""Test mixed MCP and function tools (streaming)."""
tools = [
{
"type": "mcp",
"server_url": "https://mcp.deepwiki.com/mcp",
"server_label": "deepwiki",
"require_approval": "never",
},
{
"type": "function",
"name": "get_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
]
resp = self.create_response(
"What is the weather in seattle now?",
tools=tools,
stream=True,
tool_choice="auto", # Encourage tool usage
)
# Should successfully make the request
self.assertEqual(resp.status_code, 200)
events = self.parse_sse_events(resp)
self.assertGreater(len(events), 0)
event_types = [e.get("event") for e in events]
# Check for lifecycle events
self.assertIn(
"response.created", event_types, "Should have response.created event"
)
# Should have mcp_list_tools events
self.assertIn(
"response.mcp_list_tools.completed",
event_types,
"Should have mcp_list_tools.completed event",
)
# Should have function_call_arguments events (not mcp_call_arguments)
self.assertIn(
"response.function_call_arguments.delta",
event_types,
"Should have function_call_arguments.delta event for function tools",
)
self.assertIn(
"response.function_call_arguments.done",
event_types,
"Should have function_call_arguments.done event for function tools",
)
# Should NOT have mcp_call_arguments events for function tools
# (get_weather should use function_call_arguments, not mcp_call_arguments)
mcp_call_arg_events = [
e
for e in events for e in events
if e.get("event") == "response.output_text.delta" if e.get("event") == "response.mcp_call_arguments.delta"
and "get_weather" in str(e.get("data", {}))
] ]
self.assertGreater(len(text_deltas), 0, "Should have text deltas") self.assertEqual(
len(mcp_call_arg_events),
0,
"Should NOT emit mcp_call_arguments.delta for function tools (get_weather)",
)
# Get final text from output_text.done event # Verify function_call_arguments.delta event structure
text_done_events = [ func_arg_deltas = [
e for e in events if e.get("event") == "response.output_text.done" e
for e in events
if e.get("event") == "response.function_call_arguments.delta"
] ]
self.assertGreater(len(text_done_events), 0) self.assertGreater(
len(func_arg_deltas), 0, "Should have function_call_arguments.delta events"
)
final_text = text_done_events[0].get("data", {}).get("text", "") # Check that at least one delta event contains location arguments
self.assertGreater(len(final_text), 0, "Final text should not be empty") has_location = False
for event in func_arg_deltas:
data = event.get("data", {})
delta = data.get("delta", "")
if "location" in delta.lower() or "seattle" in delta.lower():
has_location = True
break
self.assertTrue(
has_location,
"function_call_arguments.delta should contain location/seattle",
)