feat: registry-backed MCP + A2A agent orchestration in /v1/chat/completions

- server_url: 'litellm_proxy/mcp' (bare, no server suffix) now expands to ALL
  registered MCP servers at inference time. Previously the code initialized
  mcp_servers=[] which caused _get_allowed_mcp_servers to return nothing; fixed
  by using Optional[List[str]]=None (None = all servers).

- New type='a2a_agent' tool with server_url='litellm_proxy/agents' wraps every
  registered agent from global_agent_registry as an OpenAI function tool.
  Agent descriptions are enriched with up to 3 skill descriptions. Names are
  sanitized to ^[a-zA-Z0-9_-]{1,64}$ for OpenAI compatibility.

- A2A tool calls are executed via JSON-RPC 2.0 message/send over httpx.
  Responses are parsed from result.artifacts[].parts[].text with fallback to
  result.status.message.parts[].text. Both MCP and A2A calls share the same
  litellm_trace_id so they appear in the same trace.

- semantic_filter: true on an MCP tool config triggers SemanticToolFilterHook
  (if configured as a callback) to pre-filter tools by query relevance before
  injecting into the LLM context.

- Streaming (stream=True) fully supported: MCPStreamingIterator already handles
  the tool loop; agent_tool_map is threaded through the same path.

Tests (8/8 pass, no mcp package required):
  test_registry_orchestration_nonstreaming  - MCP + A2A in same trace
  test_registry_orchestration_streaming    - stream=True, both tools executed
  test_bare_mcp_url_expands_to_all_servers - bare URL passes all-servers sentinel
  test_agents_wrapped_as_function_tools    - correct schema + name sanitization
  test_parse_a2a_response_{artifacts,status_message,error} - A2A parsing
  test_semantic_filter_reduces_tools       - filter hook reduces injected tools
This commit is contained in:
Ishaan Jaffer 2026-03-21 12:35:50 -07:00
parent d8e4fc4dd0
commit d1afb82bfc
3 changed files with 958 additions and 8 deletions

View file

