mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
6cc394efa9
commit
e44d23892f
4 changed files with 652 additions and 655 deletions
265
litellm/proxy/agent_endpoints/registry_orchestrator.py
Normal file
265
litellm/proxy/agent_endpoints/registry_orchestrator.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue