mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: address greptile P1/P2 review comments
- registry_orchestrator: exact URL match for agent routing (no substring), skip agents with empty URLs, deduplicate colliding sanitized names, wrap a2a import in ImportError guard, hoist SemanticToolFilterHook import out of callback loop, note user_api_key_auth auth filtering as TODO - litellm_proxy_mcp_handler: guard FastAPI import with try/except ImportError, use .get() on tool_server_map to handle hallucinated tool names gracefully, add debug logging around A2A tool call dispatch - chat_completions_handler: lazy-import RegistryOrchestrator inside function to avoid hard proxy coupling at module load, extend semantic_filter check to cover agent_tool_configs as well as mcp_tools, handle non-streaming ModelResponse follow-up by emitting as synthetic final chunk instead of silently dropping
This commit is contained in:
parent
bf863cb6c8
commit
3e41b0d716
3 changed files with 107 additions and 30 deletions
|
|
@ -21,6 +21,14 @@ ToolParam = Any
|
|||
|
||||
LITELLM_PROXY_AGENTS_URL = "litellm_proxy/agents"
|
||||
|
||||
# Import hoisted out of the callback loop to avoid re-evaluating on every iteration.
|
||||
try:
|
||||
from litellm.proxy.hooks.mcp_semantic_filter.hook import ( # noqa: E501
|
||||
SemanticToolFilterHook as _SemanticToolFilterHook,
|
||||
)
|
||||
except ImportError:
|
||||
_SemanticToolFilterHook = None # type: ignore
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Module-level helper
|
||||
|
|
@ -95,7 +103,7 @@ class RegistryOrchestrator:
|
|||
server_url = tool.get("server_url", "")
|
||||
if (
|
||||
isinstance(server_url, str)
|
||||
and LITELLM_PROXY_AGENTS_URL in server_url
|
||||
and server_url == LITELLM_PROXY_AGENTS_URL
|
||||
):
|
||||
agent_tool_configs.append(tool)
|
||||
else:
|
||||
|
|
@ -120,6 +128,11 @@ class RegistryOrchestrator:
|
|||
"""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
# NOTE: user_api_key_auth is accepted for future per-key agent filtering
|
||||
# (mirroring get_allowed_mcp_servers). Agent-level access control is not yet
|
||||
# implemented in AgentRegistry; all registered agents are returned for now.
|
||||
_ = user_api_key_auth
|
||||
|
||||
agents = global_agent_registry.get_agent_list()
|
||||
function_tools: List[Dict[str, Any]] = []
|
||||
agent_tool_map: Dict[str, Dict[str, str]] = {}
|
||||
|
|
@ -128,6 +141,13 @@ class RegistryOrchestrator:
|
|||
card = agent.agent_card_params or {}
|
||||
agent_url = card.get("url", "")
|
||||
agent_name = card.get("name") or agent.agent_name
|
||||
|
||||
if not agent_url:
|
||||
verbose_logger.warning(
|
||||
"Agent '%s' has no URL configured, skipping", agent_name
|
||||
)
|
||||
continue
|
||||
|
||||
description = card.get("description") or f"A2A agent: {agent_name}"
|
||||
|
||||
# Enrich description with up to 3 skill descriptions
|
||||
|
|
@ -146,6 +166,15 @@ class RegistryOrchestrator:
|
|||
or f"agent_{agent.agent_id[:8]}"
|
||||
)
|
||||
|
||||
# Deduplicate: if two agents produce the same sanitized name, append the
|
||||
# agent_id suffix so neither is silently dropped.
|
||||
if func_name in agent_tool_map:
|
||||
func_name = f"{func_name}_{agent.agent_id[:8]}"[:64]
|
||||
verbose_logger.warning(
|
||||
"Agent name collision: renamed to '%s' to avoid overwrite",
|
||||
func_name,
|
||||
)
|
||||
|
||||
function_tools.append(
|
||||
{
|
||||
"type": "function",
|
||||
|
|
@ -187,14 +216,20 @@ class RegistryOrchestrator:
|
|||
"""Send a message to an A2A agent via LiteLLM's asend_message and return the result."""
|
||||
import uuid
|
||||
|
||||
from a2a.types import (
|
||||
Message,
|
||||
MessageSendParams,
|
||||
Part,
|
||||
Role,
|
||||
SendMessageRequest,
|
||||
TextPart,
|
||||
)
|
||||
try:
|
||||
from a2a.types import (
|
||||
Message,
|
||||
MessageSendParams,
|
||||
Part,
|
||||
Role,
|
||||
SendMessageRequest,
|
||||
TextPart,
|
||||
)
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"The 'a2a' package is required for A2A agent calls. "
|
||||
"Install it with: pip install a2a-sdk"
|
||||
) from exc
|
||||
|
||||
from litellm.a2a_protocol.main import asend_message
|
||||
|
||||
|
|
@ -249,11 +284,9 @@ class RegistryOrchestrator:
|
|||
import litellm
|
||||
|
||||
for callback in litellm.callbacks or []:
|
||||
from litellm.proxy.hooks.mcp_semantic_filter.hook import (
|
||||
SemanticToolFilterHook,
|
||||
)
|
||||
|
||||
if isinstance(callback, SemanticToolFilterHook):
|
||||
if _SemanticToolFilterHook is None:
|
||||
break
|
||||
if isinstance(callback, _SemanticToolFilterHook):
|
||||
query = callback.filter.extract_user_query(messages)
|
||||
if query:
|
||||
filtered = await callback.filter.filter_tools(
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ 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,
|
||||
)
|
||||
|
|
@ -152,6 +151,7 @@ class MCPStreamingIterator:
|
|||
self.stream_exhausted = False
|
||||
self.tool_execution_done = False
|
||||
self.follow_up_stream: Optional[CustomStreamWrapper] = None
|
||||
self.follow_up_non_stream: Optional[ModelResponse] = None
|
||||
self.follow_up_iterator: Any = None
|
||||
self.follow_up_exhausted = False
|
||||
|
||||
|
|
@ -268,15 +268,28 @@ class MCPStreamingIterator:
|
|||
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"
|
||||
# Phase 3: emit non-streaming follow-up answer as a synthetic final chunk
|
||||
if self.follow_up_non_stream is not None:
|
||||
from litellm.types.utils import ModelResponseStream, StreamingChoices
|
||||
|
||||
non_stream = self.follow_up_non_stream
|
||||
self.follow_up_non_stream = None
|
||||
# Build a minimal streaming chunk from the ModelResponse
|
||||
content = ""
|
||||
if non_stream.choices:
|
||||
content = getattr(non_stream.choices[0].message, "content", "") or ""
|
||||
synthetic = ModelResponseStream(
|
||||
id=non_stream.id,
|
||||
model=non_stream.model or "",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta={"role": "assistant", "content": content}, # type: ignore[arg-type]
|
||||
)
|
||||
],
|
||||
)
|
||||
return synthetic
|
||||
|
||||
raise StopAsyncIteration
|
||||
|
||||
|
|
@ -349,12 +362,19 @@ class MCPStreamingIterator:
|
|||
if isinstance(follow_up_response, CustomStreamWrapper):
|
||||
self.follow_up_stream = follow_up_response
|
||||
verbose_logger.debug("Follow-up stream created successfully")
|
||||
elif isinstance(follow_up_response, ModelResponse):
|
||||
# Provider returned a non-streaming response despite stream=True.
|
||||
# Store it so __anext__ can yield it as a synthetic final chunk rather
|
||||
# than silently dropping the follow-up answer.
|
||||
self.follow_up_non_stream = follow_up_response
|
||||
verbose_logger.debug(
|
||||
"Follow-up response is non-streaming ModelResponse; will emit as final chunk"
|
||||
)
|
||||
else:
|
||||
verbose_logger.warning(
|
||||
"Follow-up response is not a CustomStreamWrapper: %s",
|
||||
"Follow-up response is unexpected type %s, answer may be dropped",
|
||||
type(follow_up_response),
|
||||
)
|
||||
self.follow_up_stream = None
|
||||
|
||||
|
||||
class MCPStreamWrapper(CustomStreamWrapper):
|
||||
|
|
@ -432,6 +452,7 @@ async def acompletion_with_mcp( # noqa: PLR0915
|
|||
5. Make a follow-up call with the tool results
|
||||
"""
|
||||
from litellm import acompletion as litellm_acompletion
|
||||
from litellm.proxy.agent_endpoints.registry_orchestrator import RegistryOrchestrator
|
||||
|
||||
# Parse MCP tools and separate from other tools
|
||||
(
|
||||
|
|
@ -482,10 +503,10 @@ async def acompletion_with_mcp( # noqa: PLR0915
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
)
|
||||
|
||||
# Apply per-tool semantic filter if any MCP tool has semantic_filter=true
|
||||
# Apply per-tool semantic filter if any MCP or agent tool config has semantic_filter=true
|
||||
if any(
|
||||
isinstance(t, dict) and t.get("semantic_filter")
|
||||
for t in mcp_tools_with_litellm_proxy
|
||||
for t in list(mcp_tools_with_litellm_proxy) + list(agent_tool_configs)
|
||||
):
|
||||
deduplicated_mcp_tools = await RegistryOrchestrator.apply_semantic_filter(
|
||||
tools=deduplicated_mcp_tools,
|
||||
|
|
|
|||
|
|
@ -548,7 +548,13 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
agent_tool_map: Optional[Dict[str, Dict[str, str]]] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
try:
|
||||
from fastapi import HTTPException
|
||||
except ImportError:
|
||||
# FastAPI is a proxy-only dependency; fall back to a plain exception so
|
||||
# SDK users (without fastapi installed) can still call MCP tools.
|
||||
# The .detail access below is already guarded with hasattr().
|
||||
HTTPException = Exception # type: ignore[assignment,misc]
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
|
|
@ -585,6 +591,11 @@ 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)
|
||||
verbose_logger.debug(
|
||||
"Executing A2A agent tool call: tool=%s agent=%s",
|
||||
tool_name,
|
||||
agent_info["agent_name"],
|
||||
)
|
||||
result = await RegistryOrchestrator.execute_a2a_tool_call(
|
||||
agent_url=agent_info["url"],
|
||||
agent_name=agent_info["agent_name"],
|
||||
|
|
@ -593,10 +604,22 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
tool_name=tool_name,
|
||||
litellm_trace_id=litellm_trace_id,
|
||||
)
|
||||
verbose_logger.debug(
|
||||
"A2A tool call complete: tool=%s result_len=%d",
|
||||
tool_name,
|
||||
len(str(result.get("result", ""))),
|
||||
)
|
||||
tool_results.append(result)
|
||||
continue
|
||||
|
||||
server_name = tool_server_map[tool_name]
|
||||
server_name = tool_server_map.get(tool_name)
|
||||
if server_name is None:
|
||||
verbose_logger.warning(
|
||||
"Tool '%s' not found in tool_server_map — skipping (possible "
|
||||
"hallucinated tool name)",
|
||||
tool_name,
|
||||
)
|
||||
continue
|
||||
|
||||
# Remove the server name prefix if the tool name includes it.
|
||||
sanitized_tool_name = tool_name
|
||||
|
|
@ -803,7 +826,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
error=e,
|
||||
)
|
||||
verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}")
|
||||
error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}"
|
||||
error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}" # type: ignore[union-attr]
|
||||
tool_results.append(
|
||||
{
|
||||
"tool_call_id": tool_call_id,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue