diff --git a/litellm/proxy/agent_endpoints/registry_orchestrator.py b/litellm/proxy/agent_endpoints/registry_orchestrator.py new file mode 100644 index 00000000000..21824e5d4bf --- /dev/null +++ b/litellm/proxy/agent_endpoints/registry_orchestrator.py @@ -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 diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index a8c236cc03c..07f2f051dac 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -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 diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index ebaf14547e3..ac6d31ab109 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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/ — 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 diff --git a/tests/test_litellm/responses/mcp/test_registry_orchestration.py b/tests/test_litellm/responses/mcp/test_registry_orchestration.py index f8ad913ab62..54a504c43b7 100644 --- a/tests/test_litellm/responses/mcp/test_registry_orchestration.py +++ b/tests/test_litellm/responses/mcp/test_registry_orchestration.py @@ -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}"