feat: registry-backed MCP + A2A agent orchestration in /v1/chat/completions (#24321)

* 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

* docs: add MCP + A2A orchestration doc for /chat/completions

* docs: move orchestration doc to A2A Agent Gateway sidebar

* refactor: extract RegistryOrchestrator and MCPStreamingIterator to module-level

- Move A2A orchestration logic (parse, wrap, execute, semantic filter) into
  RegistryOrchestrator class in litellm/proxy/agent_endpoints/registry_orchestrator.py
- Extract MCPStreamingIterator, MCPStreamWrapper, _SyncIteratorWrapper from nested
  closures inside acompletion_with_mcp() to module-level classes in
  chat_completions_handler.py
- Fix 7x repeated verbose_logger inline imports in MCPStreamingIterator methods
- Fix async __aiter__ (should be sync) in MCPStreamingIterator
- Remove dead rules_obj and duplicate proxy_logging_obj import from _execute_tool_calls
- Update tests to import from new locations

* refactor: execute_a2a_tool_call reuses litellm asend_message instead of raw httpx

* fix: _parse_a2a_response handles kind=text (A2A SDK >= 0.3) in addition to type=text

* docs: add finance+compliance demo table and client code snippet

* fix: address greptile P1/P2 review comments

- registry_orchestrator: exact URL match for agent routing (no substring),
  skip agents with empty URLs, deduplicate colliding sanitized names,
  wrap a2a import in ImportError guard, hoist SemanticToolFilterHook import
  out of callback loop, note user_api_key_auth auth filtering as TODO
- litellm_proxy_mcp_handler: guard FastAPI import with try/except ImportError,
  use .get() on tool_server_map to handle hallucinated tool names gracefully,
  add debug logging around A2A tool call dispatch
- chat_completions_handler: lazy-import RegistryOrchestrator inside function
  to avoid hard proxy coupling at module load, extend semantic_filter check to
  cover agent_tool_configs as well as mcp_tools, handle non-streaming ModelResponse
  follow-up by emitting as synthetic final chunk instead of silently dropping

* fix: add LiteLLMLoggingObj pre/post hooks to A2A tool call path

* fix: add async hooks and standard metadata to A2A logging path; fix asyncio.get_event_loop deprecation
This commit is contained in:
Ishaan Jaff 2026-03-21 19:02:02 -07:00 committed by GitHub
parent f5194b5ce3
commit 966bb40845
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1817 additions and 407 deletions

View file

@ -0,0 +1,307 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# MCP + Agent Orchestration in /chat/completions
Use your registered MCP servers and A2A agents directly from `/v1/chat/completions` — no hardcoded URLs required.
## The problem it solves
Without this, every client has to know which MCP server to call:
```json
// ❌ Every request hardcodes a specific server URL
{
"tools": [{"type": "mcp", "server_url": "http://my-zapier-server/mcp", ...}]
}
```
That means updating every client when servers change, no central access control, and no way to let the LLM pick across multiple servers.
## How it works
Register your servers once via `POST /v1/mcp/server`. Then point any chat request at the proxy's registry using `"server_url": "litellm_proxy/mcp"` — the proxy fetches available tools, injects them into the LLM context, and executes tool calls on the model's behalf.
```
POST /v1/chat/completions
├── type:"mcp", server_url:"litellm_proxy/mcp"
│ └── expand → all servers registered via POST /v1/mcp/server
│ └── fetch tool schemas from each server
│ └── inject into LLM context
└── type:"a2a_agent", server_url:"litellm_proxy/agents"
└── expand → all agents registered via POST /v1/agents
└── wrap each agent as a callable function tool
LLM decides which tools to call
Proxy executes tool calls, returns results
Follow-up LLM call → final answer
```
## Quickstart
### 1. Register an MCP server
```bash
curl -X POST http://localhost:4000/v1/mcp/server \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"server_name": "math",
"url": "http://localhost:8001/mcp",
"transport": "http"
}'
```
### 2. Call `/v1/chat/completions` with `litellm_proxy/mcp`
<Tabs>
<TabItem value="curl" label="cURL">
```bash
curl http://localhost:4000/v1/chat/completions \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "What is 1250 + 18?"}],
"tools": [
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never"
}
]
}'
```
</TabItem>
<TabItem value="python" label="Python (OpenAI SDK)">
```python
import openai
client = openai.OpenAI(
api_key="sk-1234",
base_url="http://localhost:4000",
)
response = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "What is 1250 + 18?"}],
tools=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
}
],
)
print(response.choices[0].message.content)
```
</TabItem>
<TabItem value="litellm" label="LiteLLM SDK">
```python
import litellm
response = await litellm.acompletion(
model="gpt-4o-mini",
api_base="http://localhost:4000",
api_key="sk-1234",
messages=[{"role": "user", "content": "What is 1250 + 18?"}],
tools=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
}
],
)
print(response.choices[0].message.content)
```
</TabItem>
</Tabs>
## Target a specific server
Append the server name to `litellm_proxy/mcp/` to restrict tool injection to one server:
```json
{
"type": "mcp",
"server_url": "litellm_proxy/mcp/math",
"require_approval": "never"
}
```
## A2A Agent orchestration
Agents registered via `POST /v1/agents` are exposed as callable function tools using the same pattern.
### 1. Register an agent
```bash
curl -X POST http://localhost:4000/v1/agents \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"agent_name": "FX_Converter",
"agent_url": "http://my-fx-agent/a2a",
"description": "Converts currency amounts using live exchange rates."
}'
```
### 2. Use MCP tools and agents together
```json
{
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "Add 500 + 250, then convert the result to EUR."}],
"tools": [
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never"
},
{
"type": "a2a_agent",
"server_url": "litellm_proxy/agents",
"require_approval": "never"
}
]
}
```
The proxy wraps each registered agent as a function tool. When the LLM calls it, the proxy sends a JSON-RPC `message/send` to the agent and returns the result as a tool message.
## Real-world example — Finance MCP + Compliance Agent
Register a finance calculation MCP server and a compliance analyst A2A agent once. Every request can then use both without knowing any server URLs.
```python
import openai
client = openai.OpenAI(
api_key="sk-1234",
base_url="http://localhost:4000",
)
# MCP only — financial calculation
response = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{
"role": "user",
"content": "What is the monthly repayment on a £250,000 mortgage at 4.5% APR over 25 years?"
}],
tools=[{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
}],
)
# → calls calculate_loan_payment tool → £1,389.58/mo
# Both MCP + Agent in a single call
response = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{
"role": "user",
"content": (
"Calculate compound interest on £100,000 at 2.8% over 3 years, "
"then draft a compliance note summarising the outcome for the audit file."
)
}],
tools=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
},
{
"type": "a2a_agent",
"server_url": "litellm_proxy/agents",
"require_approval": "never",
},
],
)
# → calls calculate_compound_interest (£8,637.40 interest) AND compliance_analyst agent
# → final answer includes both the numbers and the audit-ready compliance note
```
### Demo results (10 scenarios, local proxy, gpt-4o-mini)
MCP server registered: `finance``calculate_compound_interest`, `convert_currency`, `calculate_loan_payment`, `calculate_var`
Agent registered: `compliance_analyst` — Basel III, KYC, VaR, earnings, trade summaries
| # | Scenario | MCP | Agent | Tool Called | Result |
|---|----------|:---:|:---:|-------------|--------|
| 1 | Mortgage repayment | ✓ | — | `calculate_loan_payment(£250k, 4.5%, 25yr)` | **£1,389.58/mo** |
| 2 | FX conversion GBP→USD | ✓ | — | `convert_currency(£1.25M, 1.2738)` | £1,592,250 USD |
| 3 | Compound interest | ✓ | — | `calculate_compound_interest(£50k, 3.5%, 5yr)` | **£9,384 interest** |
| 4 | Basel III notice | — | ✓ | `compliance_analyst` | CET1 ≥4.5%, Tier1 ≥6% — review capital position |
| 5 | KYC note | — | ✓ | `compliance_analyst` | Entity verified, no sanctions, onboarding approved |
| 6 | VaR calculation | ✓ | ✓ | `calculate_var(£5M, 0.8% vol, 99%)` | 1-day VaR **£93,040**, 10-day **£294,218** |
| 7 | Interest calc + audit note | ✓ | ✓ | `calculate_compound_interest` + `compliance_analyst` | **£8,637 interest** + audit-ready compliance note |
| 8 | Mortgage refinance | ✓ | ✓ | `calculate_loan_payment(£180k, 3.9%, 20yr)` | **£1,081.30/mo** |
| 9 | Large FX GBP→JPY | ✓ | — | `convert_currency(£2.5M, 191.45)` | **¥478,625,000** |
| 10 | Earnings summary | — | ✓ | `compliance_analyst` | NII +8% YoY, CET1=13.8%, guidance reaffirmed |
Row 7 demonstrates the orchestrator routing a single request to **both** the MCP finance server and the compliance analyst agent — the LLM received the calculation result from MCP and the formatted audit note from the agent in one turn, with no URL configuration in the client.
## Semantic filter
Add `"semantic_filter": true` to only inject tools relevant to the user's query. Useful when you have many registered servers and want to keep the LLM context lean.
```json
{
"type": "mcp",
"server_url": "litellm_proxy/mcp",
"require_approval": "never",
"semantic_filter": true
}
```
Configure top-k and similarity threshold in your proxy config:
```yaml
mcp_semantic_tool_filter:
top_k: 10
similarity_threshold: 0.3
embedding_model: "text-embedding-3-small"
```
See [MCP Semantic Filter](./mcp_semantic_filter) for setup details.
## Streaming
Works with `"stream": true` — tokens arrive as they're generated, tool execution happens between LLM turns.
```python
stream = await litellm.acompletion(
model="gpt-4o-mini",
api_base="http://localhost:4000",
api_key="sk-1234",
messages=[{"role": "user", "content": "What is 42 × 13?"}],
tools=[{"type": "mcp", "server_url": "litellm_proxy/mcp", "require_approval": "never"}],
stream=True,
)
async for chunk in stream:
delta = chunk.choices[0].delta
if delta.content:
print(delta.content, end="", flush=True)
```
## Access control
Tool visibility follows your existing key and team permissions. A virtual key scoped to specific MCP servers will only see those servers when it calls `litellm_proxy/mcp` — no extra config needed.
See [MCP Zero Trust](./mcp_zero_trust) for per-key and per-team tool restrictions.

View file

@ -308,6 +308,7 @@ const sidebars = {
"a2a_cost_tracking",
"a2a_agent_permissions",
"a2a_iteration_budgets",
"mcp_chat_completions_orchestration",
],
},
{

View file

@ -0,0 +1,307 @@
"""
RegistryOrchestrator centralises all registry-backed orchestration logic.
Responsibilities
----------------
- Parsing ``a2a_agent`` tool configs out of a request's ``tools`` list.
- Resolving registered A2A agents from the global agent registry and wrapping each
one as an OpenAI function tool that the LLM can call.
- Executing a single A2A tool call via JSON-RPC 2.0 ``message/send``.
- Applying the semantic MCP tool filter when the caller opts in via
``"semantic_filter": true`` on an MCP tool config.
"""
import re
from typing import Any, Dict, Iterable, List, Optional, Tuple
from litellm._logging import verbose_logger
# NOTE: Kept broad to avoid coupling to optional OpenAI SDK typing symbols.
ToolParam = Any
LITELLM_PROXY_AGENTS_URL = "litellm_proxy/agents"
# Import hoisted out of the callback loop to avoid re-evaluating on every iteration.
try:
from litellm.proxy.hooks.mcp_semantic_filter.hook import ( # noqa: E501
SemanticToolFilterHook as _SemanticToolFilterHook,
)
except ImportError:
_SemanticToolFilterHook = None # type: ignore
# ---------------------------------------------------------------------------
# Module-level helper
# ---------------------------------------------------------------------------
def _parse_a2a_response(data: Dict[str, Any]) -> str:
"""Extract text content from an A2A JSON-RPC message/send response.
Handles both ``"type": "text"`` (older A2A SDK) and ``"kind": "text"``
(A2A SDK >= 0.3) part schemas.
"""
if "error" in data:
err = data["error"]
return f"Agent error: {err.get('message', str(err))}"
result = data.get("result", {})
def _is_text_part(p: Dict[str, Any]) -> bool:
return (p.get("kind") == "text" or p.get("type") == "text") and bool(
p.get("text")
)
# A2A spec: result.artifacts[].parts[].text
for artifact in result.get("artifacts", []):
texts = [p["text"] for p in artifact.get("parts", []) if _is_text_part(p)]
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 _is_text_part(p):
return p["text"]
return str(result) if result else "Agent executed successfully"
# ---------------------------------------------------------------------------
# RegistryOrchestrator
# ---------------------------------------------------------------------------
class RegistryOrchestrator:
"""
Static-method class that owns all registry-backed orchestration concerns:
* Parsing A2A agent tool configs from a request.
* Resolving registered agents and wrapping them as function tools.
* Executing A2A tool calls via JSON-RPC.
* Applying the per-request semantic MCP tool filter.
"""
@staticmethod
def parse_agent_tool_configs(
tools: Optional[Iterable[ToolParam]],
) -> Tuple[List[ToolParam], List[Any]]:
"""
Separate ``a2a_agent`` registry tool configs from all other tools.
Returns:
(agent_tool_configs, other_tools)
"""
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 server_url == LITELLM_PROXY_AGENTS_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 resolve_agent_tools(
user_api_key_auth: Any,
) -> Tuple[List[Dict[str, Any]], Dict[str, Dict[str, str]]]:
"""
Read all registered A2A agents and expose each as an OpenAI function tool.
Returns:
(function_tools, agent_tool_map)
* ``function_tools``: list of ``{"type": "function", "function": {...}}`` dicts
* ``agent_tool_map``: mapping of sanitized function name ``{"url": str, "agent_name": str}``
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
# NOTE: user_api_key_auth is accepted for future per-key agent filtering
# (mirroring get_allowed_mcp_servers). Agent-level access control is not yet
# implemented in AgentRegistry; all registered agents are returned for now.
_ = user_api_key_auth
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
if not agent_url:
verbose_logger.warning(
"Agent '%s' has no URL configured, skipping", agent_name
)
continue
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]}"
)
# Deduplicate: if two agents produce the same sanitized name, append the
# agent_id suffix so neither is silently dropped.
if func_name in agent_tool_map:
func_name = f"{func_name}_{agent.agent_id[:8]}"[:64]
verbose_logger.warning(
"Agent name collision: renamed to '%s' to avoid overwrite",
func_name,
)
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 LiteLLM's asend_message and return the result."""
import uuid
try:
from a2a.types import (
Message,
MessageSendParams,
Part,
Role,
SendMessageRequest,
TextPart,
)
except ImportError as exc:
raise ImportError(
"The 'a2a' package is required for A2A agent calls. "
"Install it with: pip install a2a-sdk"
) from exc
from litellm.a2a_protocol.main import asend_message
a2a_message = Message(
role=Role.user,
parts=[Part(root=TextPart(text=message))],
message_id=uuid.uuid4().hex,
context_id=litellm_trace_id,
)
request = SendMessageRequest(
id=str(uuid.uuid4()),
params=MessageSendParams(message=a2a_message),
)
try:
response = await asend_message(
api_base=agent_url,
request=request,
agent_id=agent_name,
)
result_text = _parse_a2a_response(
response.model_dump(mode="json", exclude_none=True)
)
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 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 []:
if _SemanticToolFilterHook is None:
break
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

View file

@ -2,17 +2,19 @@
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,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.utils import ModelResponse
from litellm.types.utils import ModelResponse, ModelResponseStream, StreamingChoices
from litellm.utils import CustomStreamWrapper
@ -77,6 +79,356 @@ def _add_mcp_metadata_to_response(
setattr(message, "provider_specific_fields", provider_fields)
class _SyncIteratorWrapper:
"""Wraps an async iterator for synchronous iteration."""
def __init__(self, async_iterator: Any, loop: Any) -> None:
self._async_iterator = async_iterator
self._loop = loop
self._iterator: Any = None
def __iter__(self) -> "_SyncIteratorWrapper":
return self
def __next__(self) -> Any:
import asyncio
if self._iterator is None:
aiter_result = self._async_iterator.__aiter__()
if hasattr(aiter_result, "__await__"):
self._iterator = self._loop.run_until_complete(aiter_result)
else:
self._iterator = aiter_result
try:
return self._loop.run_until_complete(self._iterator.__anext__())
except StopAsyncIteration:
raise StopIteration
class MCPStreamingIterator:
"""
Async iterator that drives the MCP tool-execution loop for streaming responses.
Phases:
1. Yield chunks from the initial LLM stream.
2. When the stream ends, execute any tool calls (MCP or A2A).
3. Yield chunks from the follow-up LLM stream.
"""
def __init__(
self,
stream_wrapper: Any,
messages: List,
tool_server_map: Any,
user_api_key_auth: Any,
mcp_auth_header: Optional[str],
mcp_server_auth_headers: Any,
oauth2_headers: Any,
raw_headers: Any,
litellm_call_id: Optional[str],
litellm_trace_id: Optional[str],
openai_tools: List,
base_call_args: Dict[str, Any],
agent_tool_map: Optional[Dict[str, Any]] = None,
) -> None:
self.stream_wrapper = stream_wrapper
self.messages = messages
self.tool_server_map = tool_server_map
self.user_api_key_auth = user_api_key_auth
self.mcp_auth_header = mcp_auth_header
self.mcp_server_auth_headers = mcp_server_auth_headers
self.oauth2_headers = oauth2_headers
self.raw_headers = raw_headers
self.litellm_call_id = litellm_call_id
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
self.complete_response: Optional[ModelResponse] = None
self.stream_exhausted = False
self.tool_execution_done = False
self.follow_up_stream: Optional[CustomStreamWrapper] = None
self.follow_up_non_stream: Optional[ModelResponse] = None
self.follow_up_iterator: Any = None
self.follow_up_exhausted = False
def __aiter__(self) -> "MCPStreamingIterator":
return self
def _add_mcp_list_tools_to_chunk(
self, chunk: ModelResponseStream
) -> ModelResponseStream:
"""Add mcp_list_tools to the first chunk."""
from litellm.types.utils import add_provider_specific_fields
if not self.openai_tools:
return chunk
if hasattr(chunk, "choices") and chunk.choices:
for choice in chunk.choices:
if (
isinstance(choice, StreamingChoices)
and hasattr(choice, "delta")
and choice.delta
):
provider_fields = dict(
getattr(choice.delta, "provider_specific_fields", None) or {}
)
provider_fields["mcp_list_tools"] = self.openai_tools
add_provider_specific_fields(choice.delta, provider_fields)
return chunk
def _add_mcp_tool_metadata_to_final_chunk(
self, chunk: ModelResponseStream
) -> ModelResponseStream:
"""Add mcp_tool_calls and mcp_call_results to the final chunk."""
from litellm.types.utils import add_provider_specific_fields
if hasattr(chunk, "choices") and chunk.choices:
for choice in chunk.choices:
if (
isinstance(choice, StreamingChoices)
and hasattr(choice, "delta")
and choice.delta
):
attr_value = getattr(choice.delta, "provider_specific_fields", None)
provider_fields = (
dict(attr_value) if isinstance(attr_value, dict) else {}
)
if self.tool_calls:
provider_fields["mcp_tool_calls"] = self.tool_calls
if self.tool_results:
provider_fields["mcp_call_results"] = self.tool_results
add_provider_specific_fields(choice.delta, provider_fields)
return chunk
async def __anext__(self) -> Any:
# Phase 1: Collect and yield initial stream chunks
if not self.stream_exhausted:
if not hasattr(self, "_stream_iterator"):
self._stream_iterator = self.stream_wrapper.__aiter__()
_add_mcp_metadata_to_response(
response=self.stream_wrapper,
openai_tools=self.openai_tools,
)
try:
chunk = await self._stream_iterator.__anext__()
self.collected_chunks.append(chunk)
if len(self.collected_chunks) == 1:
chunk = self._add_mcp_list_tools_to_chunk(chunk)
is_final = (
hasattr(chunk, "choices")
and chunk.choices
and hasattr(chunk.choices[0], "finish_reason")
and chunk.choices[0].finish_reason is not None
)
if is_final:
self.stream_exhausted = True
await self._process_tool_calls()
chunk = self._add_mcp_tool_metadata_to_final_chunk(chunk)
if self.tool_results and self.complete_response:
await self._prepare_follow_up_call()
return chunk
except StopAsyncIteration:
self.stream_exhausted = True
await self._process_tool_calls()
if self.collected_chunks:
final_chunk = self.collected_chunks[-1]
final_chunk = self._add_mcp_tool_metadata_to_final_chunk(
final_chunk
)
if self.tool_results and self.complete_response:
await self._prepare_follow_up_call()
return final_chunk
# Phase 2: Yield follow-up stream chunks if available
if self.follow_up_stream and not self.follow_up_exhausted:
if not self.follow_up_iterator:
self.follow_up_iterator = self.follow_up_stream.__aiter__()
verbose_logger.debug("Follow-up stream iterator created")
try:
chunk = await self.follow_up_iterator.__anext__()
verbose_logger.debug("Follow-up chunk yielded: %s", chunk)
return chunk
except StopAsyncIteration:
self.follow_up_exhausted = True
verbose_logger.debug("Follow-up stream exhausted")
raise StopAsyncIteration
# Phase 3: emit non-streaming follow-up answer as a synthetic final chunk
if self.follow_up_non_stream is not None:
from litellm.types.utils import ModelResponseStream, StreamingChoices
non_stream = self.follow_up_non_stream
self.follow_up_non_stream = None
# Build a minimal streaming chunk from the ModelResponse
content = ""
if non_stream.choices:
content = getattr(non_stream.choices[0].message, "content", "") or ""
synthetic = ModelResponseStream(
id=non_stream.id,
model=non_stream.model or "",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta={"role": "assistant", "content": content}, # type: ignore[arg-type]
)
],
)
return synthetic
raise StopAsyncIteration
async def _process_tool_calls(self) -> None:
"""Build complete response from collected chunks and execute any tool calls."""
from litellm.main import stream_chunk_builder
if self.tool_execution_done:
return
self.tool_execution_done = True
if not self.collected_chunks:
return
complete_response = stream_chunk_builder(
chunks=self.collected_chunks,
messages=self.messages,
)
if isinstance(complete_response, ModelResponse):
self.complete_response = complete_response
self.tool_calls = (
LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(
response=complete_response
)
)
if self.tool_calls:
self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=self.tool_server_map,
tool_calls=self.tool_calls,
user_api_key_auth=self.user_api_key_auth,
mcp_auth_header=self.mcp_auth_header,
mcp_server_auth_headers=self.mcp_server_auth_headers,
oauth2_headers=self.oauth2_headers,
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,
)
async def _prepare_follow_up_call(self) -> None:
"""Initiate the follow-up streaming call with tool results."""
if self.follow_up_stream is not None:
return
if not self.tool_results or not self.complete_response:
return
follow_up_messages = (
LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat(
original_messages=self.messages,
response=self.complete_response,
tool_results=self.tool_results,
)
)
follow_up_call_args = {
**self.base_call_args,
"messages": follow_up_messages,
"stream": True,
"_skip_mcp_handler": True,
}
import litellm
follow_up_response = await litellm.acompletion(**follow_up_call_args)
if isinstance(follow_up_response, CustomStreamWrapper):
self.follow_up_stream = follow_up_response
verbose_logger.debug("Follow-up stream created successfully")
elif isinstance(follow_up_response, ModelResponse):
# Provider returned a non-streaming response despite stream=True.
# Store it so __anext__ can yield it as a synthetic final chunk rather
# than silently dropping the follow-up answer.
self.follow_up_non_stream = follow_up_response
verbose_logger.debug(
"Follow-up response is non-streaming ModelResponse; will emit as final chunk"
)
else:
verbose_logger.warning(
"Follow-up response is unexpected type %s, answer may be dropped",
type(follow_up_response),
)
class MCPStreamWrapper(CustomStreamWrapper):
"""
Thin ``CustomStreamWrapper`` subclass that delegates async iteration to
an ``MCPStreamingIterator`` so that the MCP tool-execution loop is
transparent to callers that consume the stream normally.
"""
def __init__(
self,
original_wrapper: CustomStreamWrapper,
custom_iterator: MCPStreamingIterator,
) -> None:
super().__init__(
completion_stream=None,
model=getattr(original_wrapper, "model", "unknown"),
logging_obj=getattr(original_wrapper, "logging_obj", None),
custom_llm_provider=getattr(original_wrapper, "custom_llm_provider", None),
stream_options=getattr(original_wrapper, "stream_options", None),
make_call=getattr(original_wrapper, "make_call", None),
_response_headers=getattr(original_wrapper, "_response_headers", None),
)
self._original_wrapper = original_wrapper
self._custom_iterator = custom_iterator
if hasattr(original_wrapper, "_hidden_params"):
self._hidden_params = original_wrapper._hidden_params
self._sync_iterator: Optional[_SyncIteratorWrapper] = None
self._sync_loop: Any = None
def __aiter__(self) -> MCPStreamingIterator:
return self._custom_iterator
def __iter__(self) -> _SyncIteratorWrapper:
import asyncio
if self._sync_iterator is None:
self._sync_loop = asyncio.new_event_loop()
asyncio.set_event_loop(self._sync_loop)
self._sync_iterator = _SyncIteratorWrapper(
self._custom_iterator, self._sync_loop
)
return self._sync_iterator
def __next__(self) -> Any:
if self._sync_iterator is None:
self.__iter__()
assert self._sync_iterator is not None
return next(self._sync_iterator)
def __getattr__(self, name: str) -> Any:
return getattr(self._original_wrapper, name)
async def acompletion_with_mcp( # noqa: PLR0915
model: str,
messages: List,
@ -97,6 +449,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
5. Make a follow-up call with the tool results
"""
from litellm import acompletion as litellm_acompletion
from litellm.proxy.agent_endpoints.registry_orchestrator import RegistryOrchestrator
# Parse MCP tools and separate from other tools
(
@ -104,8 +457,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,
) = RegistryOrchestrator.parse_agent_tool_configs(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 +500,41 @@ async def acompletion_with_mcp( # noqa: PLR0915
mcp_server_auth_headers=mcp_server_auth_headers,
)
# Apply per-tool semantic filter if any MCP or agent tool config has semantic_filter=true
if any(
isinstance(t, dict) and t.get("semantic_filter")
for t in list(mcp_tools_with_litellm_proxy) + list(agent_tool_configs)
):
deduplicated_mcp_tools = await RegistryOrchestrator.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 RegistryOrchestrator.resolve_agent_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 +569,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)
@ -199,312 +583,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
)
return initial_stream
# Create a custom async generator that collects chunks and handles tool execution
from litellm.main import stream_chunk_builder
from litellm.types.utils import ModelResponseStream
class MCPStreamingIterator:
"""Custom iterator that collects chunks, detects tool calls, and adds MCP metadata to final chunk."""
def __init__(
self,
stream_wrapper,
messages,
tool_server_map,
user_api_key_auth,
mcp_auth_header,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
litellm_call_id,
litellm_trace_id,
openai_tools,
base_call_args,
):
self.stream_wrapper = stream_wrapper
self.messages = messages
self.tool_server_map = tool_server_map
self.user_api_key_auth = user_api_key_auth
self.mcp_auth_header = mcp_auth_header
self.mcp_server_auth_headers = mcp_server_auth_headers
self.oauth2_headers = oauth2_headers
self.raw_headers = raw_headers
self.litellm_call_id = litellm_call_id
self.litellm_trace_id = litellm_trace_id
self.openai_tools = openai_tools
self.base_call_args = base_call_args
self.collected_chunks: List[ModelResponseStream] = []
self.tool_calls: Optional[List] = None
self.tool_results: Optional[List] = None
self.complete_response: Optional[ModelResponse] = None
self.stream_exhausted = False
self.tool_execution_done = False
self.follow_up_stream = None
self.follow_up_iterator = None
self.follow_up_exhausted = False
async def __aiter__(self):
return self
def _add_mcp_list_tools_to_chunk(
self, chunk: ModelResponseStream
) -> ModelResponseStream:
"""Add mcp_list_tools to the first chunk."""
from litellm.types.utils import (
StreamingChoices,
add_provider_specific_fields,
)
if not self.openai_tools:
return chunk
if hasattr(chunk, "choices") and chunk.choices:
for choice in chunk.choices:
if (
isinstance(choice, StreamingChoices)
and hasattr(choice, "delta")
and choice.delta
):
# Get existing provider_specific_fields or create new dict
existing_fields = (
getattr(choice.delta, "provider_specific_fields", None)
or {}
)
provider_fields = dict(
existing_fields
) # Create a copy to avoid mutating the original
# Add only mcp_list_tools to first chunk
provider_fields["mcp_list_tools"] = self.openai_tools
# Use add_provider_specific_fields to ensure proper setting
# This function handles Pydantic model attribute setting correctly
add_provider_specific_fields(choice.delta, provider_fields)
return chunk
def _add_mcp_tool_metadata_to_final_chunk(
self, chunk: ModelResponseStream
) -> ModelResponseStream:
"""Add mcp_tool_calls and mcp_call_results to the final chunk."""
from litellm.types.utils import (
StreamingChoices,
add_provider_specific_fields,
)
if hasattr(chunk, "choices") and chunk.choices:
for choice in chunk.choices:
if (
isinstance(choice, StreamingChoices)
and hasattr(choice, "delta")
and choice.delta
):
# Get existing provider_specific_fields or create new dict
# Access the attribute directly to handle Pydantic model attributes correctly
existing_fields = {}
if hasattr(choice.delta, "provider_specific_fields"):
attr_value = getattr(
choice.delta, "provider_specific_fields", None
)
if attr_value is not None:
# Create a copy to avoid mutating the original
existing_fields = (
dict(attr_value)
if isinstance(attr_value, dict)
else {}
)
provider_fields = existing_fields
# Add tool_calls and tool_results if available
if self.tool_calls:
provider_fields["mcp_tool_calls"] = self.tool_calls
if self.tool_results:
provider_fields["mcp_call_results"] = self.tool_results
# Use add_provider_specific_fields to ensure proper setting
# This function handles Pydantic model attribute setting correctly
add_provider_specific_fields(choice.delta, provider_fields)
return chunk
async def __anext__(self):
# Phase 1: Collect and yield initial stream chunks
if not self.stream_exhausted:
# Get the iterator from the stream wrapper
if not hasattr(self, "_stream_iterator"):
self._stream_iterator = self.stream_wrapper.__aiter__()
# Add mcp_list_tools to the first chunk (available from the start)
_add_mcp_metadata_to_response(
response=self.stream_wrapper,
openai_tools=self.openai_tools,
)
try:
chunk = await self._stream_iterator.__anext__()
self.collected_chunks.append(chunk)
# Add mcp_list_tools to the first chunk
if len(self.collected_chunks) == 1:
chunk = self._add_mcp_list_tools_to_chunk(chunk)
# Check if this is the final chunk (has finish_reason)
is_final = (
hasattr(chunk, "choices")
and chunk.choices
and hasattr(chunk.choices[0], "finish_reason")
and chunk.choices[0].finish_reason is not None
)
if is_final:
# This is the final chunk, mark stream as exhausted
self.stream_exhausted = True
# Process tool calls after we've collected all chunks
await self._process_tool_calls()
# Apply MCP metadata (tool_calls and tool_results) to final chunk
chunk = self._add_mcp_tool_metadata_to_final_chunk(chunk)
# If we have tool results, prepare follow-up call immediately
if self.tool_results and self.complete_response:
await self._prepare_follow_up_call()
return chunk
except StopAsyncIteration:
self.stream_exhausted = True
# Process tool calls after stream is exhausted
await self._process_tool_calls()
# If we have chunks, yield the final one with metadata
if self.collected_chunks:
final_chunk = self.collected_chunks[-1]
final_chunk = self._add_mcp_tool_metadata_to_final_chunk(
final_chunk
)
# If we have tool results, prepare follow-up call
if self.tool_results and self.complete_response:
await self._prepare_follow_up_call()
return final_chunk
# Phase 2: Yield follow-up stream chunks if available
if self.follow_up_stream and not self.follow_up_exhausted:
if not self.follow_up_iterator:
self.follow_up_iterator = self.follow_up_stream.__aiter__()
from litellm._logging import verbose_logger
verbose_logger.debug("Follow-up stream iterator created")
try:
chunk = await self.follow_up_iterator.__anext__()
from litellm._logging import verbose_logger
verbose_logger.debug(f"Follow-up chunk yielded: {chunk}")
return chunk
except StopAsyncIteration:
self.follow_up_exhausted = True
from litellm._logging import verbose_logger
verbose_logger.debug("Follow-up stream exhausted")
# After follow-up stream is exhausted, check if we need to raise StopAsyncIteration
raise StopAsyncIteration
# If we're here and follow_up_stream is None but we expected it, log a warning
if (
self.stream_exhausted
and self.tool_results
and self.complete_response
and self.follow_up_stream is None
):
from litellm._logging import verbose_logger
verbose_logger.warning(
"Follow-up stream was not created despite having tool results"
)
raise StopAsyncIteration
async def _process_tool_calls(self):
"""Process tool calls after streaming completes."""
if self.tool_execution_done:
return
self.tool_execution_done = True
if not self.collected_chunks:
return
# Build complete response from chunks
complete_response = stream_chunk_builder(
chunks=self.collected_chunks,
messages=self.messages,
)
if isinstance(complete_response, ModelResponse):
self.complete_response = complete_response
# Extract tool calls from complete response
self.tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(
response=complete_response
)
if self.tool_calls:
# Execute tool calls
self.tool_results = (
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=self.tool_server_map,
tool_calls=self.tool_calls,
user_api_key_auth=self.user_api_key_auth,
mcp_auth_header=self.mcp_auth_header,
mcp_server_auth_headers=self.mcp_server_auth_headers,
oauth2_headers=self.oauth2_headers,
raw_headers=self.raw_headers,
litellm_call_id=self.litellm_call_id,
litellm_trace_id=self.litellm_trace_id,
)
)
async def _prepare_follow_up_call(self):
"""Prepare and initiate follow-up call with tool results."""
if self.follow_up_stream is not None:
return # Already prepared
if not self.tool_results or not self.complete_response:
return
# Create follow-up messages with tool results
follow_up_messages = (
LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat(
original_messages=self.messages,
response=self.complete_response,
tool_results=self.tool_results,
)
)
# Make follow-up call with streaming
follow_up_call_args = dict(self.base_call_args)
follow_up_call_args["messages"] = follow_up_messages
follow_up_call_args["stream"] = True
# Ensure follow-up call doesn't trigger MCP handler again
follow_up_call_args["_skip_mcp_handler"] = True
# Import litellm here to ensure we get the patched version
# This ensures the patch works correctly in tests
import litellm
follow_up_response = await litellm.acompletion(**follow_up_call_args)
# Ensure follow-up response is a CustomStreamWrapper
if isinstance(follow_up_response, CustomStreamWrapper):
self.follow_up_stream = follow_up_response
from litellm._logging import verbose_logger
verbose_logger.debug("Follow-up stream created successfully")
else:
# Unexpected response type - log and set to None
from litellm._logging import verbose_logger
verbose_logger.warning(
f"Follow-up response is not a CustomStreamWrapper: {type(follow_up_response)}"
)
self.follow_up_stream = None
# Create the custom iterator
# Create the MCP streaming iterator (module-level class)
iterator = MCPStreamingIterator(
stream_wrapper=initial_stream,
messages=messages,
@ -518,88 +597,9 @@ 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
# We'll use a simple approach: just replace the __aiter__ method
class MCPStreamWrapper(CustomStreamWrapper):
def __init__(self, original_wrapper, custom_iterator):
# Initialize with the same parameters as original wrapper
super().__init__(
completion_stream=None,
model=getattr(original_wrapper, "model", "unknown"),
logging_obj=getattr(original_wrapper, "logging_obj", None),
custom_llm_provider=getattr(
original_wrapper, "custom_llm_provider", None
),
stream_options=getattr(original_wrapper, "stream_options", None),
make_call=getattr(original_wrapper, "make_call", None),
_response_headers=getattr(
original_wrapper, "_response_headers", None
),
)
self._original_wrapper = original_wrapper
self._custom_iterator = custom_iterator
# Copy important attributes from original wrapper
if hasattr(original_wrapper, "_hidden_params"):
self._hidden_params = original_wrapper._hidden_params
# For synchronous iteration, we need to run the async iterator
self._sync_iterator = None
self._sync_loop = None
def __aiter__(self):
return self._custom_iterator
def __iter__(self):
# For synchronous iteration, create a sync wrapper
if self._sync_iterator is None:
import asyncio
try:
self._sync_loop = asyncio.get_event_loop()
except RuntimeError:
self._sync_loop = asyncio.new_event_loop()
asyncio.set_event_loop(self._sync_loop)
self._sync_iterator = _SyncIteratorWrapper(
self._custom_iterator, self._sync_loop
)
return self._sync_iterator
def __next__(self):
# Delegate to sync iterator
if self._sync_iterator is None:
self.__iter__()
return next(self._sync_iterator)
def __getattr__(self, name):
# Delegate all other attributes to original wrapper
return getattr(self._original_wrapper, name)
# Helper class to wrap async iterator for sync iteration
class _SyncIteratorWrapper:
def __init__(self, async_iterator, loop):
self._async_iterator = async_iterator
self._loop = loop
self._iterator = None
def __iter__(self):
return self
def __next__(self):
if self._iterator is None:
# __aiter__ might be async, so we need to await it
aiter_result = self._async_iterator.__aiter__()
if hasattr(aiter_result, "__await__"):
# It's a coroutine, await it
self._iterator = self._loop.run_until_complete(aiter_result)
else:
# It's already an iterator
self._iterator = aiter_result
try:
return self._loop.run_until_complete(self._iterator.__anext__())
except StopAsyncIteration:
raise StopIteration
return cast(CustomStreamWrapper, MCPStreamWrapper(initial_stream, iterator))
# Non-streaming mode: use existing logic
@ -626,7 +626,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 +637,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

@ -148,7 +148,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 +159,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,20 +545,30 @@ 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
try:
from fastapi import HTTPException
except ImportError:
# FastAPI is a proxy-only dependency; fall back to a plain exception so
# SDK users (without fastapi installed) can still call MCP tools.
# The .detail access below is already guarded with hasattr().
HTTPException = Exception # type: ignore[assignment,misc]
from litellm._uuid import uuid
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.agent_endpoints.registry_orchestrator import (
RegistryOrchestrator,
)
from litellm.proxy.proxy_server import proxy_logging_obj
rules_obj = Rules()
tool_results = []
tool_call_id: Optional[str] = None
rules_obj = Rules()
for tool_call in tool_calls:
logging_request_data: Dict[str, Any] = {}
tool_name: Optional[str] = None
@ -569,10 +587,111 @@ class LiteLLM_Proxy_MCP_Handler:
tool_arguments
)
# Import here to avoid circular import
from litellm.proxy.proxy_server import proxy_logging_obj
# Route A2A agent tool calls through the same logging path as MCP tools
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)
start_time = datetime.now()
logging_request_data = {
"model": f"A2A: {tool_name}",
"metadata": {
"tool_call_id": tool_call_id,
"tool_name": tool_name,
"agent_name": agent_info["agent_name"],
},
"input": [{"role": "tool", "content": message}],
"call_type": CallTypes.call_mcp_tool.value,
"litellm_call_id": litellm_call_id or str(uuid.uuid4()),
"proxy_server_request": {
"url": agent_info["url"],
"method": "POST",
"headers": {},
"body": {"message": message},
},
}
if litellm_trace_id:
logging_request_data["litellm_trace_id"] = litellm_trace_id
if user_api_key_auth is not None:
user_api_key = getattr(user_api_key_auth, "api_key", None)
if user_api_key:
logging_request_data["metadata"][
"user_api_key"
] = user_api_key
server_name = tool_server_map[tool_name]
litellm_logging_obj = None
try:
litellm_logging_obj, _ = function_setup(
original_function="call_mcp_tool",
rules_obj=rules_obj,
start_time=start_time,
**logging_request_data,
)
except Exception as _log_err:
verbose_logger.debug(
"Failed to init logging for A2A tool call %s: %s",
tool_name,
_log_err,
)
standard_logging_a2a_tool_call: StandardLoggingMCPToolCall = {
"name": tool_name,
"arguments": parsed_arguments,
"namespaced_tool_name": tool_name,
"mcp_server_name": agent_info["agent_name"],
}
if litellm_logging_obj:
litellm_logging_obj.model_call_details[
"mcp_tool_call_metadata"
] = standard_logging_a2a_tool_call
litellm_logging_obj.model = f"A2A: {tool_name}"
litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value
try:
litellm_logging_obj.pre_call(input=[message], api_key="")
except Exception:
pass
result = await RegistryOrchestrator.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,
)
if litellm_logging_obj:
try:
litellm_logging_obj.post_call(
original_response=result.get("result", "")
)
end_time = datetime.now()
await litellm_logging_obj.async_post_mcp_tool_call_hook(
kwargs=litellm_logging_obj.model_call_details,
response_obj=result,
start_time=start_time,
end_time=end_time,
)
await litellm_logging_obj.async_success_handler(
result=result,
start_time=start_time,
end_time=end_time,
)
except Exception:
verbose_logger.exception(
"Failed to log A2A tool call success for %s", tool_name
)
tool_results.append(result)
continue
server_name = tool_server_map.get(tool_name)
if server_name is None:
verbose_logger.warning(
"Tool '%s' not found in tool_server_map — skipping (possible "
"hallucinated tool name)",
tool_name,
)
continue
# Remove the server name prefix if the tool name includes it.
sanitized_tool_name = tool_name
@ -682,14 +801,14 @@ class LiteLLM_Proxy_MCP_Handler:
standard_logging_mcp_tool_call["mcp_server_logo_url"] = logo_url
cost_info = mcp_info.get("mcp_server_cost_info")
if cost_info:
standard_logging_mcp_tool_call[
"mcp_server_cost_info"
] = cost_info
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (
cost_info
)
if litellm_logging_obj:
litellm_logging_obj.model_call_details[
"mcp_tool_call_metadata"
] = standard_logging_mcp_tool_call
litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = (
standard_logging_mcp_tool_call
)
litellm_logging_obj.model = f"MCP: {tool_name}"
litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value
@ -779,7 +898,7 @@ class LiteLLM_Proxy_MCP_Handler:
error=e,
)
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}"
error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}" # type: ignore[union-attr]
tool_results.append(
{
"tool_call_id": tool_call_id,

View file

@ -0,0 +1,675 @@
"""
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:
from litellm.proxy.agent_endpoints.registry_orchestrator import (
RegistryOrchestrator,
)
function_tools, agent_tool_map = await RegistryOrchestrator.resolve_agent_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.proxy.agent_endpoints.registry_orchestrator 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.proxy.agent_endpoints.registry_orchestrator 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.proxy.agent_endpoints.registry_orchestrator 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 RegistryOrchestrator.apply_semantic_filter (moved from module-level)
from litellm.proxy.agent_endpoints.registry_orchestrator import RegistryOrchestrator
monkeypatch.setattr(
RegistryOrchestrator,
"apply_semantic_filter",
staticmethod(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}"