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
This commit is contained in:
Ishaan Jaffer 2026-03-21 15:55:27 -07:00
parent 6cc394efa9
commit e44d23892f
4 changed files with 652 additions and 655 deletions

View file

@ -0,0 +1,265 @@
"""
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"
# ---------------------------------------------------------------------------
# Module-level helper
# ---------------------------------------------------------------------------
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"
# ---------------------------------------------------------------------------
# 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 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 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
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 2.0 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 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

View file

@ -10,48 +10,15 @@ from typing import (
)
from litellm._logging import verbose_logger
from litellm.proxy.agent_endpoints.registry_orchestrator import RegistryOrchestrator
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
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],
@ -113,6 +80,338 @@ 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_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
if (
self.stream_exhausted
and self.tool_results
and self.complete_response
and self.follow_up_stream is None
):
verbose_logger.warning(
"Follow-up stream was not created despite having tool results"
)
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")
else:
verbose_logger.warning(
"Follow-up response is not a CustomStreamWrapper: %s",
type(follow_up_response),
)
self.follow_up_stream = None
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:
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) -> 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,
@ -144,7 +443,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
(
agent_tool_configs,
other_tools,
) = LiteLLM_Proxy_MCP_Handler._parse_agent_tools(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
@ -188,7 +487,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
isinstance(t, dict) and t.get("semantic_filter")
for t in mcp_tools_with_litellm_proxy
):
deduplicated_mcp_tools = await _apply_semantic_filter(
deduplicated_mcp_tools = await RegistryOrchestrator.apply_semantic_filter(
tools=deduplicated_mcp_tools,
messages=messages,
)
@ -203,7 +502,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
agent_tool_map: dict = {}
if agent_tool_configs:
agent_function_tools, agent_tool_map = (
await LiteLLM_Proxy_MCP_Handler._wrap_agents_as_function_tools(
await RegistryOrchestrator.resolve_agent_tools(
user_api_key_auth=user_api_key_auth,
)
)
@ -266,315 +565,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,
agent_tool_map=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 = 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,
agent_tool_map=self.agent_tool_map,
)
)
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,
@ -591,87 +582,6 @@ async def acompletion_with_mcp( # noqa: PLR0915
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__()
assert self._sync_iterator is not None
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

View file

@ -43,36 +43,6 @@ ToolParam = Any
LITELLM_PROXY_MCP_SERVER_URL = "litellm_proxy"
LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/"
LITELLM_PROXY_AGENTS_URL = "litellm_proxy/agents"
def _parse_a2a_response(data: Dict[str, Any]) -> str:
"""Extract text content from an A2A JSON-RPC message/send response."""
if "error" in data:
err = data["error"]
return f"Agent error: {err.get('message', str(err))}"
result = data.get("result", {})
# A2A spec: result.artifacts[].parts[].text
for artifact in result.get("artifacts", []):
texts = [
p["text"]
for p in artifact.get("parts", [])
if p.get("type") == "text" and p.get("text")
]
if texts:
return "\n".join(texts)
# Fallback: status.message.parts[].text
status = result.get("status", {})
if isinstance(status, dict):
msg = status.get("message") or {}
for p in msg.get("parts", []):
if p.get("type") == "text" and p.get("text"):
return p["text"]
return str(result) if result else "Agent executed successfully"
# Matches any URL whose path ends with /mcp/<server_name> — covers both root-path
# (http://host:port/mcp/name) and sub-path (http://host/base/mcp/name) proxy deployments.
@ -149,153 +119,6 @@ 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,
@ -338,7 +161,7 @@ class LiteLLM_Proxy_MCP_Handler:
):
# "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):]
server_name = server_url[len(LITELLM_PROXY_MCP_SERVER_URL_PREFIX) :]
if server_name:
if mcp_servers is None:
mcp_servers = []
@ -732,11 +555,14 @@ class LiteLLM_Proxy_MCP_Handler:
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
@ -759,7 +585,7 @@ class LiteLLM_Proxy_MCP_Handler:
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(
result = await RegistryOrchestrator.execute_a2a_tool_call(
agent_url=agent_info["url"],
agent_name=agent_info["agent_name"],
message=message,
@ -770,9 +596,6 @@ class LiteLLM_Proxy_MCP_Handler:
tool_results.append(result)
continue
# Import here to avoid circular import
from litellm.proxy.proxy_server import proxy_logging_obj
server_name = tool_server_map[tool_name]
# Remove the server name prefix if the tool name includes it.
@ -883,14 +706,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

View file

@ -117,9 +117,7 @@ async def test_registry_orchestration_nonstreaming(monkeypatch):
if name == "add":
executed.append({"type": "mcp", "tool": "add"})
results.append(
{"tool_call_id": call_id, "result": "12", "name": "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(
@ -215,9 +213,9 @@ async def test_registry_orchestration_nonstreaming(monkeypatch):
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())}"
)
assert (
"mcp_list_tools" in mcp_metadata
), f"Expected mcp_list_tools in provider_specific_fields, got: {list(mcp_metadata.keys())}"
# ---------------------------------------------------------------------------
@ -240,9 +238,7 @@ async def test_bare_mcp_url_expands_to_all_servers(monkeypatch):
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"}
]
return [{"tool_call_id": "tc-1", "result": "8", "name": "add"}]
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler,
@ -288,9 +284,9 @@ async def test_bare_mcp_url_expands_to_all_servers(monkeypatch):
# 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"
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' "
@ -316,10 +312,12 @@ async def test_agents_wrapped_as_function_tools(monkeypatch):
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
)
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
@ -331,14 +329,15 @@ async def test_agents_wrapped_as_function_tools(monkeypatch):
# 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"
)
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']}"
)
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"]
@ -358,7 +357,7 @@ async def test_agents_wrapped_as_function_tools(monkeypatch):
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
from litellm.proxy.agent_endpoints.registry_orchestrator import _parse_a2a_response
data = {
"jsonrpc": "2.0",
@ -381,7 +380,7 @@ def test_parse_a2a_response_artifacts():
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
from litellm.proxy.agent_endpoints.registry_orchestrator import _parse_a2a_response
data = {
"jsonrpc": "2.0",
@ -401,7 +400,7 @@ def test_parse_a2a_response_status_message():
def test_parse_a2a_response_error():
"""Error responses surface the error message."""
from litellm.responses.mcp.litellm_proxy_mcp_handler import _parse_a2a_response
from litellm.proxy.agent_endpoints.registry_orchestrator import _parse_a2a_response
data = {
"jsonrpc": "2.0",
@ -442,9 +441,7 @@ async def test_registry_orchestration_streaming(monkeypatch):
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"}
)
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(
@ -560,9 +557,9 @@ async def test_registry_orchestration_streaming(monkeypatch):
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}"
)
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"
@ -631,11 +628,13 @@ async def test_semantic_filter_reduces_tools(monkeypatch):
"extract_mcp_headers_from_request",
staticmethod(_no_mcp_headers),
)
# Patch the module-level _apply_semantic_filter used inside acompletion_with_mcp
# Patch RegistryOrchestrator.apply_semantic_filter (moved from module-level)
from litellm.proxy.agent_endpoints.registry_orchestrator import RegistryOrchestrator
monkeypatch.setattr(
chat_completions_handler,
"_apply_semantic_filter",
fake_semantic_filter,
RegistryOrchestrator,
"apply_semantic_filter",
staticmethod(fake_semantic_filter),
)
response = await litellm.acompletion(
@ -671,6 +670,6 @@ async def test_semantic_filter_reduces_tools(monkeypatch):
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}"
)
assert (
"multiply" not in tool_names
), f"Expected 'multiply' to be filtered out, got: {tool_names}"