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:
Ishaan Jaffer 2026-03-21 17:49:52 -07:00
parent bf863cb6c8
commit 3e41b0d716
3 changed files with 107 additions and 30 deletions

View file

@ -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(

View file

@ -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,

View file

@ -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,