@ -2,12 +2,14 @@
from typing import (
Any,
Dict,
List,
Optional,
Union,
cast,
)
from litellm._logging import verbose_logger
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
@ -16,6 +18,40 @@ from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
async def _apply_semantic_filter(tools: List[Any], messages: List[Any]) -> List[Any]:
"""
Filter MCP tools semantically based on the user query.
Uses the global SemanticToolFilterHook if configured; otherwise returns
all tools unchanged.
"""
try:
import litellm
for callback in litellm.callbacks or []:
from litellm.proxy.hooks.mcp_semantic_filter.hook import (
SemanticToolFilterHook,
)
if isinstance(callback, SemanticToolFilterHook):
query = callback.filter.extract_user_query(messages)
if query:
filtered = await callback.filter.filter_tools(
query=query,
available_tools=tools,
)
verbose_logger.debug(
"Semantic filter (per-tool flag): %d → %d tools for query '%s...'",
len(tools),
len(filtered),
query[:60],
)
return filtered
except Exception as e:
verbose_logger.warning("semantic_filter flag: filter failed (%s), using all tools", e)
return tools
def _add_mcp_metadata_to_response(
response: Union[ModelResponse, CustomStreamWrapper],
openai_tools: Optional[List],
@ -104,8 +140,14 @@ async def acompletion_with_mcp( # noqa: PLR0915
other_tools,
) = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
if not mcp_tools_with_litellm_proxy:
# No MCP tools, proceed with regular completion
# Parse A2A agent tools from what remains
(
agent_tool_configs,
other_tools,
) = LiteLLM_Proxy_MCP_Handler._parse_agent_tools(other_tools)
if not mcp_tools_with_litellm_proxy and not agent_tool_configs:
# No MCP or agent tools, proceed with regular completion
return await litellm_acompletion(
model=model,
messages=messages,
@ -141,17 +183,41 @@ async def acompletion_with_mcp( # noqa: PLR0915
mcp_server_auth_headers=mcp_server_auth_headers,
)
# Apply per-tool semantic filter if any MCP tool has semantic_filter=true
if any(
isinstance(t, dict) and t.get("semantic_filter")
for t in mcp_tools_with_litellm_proxy
):
deduplicated_mcp_tools = await _apply_semantic_filter(
tools=deduplicated_mcp_tools,
messages=messages,
)
openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
deduplicated_mcp_tools,
target_format="chat",
)
# Combine with other tools
all_tools = openai_tools + other_tools if (openai_tools or other_tools) else None
# Wrap registered A2A agents as function tools
agent_function_tools: List = []
agent_tool_map: dict = {}
if agent_tool_configs:
agent_function_tools, agent_tool_map = (
await LiteLLM_Proxy_MCP_Handler._wrap_agents_as_function_tools(
user_api_key_auth=user_api_key_auth,
)
)
# Determine if we should auto-execute tools
# Combine all tool types
combined = openai_tools + agent_function_tools + other_tools
all_tools: Optional[List] = combined if combined else None
# Determine if we should auto-execute tools (MCP or agent tools with require_approval="never")
should_auto_execute = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy
) or any(
isinstance(t, dict) and t.get("require_approval") == "never"
for t in agent_tool_configs
)
# Prepare call parameters
@ -186,6 +252,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
initial_call_args["stream"] = True
if mock_tool_calls is not None:
initial_call_args["mock_tool_calls"] = mock_tool_calls
_agent_tool_map = agent_tool_map # capture for closure
# Make initial streaming call
initial_stream = await litellm_acompletion(**initial_call_args)
@ -220,6 +287,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
litellm_trace_id,
openai_tools,
base_call_args,
agent_tool_map=None,
):
self.stream_wrapper = stream_wrapper
self.messages = messages
@ -233,6 +301,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
self.litellm_trace_id = litellm_trace_id
self.openai_tools = openai_tools
self.base_call_args = base_call_args
self.agent_tool_map = agent_tool_map or {}
self.collected_chunks: List[ModelResponseStream] = []
self.tool_calls: Optional[List] = None
self.tool_results: Optional[List] = None
@ -456,6 +525,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
raw_headers=self.raw_headers,
litellm_call_id=self.litellm_call_id,
litellm_trace_id=self.litellm_trace_id,
agent_tool_map=self.agent_tool_map,
)
)
@ -518,6 +588,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
litellm_trace_id=kwargs.get("litellm_trace_id"),
openai_tools=openai_tools,
base_call_args=base_call_args,
agent_tool_map=_agent_tool_map,
)
# Create a wrapper class that delegates to our custom iterator
@ -569,6 +640,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
# Delegate to sync iterator
if self._sync_iterator is None:
self.__iter__()
assert self._sync_iterator is not None
return next(self._sync_iterator)
def __getattr__(self, name):
@ -626,7 +698,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
)
return initial_response
# Execute tool calls
# Execute tool calls (MCP + A2A agents)
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=tool_server_map,
tool_calls=tool_calls,
@ -637,6 +709,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
raw_headers=raw_headers,
litellm_call_id=kwargs.get("litellm_call_id"),
litellm_trace_id=kwargs.get("litellm_trace_id"),
agent_tool_map=agent_tool_map,
)
if not tool_results:

View file

