diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 24b5db28571..a8c236cc03c 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -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: diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 7a3934ffdaa..ebaf14547e3 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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/ — 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 diff --git a/tests/test_litellm/responses/mcp/test_registry_orchestration.py b/tests/test_litellm/responses/mcp/test_registry_orchestration.py new file mode 100644 index 00000000000..f8ad913ab62 --- /dev/null +++ b/tests/test_litellm/responses/mcp/test_registry_orchestration.py @@ -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}" + )