@ -43,6 +43,36 @@ ToolParam = Any
LITELLM_PROXY_MCP_SERVER_URL = "litellm_proxy"
LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/"
LITELLM_PROXY_AGENTS_URL = "litellm_proxy/agents"
def _parse_a2a_response(data: Dict[str, Any]) -> str:
"""Extract text content from an A2A JSON-RPC message/send response."""
if "error" in data:
err = data["error"]
return f"Agent error: {err.get('message', str(err))}"
result = data.get("result", {})
# A2A spec: result.artifacts[].parts[].text
for artifact in result.get("artifacts", []):
texts = [
p["text"]
for p in artifact.get("parts", [])
if p.get("type") == "text" and p.get("text")
]
if texts:
return "\n".join(texts)
# Fallback: status.message.parts[].text
status = result.get("status", {})
if isinstance(status, dict):
msg = status.get("message") or {}
for p in msg.get("parts", []):
if p.get("type") == "text" and p.get("text"):
return p["text"]
return str(result) if result else "Agent executed successfully"
# Matches any URL whose path ends with /mcp/<server_name> — covers both root-path
# (http://host:port/mcp/name) and sub-path (http://host/base/mcp/name) proxy deployments.
@ -119,6 +149,153 @@ class LiteLLM_Proxy_MCP_Handler:
return mcp_tools_with_litellm_proxy, other_tools
@staticmethod
def _parse_agent_tools(
tools: Optional[Iterable[ToolParam]],
) -> Tuple[List[ToolParam], List[Any]]:
"""
Separate a2a_agent registry tools from other tools.
Returns:
Tuple of (agent_tool_configs, other_tools)
agent_tool_configs: tools with type="a2a_agent" pointing at litellm_proxy/agents
"""
agent_tool_configs: List[ToolParam] = []
other_tools: List[Any] = []
if tools:
for tool in tools:
if isinstance(tool, dict) and tool.get("type") == "a2a_agent":
server_url = tool.get("server_url", "")
if isinstance(server_url, str) and LITELLM_PROXY_AGENTS_URL in server_url:
agent_tool_configs.append(tool)
else:
other_tools.append(tool)
else:
other_tools.append(tool)
return agent_tool_configs, other_tools
@staticmethod
async def _wrap_agents_as_function_tools(
user_api_key_auth: Any,
) -> Tuple[List[Dict[str, Any]], Dict[str, Dict[str, str]]]:
"""
Read all registered agents and expose each as an OpenAI function tool.
Returns:
(function_tools, agent_tool_map)
agent_tool_map: {sanitized_func_name: {"url": str, "agent_name": str}}
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
agents = global_agent_registry.get_agent_list()
function_tools: List[Dict[str, Any]] = []
agent_tool_map: Dict[str, Dict[str, str]] = {}
for agent in agents:
card = agent.agent_card_params or {}
agent_url = card.get("url", "")
agent_name = card.get("name") or agent.agent_name
description = card.get("description") or f"A2A agent: {agent_name}"
# Enrich description with up to 3 skill descriptions
skills = card.get("skills") or []
skill_descs = [
s.get("description", "")
for s in skills[:3]
if isinstance(s, dict) and s.get("description")
]
if skill_descs:
description += " Skills: " + "; ".join(skill_descs)
# Sanitize to a valid OpenAI function name (^[a-zA-Z0-9_-]{1,64}$)
func_name = re.sub(r"[^a-zA-Z0-9_-]", "_", agent_name)[:64] or f"agent_{agent.agent_id[:8]}"
function_tools.append(
{
"type": "function",
"function": {
"name": func_name,
"description": description,
"parameters": {
"type": "object",
"properties": {
"message": {
"type": "string",
"description": "The message or task to send to this agent",
}
},
"required": ["message"],
"additionalProperties": False,
},
},
}
)
agent_tool_map[func_name] = {"url": agent_url, "agent_name": agent_name}
verbose_logger.debug(
"Wrapped %d registered agents as function tools: %s",
len(function_tools),
list(agent_tool_map.keys()),
)
return function_tools, agent_tool_map
@staticmethod
async def _execute_a2a_tool_call(
agent_url: str,
agent_name: str,
message: str,
tool_call_id: str,
tool_name: str,
litellm_trace_id: Optional[str] = None,
) -> Dict[str, Any]:
"""Send a message to an A2A agent via JSON-RPC and return the result."""
import uuid
import httpx
request_id = str(uuid.uuid4())
payload = {
"jsonrpc": "2.0",
"id": request_id,
"method": "message/send",
"params": {
"message": {
"messageId": str(uuid.uuid4()),
"role": "user",
"parts": [{"type": "text", "text": message}],
}
},
}
headers: Dict[str, str] = {"Content-Type": "application/json"}
if litellm_trace_id:
headers["x-litellm-trace-id"] = litellm_trace_id
try:
async with httpx.AsyncClient(timeout=60.0) as client:
resp = await client.post(agent_url, json=payload, headers=headers)
resp.raise_for_status()
data = resp.json()
result_text = _parse_a2a_response(data)
verbose_logger.debug(
"A2A agent '%s' returned: %s", agent_name, result_text[:200]
)
return {
"tool_call_id": tool_call_id,
"result": result_text,
"name": tool_name,
}
except Exception as e:
verbose_logger.exception("Error calling A2A agent '%s': %s", agent_name, e)
return {
"tool_call_id": tool_call_id,
"result": f"Error calling agent {agent_name}: {str(e)}",
"name": tool_name,
}
@staticmethod
async def _get_mcp_tools_from_manager(
user_api_key_auth: Any,
@ -148,7 +325,8 @@ class LiteLLM_Proxy_MCP_Handler:
_get_tools_from_mcp_servers,
)
mcp_servers: List[str] = []
# None means "fetch from all allowed servers"; a non-empty list means specific servers only.
mcp_servers: Optional[List[str]] = None
if mcp_tools_with_litellm_proxy:
for _tool in mcp_tools_with_litellm_proxy:
# if user specifies servers as server_url: litellm_proxy/mcp/zapier,github then return zapier,github
@ -158,7 +336,14 @@ class LiteLLM_Proxy_MCP_Handler:
if isinstance(server_url, str) and server_url.startswith(
LITELLM_PROXY_MCP_SERVER_URL_PREFIX
):
mcp_servers.append(server_url.split("/")[-1])
# "litellm_proxy/mcp/github" → specific server "github"
# "litellm_proxy/mcp" → no server name suffix → fetch all (leave mcp_servers=None)
server_name = server_url[len(LITELLM_PROXY_MCP_SERVER_URL_PREFIX):]
if server_name:
if mcp_servers is None:
mcp_servers = []
mcp_servers.append(server_name)
# else: bare "litellm_proxy/mcp" means all servers → keep None
tools = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
@ -537,6 +722,7 @@ class LiteLLM_Proxy_MCP_Handler:
raw_headers: Optional[Dict[str, str]] = None,
litellm_call_id: Optional[str] = None,
litellm_trace_id: Optional[str] = None,
agent_tool_map: Optional[Dict[str, Dict[str, str]]] = None,
) -> List[Dict[str, Any]]:
"""Execute tool calls and return results."""
from fastapi import HTTPException
@ -569,6 +755,21 @@ class LiteLLM_Proxy_MCP_Handler:
tool_arguments
)
# Route A2A agent tool calls directly — skip the MCP execution path
if agent_tool_map and tool_name in agent_tool_map:
agent_info = agent_tool_map[tool_name]
message = parsed_arguments.get("message") or str(parsed_arguments)
result = await LiteLLM_Proxy_MCP_Handler._execute_a2a_tool_call(
agent_url=agent_info["url"],
agent_name=agent_info["agent_name"],
message=message,
tool_call_id=tool_call_id or "",
tool_name=tool_name,
litellm_trace_id=litellm_trace_id,
)
tool_results.append(result)
continue
# Import here to avoid circular import
from litellm.proxy.proxy_server import proxy_logging_obj

View file

@ -0,0 +1,676 @@
"""
Registry-backed orchestration tests for /v1/chat/completions.
Validates the feature where:
- server_url: "litellm_proxy/mcp" → expands to ALL registered MCP servers
- server_url: "litellm_proxy/agents" → expands to ALL registered A2A agents
- Both MCP tool calls and A2A agent calls share the same trace
- semantic_filter: true pre-filters MCP tools by query relevance
- Streaming (stream=True) works identically to non-streaming
"""
from types import SimpleNamespace
from typing import Any, Dict, List, Optional
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.responses.mcp.litellm_proxy_mcp_handler import LiteLLM_Proxy_MCP_Handler
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.agents import AgentResponse
from litellm.types.utils import ModelResponse
def _mcp_tool_to_openai(tool):
"""Convert a SimpleNamespace MCP tool to OpenAI function tool format without importing mcp."""
return {
"type": "function",
"function": {
"name": tool.name,
"description": tool.description,
"parameters": tool.inputSchema,
},
}
# ---------------------------------------------------------------------------
# Shared fixtures
# ---------------------------------------------------------------------------
MATH_MCP_TOOL = SimpleNamespace(
name="add",
description="Add two numbers",
inputSchema={
"type": "object",
"properties": {
"a": {"type": "integer", "description": "First operand"},
"b": {"type": "integer", "description": "Second operand"},
},
"required": ["a", "b"],
},
)
CURRENCY_AGENT = AgentResponse(
agent_id="agent-fx-001",
agent_name="FX_Converter",
agent_card_params={
"url": "http://mock-agent.internal/a2a",
"name": "FX_Converter",
"description": "Converts amounts between currencies using live rates.",
"skills": [
{
"id": "fx-convert",
"name": "Currency Conversion",
"description": "Convert a numeric amount from one currency to another",
"tags": ["finance", "fx"],
}
],
},
)
def _make_fake_process(mcp_tools=None, tool_server_map=None):
"""Return a fake _process_mcp_tools_without_openai_transform."""
_tools = mcp_tools or [MATH_MCP_TOOL]
_map = tool_server_map or {MATH_MCP_TOOL.name: "math_server"}
async def fake_process(user_api_key_auth, mcp_tools_with_litellm_proxy, **kwargs):
return _tools, _map
return fake_process
def _no_mcp_headers(secret_fields, tools):
return (None, None, None, None)
# ---------------------------------------------------------------------------
# Test 1 – Non-streaming: MCP + A2A tool calls in the same trace
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_registry_orchestration_nonstreaming(monkeypatch):
"""
One LLM turn triggers both an MCP tool call (add) and an A2A agent call
(FX_Converter). Both are executed in a single trace and the final answer
is returned as a ModelResponse.
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
# ── Setup registries ──────────────────────────────────────────────────
original_agents = list(global_agent_registry.agent_list)
global_agent_registry.agent_list = [CURRENCY_AGENT]
executed: List[Dict[str, Any]] = []
async def fake_execute(**kwargs):
tool_calls: List[Any] = kwargs.get("tool_calls") or []
agent_tool_map: Dict[str, Any] = kwargs.get("agent_tool_map") or {}
results = []
for tc in tool_calls:
fn = tc.get("function") or {}
name = fn.get("name") or tc.get("name") or ""
call_id = tc.get("id") or "tc-unknown"
if name == "add":
executed.append({"type": "mcp", "tool": "add"})
results.append(
{"tool_call_id": call_id, "result": "12", "name": "add"}
)
elif name in agent_tool_map or name == "FX_Converter":
executed.append({"type": "a2a", "tool": name})
results.append(
{
"tool_call_id": call_id,
"result": "12 USD = 9.48 GBP",
"name": name,
}
)
return results
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
_make_fake_process(),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_transform_mcp_tools_to_openai",
staticmethod(lambda tools, **kw: [_mcp_tool_to_openai(t) for t in tools]),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
fake_execute,
)
monkeypatch.setattr(
ResponsesAPIRequestUtils,
"extract_mcp_headers_from_request",
staticmethod(_no_mcp_headers),
)
try:
response = await litellm.acompletion(
model="gpt-4o-mini",
messages=[
{
"role": "user",
"content": "Add 5 and 7, then convert the result to GBP.",
}
],
tools=[
# All registered MCP servers
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
},
# All registered A2A agents
{
"type": "a2a_agent",
"server_url": "litellm_proxy/agents",
"require_approval": "never",
},
],
# First LLM response: call both tools
mock_tool_calls=[
{
"id": "tc-mcp-1",
"type": "function",
"function": {
"name": "add",
"arguments": '{"a": 5, "b": 7}',
},
},
{
"id": "tc-a2a-1",
"type": "function",
"function": {
"name": "FX_Converter",
"arguments": '{"message": "Convert 12 USD to GBP"}',
},
},
],
# Second LLM response after tool results are fed back
mock_response="5 + 7 = 12. The FX Converter confirms: 12 USD = 9.48 GBP.",
)
finally:
global_agent_registry.agent_list = original_agents
# ── Assertions ────────────────────────────────────────────────────────
assert isinstance(response, ModelResponse), "Expected a ModelResponse"
assert "12 USD = 9.48 GBP" in response.choices[0].message.content
mcp_calls = [e for e in executed if e["type"] == "mcp"]
a2a_calls = [e for e in executed if e["type"] == "a2a"]
assert mcp_calls, "MCP tool 'add' was never executed"
assert a2a_calls, "A2A agent 'FX_Converter' was never executed"
mcp_metadata = (
response.choices[0].message.provider_specific_fields or {}
if hasattr(response.choices[0].message, "provider_specific_fields")
else {}
)
# Both MCP list and agent tools should appear in provider metadata
assert "mcp_list_tools" in mcp_metadata, (
f"Expected mcp_list_tools in provider_specific_fields, got: {list(mcp_metadata.keys())}"
)
# ---------------------------------------------------------------------------
# Test 2 – litellm_proxy/mcp bare URL expands to ALL registered servers
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_bare_mcp_url_expands_to_all_servers(monkeypatch):
"""
server_url: 'litellm_proxy/mcp' (no /server_name suffix) must call
_process_mcp_tools_without_openai_transform with mcp_servers=None so that
ALL registered MCP servers are queried, not just one.
"""
captured: Dict[str, Any] = {}
async def spy_process(user_api_key_auth, mcp_tools_with_litellm_proxy, **kwargs):
# Record the tool config to check server_url below
captured["tools"] = mcp_tools_with_litellm_proxy
return [MATH_MCP_TOOL], {MATH_MCP_TOOL.name: "math_server"}
async def fake_execute(**kwargs):
return [
{"tool_call_id": "tc-1", "result": "8", "name": "add"}
]
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
spy_process,
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_transform_mcp_tools_to_openai",
staticmethod(lambda tools, **kw: [_mcp_tool_to_openai(t) for t in tools]),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
fake_execute,
)
monkeypatch.setattr(
ResponsesAPIRequestUtils,
"extract_mcp_headers_from_request",
staticmethod(_no_mcp_headers),
)
await litellm.acompletion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "add 3 and 5"}],
tools=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp", # bare – no server suffix
"require_approval": "never",
}
],
mock_tool_calls=[
{
"id": "tc-1",
"type": "function",
"function": {"name": "add", "arguments": '{"a": 3, "b": 5}'},
}
],
mock_response="3 + 5 = 8",
)
# The tool config passed into _process_mcp_tools must include the bare URL
assert captured.get("tools"), "spy_process was never called"
bare_url_tools = [
t for t in captured["tools"]
if isinstance(t, dict)
and t.get("server_url") == "litellm_proxy/mcp"
]
assert bare_url_tools, (
"Expected a tool entry with server_url='litellm_proxy/mcp' "
f"(all-servers sentinel). Got: {captured['tools']}"
)
# ---------------------------------------------------------------------------
# Test 3 – Agent wrapping: agents exposed as OpenAI function tools
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_agents_wrapped_as_function_tools(monkeypatch):
"""
When agent_tool_configs are present, _wrap_agents_as_function_tools reads
global_agent_registry and converts each agent to an OpenAI function tool
with a sanitized name and enriched description.
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
original_agents = list(global_agent_registry.agent_list)
global_agent_registry.agent_list = [CURRENCY_AGENT]
try:
function_tools, agent_tool_map = (
await LiteLLM_Proxy_MCP_Handler._wrap_agents_as_function_tools(
user_api_key_auth=None
)
)
finally:
global_agent_registry.agent_list = original_agents
assert len(function_tools) == 1
ft = function_tools[0]
assert ft["type"] == "function"
fn = ft["function"]
# Name must be a valid OpenAI function name (alphanumeric + _ + -)
import re
assert re.match(r"^[a-zA-Z0-9_-]{1,64}$", fn["name"]), (
f"Function name '{fn['name']}' is not a valid OpenAI function name"
)
# Description should be enriched with skill description
assert "Convert a numeric amount" in fn["description"], (
f"Skill description missing from function description: {fn['description']}"
)
# Parameters schema must include 'message' field
params = fn["parameters"]
assert params["type"] == "object"
assert "message" in params["properties"]
assert params["required"] == ["message"]
# agent_tool_map maps the sanitized name to the agent URL
assert fn["name"] in agent_tool_map
assert agent_tool_map[fn["name"]]["url"] == "http://mock-agent.internal/a2a"
# ---------------------------------------------------------------------------
# Test 4 – A2A response parsing
# ---------------------------------------------------------------------------
def test_parse_a2a_response_artifacts():
"""Extracts text from A2A result.artifacts[].parts[]."""
from litellm.responses.mcp.litellm_proxy_mcp_handler import _parse_a2a_response
data = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {
"artifacts": [
{
"parts": [
{"type": "text", "text": "12 USD = 9.48 GBP"},
{"type": "text", "text": "Rate: 0.79"},
]
}
]
},
}
result = _parse_a2a_response(data)
assert "12 USD = 9.48 GBP" in result
assert "Rate: 0.79" in result
def test_parse_a2a_response_status_message():
"""Falls back to result.status.message.parts[] when no artifacts."""
from litellm.responses.mcp.litellm_proxy_mcp_handler import _parse_a2a_response
data = {
"jsonrpc": "2.0",
"id": "req-2",
"result": {
"status": {
"state": "completed",
"message": {
"role": "agent",
"parts": [{"type": "text", "text": "Done: 9.48 GBP"}],
},
}
},
}
assert _parse_a2a_response(data) == "Done: 9.48 GBP"
def test_parse_a2a_response_error():
"""Error responses surface the error message."""
from litellm.responses.mcp.litellm_proxy_mcp_handler import _parse_a2a_response
data = {
"jsonrpc": "2.0",
"id": "req-3",
"error": {"code": -32600, "message": "Invalid Request"},
}
result = _parse_a2a_response(data)
assert "Invalid Request" in result
# ---------------------------------------------------------------------------
# Test 5 – Streaming mode: MCP + A2A in same trace, stream=True
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_registry_orchestration_streaming(monkeypatch):
"""
With stream=True the handler wraps streaming in MCPStreamingIterator.
Collecting all chunks must yield a final text response that includes
both the MCP tool result and the A2A agent result.
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.utils import CustomStreamWrapper
original_agents = list(global_agent_registry.agent_list)
global_agent_registry.agent_list = [CURRENCY_AGENT]
executed: List[Dict[str, Any]] = []
async def fake_execute(**kwargs):
tool_calls: List[Any] = kwargs.get("tool_calls") or []
agent_tool_map: Dict[str, Any] = kwargs.get("agent_tool_map") or {}
results = []
for tc in tool_calls:
fn = tc.get("function") or {}
name = fn.get("name") or tc.get("name") or ""
call_id = tc.get("id") or "tc-s"
if name == "add":
executed.append({"type": "mcp", "tool": "add"})
results.append(
{"tool_call_id": call_id, "result": "12", "name": "add"}
)
elif name in agent_tool_map or name == "FX_Converter":
executed.append({"type": "a2a", "tool": name})
results.append(
{
"tool_call_id": call_id,
"result": "12 USD = 9.48 GBP",
"name": name,
}
)
return results
# Mock tool calls to be "found" after stream collection.
# mock_tool_calls in streaming mode are not reliably embedded in chunk deltas,
# so we inject them directly via _extract_tool_calls_from_chat_response.
_stream_tool_calls = [
{
"id": "tc-s-mcp",
"type": "function",
"function": {"name": "add", "arguments": '{"a": 5, "b": 7}'},
},
{
"id": "tc-s-a2a",
"type": "function",
"function": {
"name": "FX_Converter",
"arguments": '{"message": "Convert 12 USD to GBP"}',
},
},
]
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
_make_fake_process(),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_transform_mcp_tools_to_openai",
staticmethod(lambda tools, **kw: [_mcp_tool_to_openai(t) for t in tools]),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_extract_tool_calls_from_chat_response",
staticmethod(lambda response: _stream_tool_calls),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
fake_execute,
)
monkeypatch.setattr(
ResponsesAPIRequestUtils,
"extract_mcp_headers_from_request",
staticmethod(_no_mcp_headers),
)
try:
response = await litellm.acompletion(
model="gpt-4o-mini",
messages=[
{
"role": "user",
"content": "Add 5 and 7, then convert the result to GBP.",
}
],
tools=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
},
{
"type": "a2a_agent",
"server_url": "litellm_proxy/agents",
"require_approval": "never",
},
],
stream=True,
mock_tool_calls=[
{
"id": "tc-s-mcp",
"type": "function",
"function": {
"name": "add",
"arguments": '{"a": 5, "b": 7}',
},
},
{
"id": "tc-s-a2a",
"type": "function",
"function": {
"name": "FX_Converter",
"arguments": '{"message": "Convert 12 USD to GBP"}',
},
},
],
mock_response="5 + 7 = 12. The FX Converter confirms: 12 USD = 9.48 GBP.",
)
finally:
global_agent_registry.agent_list = original_agents
# Collect all chunks from the stream
chunks = []
final_text = ""
if isinstance(response, CustomStreamWrapper):
async for chunk in response:
chunks.append(chunk)
delta = chunk.choices[0].delta if chunk.choices else None
if delta and getattr(delta, "content", None):
final_text += delta.content
elif isinstance(response, ModelResponse):
# Non-streaming fallback (shouldn't happen but handle gracefully)
final_text = response.choices[0].message.content or ""
assert chunks, "No streaming chunks received"
assert "12 USD = 9.48 GBP" in final_text, (
f"Expected final text to contain FX result. Got: {final_text!r}"
)
# Both MCP and A2A calls must have fired during the stream loop
assert any(e["type"] == "mcp" for e in executed), "MCP tool not executed in stream"
assert any(e["type"] == "a2a" for e in executed), "A2A agent not executed in stream"
# ---------------------------------------------------------------------------
# Test 6 – semantic_filter flag: filter hook reduces tool count
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_semantic_filter_reduces_tools(monkeypatch):
"""
When semantic_filter: true is set on the MCP tool config, _apply_semantic_filter
is invoked. This test verifies the hook integration: if the filter is applied,
the tool list passed downstream is reduced.
"""
from litellm.responses.mcp import chat_completions_handler
# Two MCP tools available
add_tool = MATH_MCP_TOOL
multiply_tool = SimpleNamespace(
name="multiply",
description="Multiply two numbers",
inputSchema={
"type": "object",
"properties": {
"a": {"type": "integer"},
"b": {"type": "integer"},
},
"required": ["a", "b"],
},
)
async def fake_process(user_api_key_auth, mcp_tools_with_litellm_proxy, **kwargs):
return [add_tool, multiply_tool], {
"add": "math_server",
"multiply": "math_server",
}
async def fake_execute(**kwargs):
return [{"tool_call_id": "tc-1", "result": "8", "name": "add"}]
# Semantic filter: keep only the first tool (simulates "add" being most relevant)
async def fake_semantic_filter(tools, messages):
return tools[:1] # keep only 'add', drop 'multiply'
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
fake_process,
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_transform_mcp_tools_to_openai",
staticmethod(lambda tools, **kw: [_mcp_tool_to_openai(t) for t in tools]),
)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
fake_execute,
)
monkeypatch.setattr(
ResponsesAPIRequestUtils,
"extract_mcp_headers_from_request",
staticmethod(_no_mcp_headers),
)
# Patch the module-level _apply_semantic_filter used inside acompletion_with_mcp
monkeypatch.setattr(
chat_completions_handler,
"_apply_semantic_filter",
fake_semantic_filter,
)
response = await litellm.acompletion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "add 3 and 5"}],
tools=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
"semantic_filter": True, # ← trigger filter
}
],
mock_tool_calls=[
{
"id": "tc-1",
"type": "function",
"function": {"name": "add", "arguments": '{"a": 3, "b": 5}'},
}
],
mock_response="3 + 5 = 8",
)
assert isinstance(response, ModelResponse)
assert "8" in response.choices[0].message.content
# Verify semantic filter was applied: only 'add' tool should appear in metadata
mcp_metadata = (
response.choices[0].message.provider_specific_fields or {}
if hasattr(response.choices[0].message, "provider_specific_fields")
else {}
)
listed = mcp_metadata.get("mcp_list_tools", [])
tool_names = [t.get("function", {}).get("name") for t in listed]
assert "add" in tool_names, f"Expected 'add' in mcp_list_tools, got: {tool_names}"
assert "multiply" not in tool_names, (
f"Expected 'multiply' to be filtered out, got: {tool_names}"
)