mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
fix mypy
This commit is contained in:
parent
002c3ad9da
commit
e7c8e55535
1 changed files with 102 additions and 92 deletions
|
|
@ -2818,6 +2818,103 @@ class MCPServerManager:
|
|||
|
||||
return cast(CallToolResult, result)
|
||||
|
||||
def _resolve_mcp_server_for_tool_call(
|
||||
self,
|
||||
server_name: str,
|
||||
name: str,
|
||||
) -> MCPServer:
|
||||
"""Resolve MCP server for call_tool (prefixed name, registry, fallback)."""
|
||||
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
|
||||
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
|
||||
resolved_by_server_name_only = False
|
||||
if mcp_server is None:
|
||||
for candidate in self.get_registry().values():
|
||||
if normalize_server_name(candidate.name) == normalize_server_name(
|
||||
server_name
|
||||
):
|
||||
mcp_server = candidate
|
||||
resolved_by_server_name_only = True
|
||||
break
|
||||
if mcp_server is None:
|
||||
fallback = self._get_mcp_server_from_tool_name(name)
|
||||
if fallback is not None and (
|
||||
not server_name
|
||||
or normalize_server_name(fallback.name)
|
||||
== normalize_server_name(server_name)
|
||||
):
|
||||
mcp_server = fallback
|
||||
if mcp_server is None:
|
||||
raise ValueError(f"Tool {name} not found")
|
||||
|
||||
if resolved_by_server_name_only:
|
||||
tool_known = (
|
||||
name in self.tool_name_to_mcp_server_name_mapping
|
||||
or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping
|
||||
)
|
||||
if not tool_known and self._mapping_has_tools_for_server(mcp_server):
|
||||
raise ValueError(f"Tool {name} not found")
|
||||
|
||||
return mcp_server
|
||||
|
||||
async def _resolve_oauth2_headers_for_tool_call(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: Optional[Dict[str, str]],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""Look up per-user OAuth headers when the client did not supply a token."""
|
||||
if (
|
||||
not mcp_server.needs_user_oauth_token
|
||||
or oauth2_headers
|
||||
or user_api_key_auth is None
|
||||
):
|
||||
return oauth2_headers
|
||||
|
||||
user_id = getattr(user_api_key_auth, "user_id", None)
|
||||
if not user_id:
|
||||
return oauth2_headers
|
||||
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
|
||||
_get_user_oauth_extra_headers_from_db,
|
||||
)
|
||||
|
||||
stored_headers = await _get_user_oauth_extra_headers_from_db(
|
||||
server=mcp_server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if stored_headers:
|
||||
return stored_headers
|
||||
except Exception as _lookup_exc:
|
||||
verbose_logger.debug(
|
||||
"call_tool: per-user token lookup failed for "
|
||||
"user=%s server=%s: %s",
|
||||
user_id,
|
||||
mcp_server.server_id,
|
||||
_lookup_exc,
|
||||
)
|
||||
return oauth2_headers
|
||||
|
||||
async def _gather_openapi_tool_tasks(
|
||||
self,
|
||||
tasks: List[Any],
|
||||
proxy_logging_obj: Optional[ProxyLogging],
|
||||
) -> CallToolResult:
|
||||
"""Await OpenAPI tool tasks and return the tool call result."""
|
||||
try:
|
||||
mcp_responses = await asyncio.gather(*tasks)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
return cast(CallToolResult, mcp_responses[result_index])
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
server_name: str,
|
||||
|
|
@ -2848,48 +2945,7 @@ class MCPServerManager:
|
|||
CallToolResult from the MCP server
|
||||
"""
|
||||
start_time = datetime.datetime.now()
|
||||
|
||||
# Resolve server (REST may pass server_name + prefixed or unprefixed tool name).
|
||||
# Prefer prefixed-name lookup first so the resolution validates that the
|
||||
# tool actually belongs to server_name. Fall back to name-based server
|
||||
# resolution only when the tool->server mapping isn't populated (e.g.
|
||||
# OAuth2 servers skipped during init).
|
||||
mcp_server: Optional[MCPServer] = None
|
||||
resolved_by_server_name_only = False
|
||||
prefixed_tool_name = add_server_prefix_to_name(name, server_name)
|
||||
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
|
||||
if mcp_server is None:
|
||||
for candidate in self.get_registry().values():
|
||||
if normalize_server_name(candidate.name) == normalize_server_name(
|
||||
server_name
|
||||
):
|
||||
mcp_server = candidate
|
||||
resolved_by_server_name_only = True
|
||||
break
|
||||
if mcp_server is None:
|
||||
# Last resort: lookup by unprefixed tool name. Only accept the
|
||||
# match when it agrees with the caller's server_name — otherwise
|
||||
# we'd silently dispatch to a different server than requested.
|
||||
fallback = self._get_mcp_server_from_tool_name(name)
|
||||
if fallback is not None and (
|
||||
not server_name
|
||||
or normalize_server_name(fallback.name)
|
||||
== normalize_server_name(server_name)
|
||||
):
|
||||
mcp_server = fallback
|
||||
if mcp_server is None:
|
||||
raise ValueError(f"Tool {name} not found")
|
||||
|
||||
# Server-name fallback is for mapping-not-yet-populated paths (e.g. REST).
|
||||
# If this server already has tools in the mapping, unknown names should fail
|
||||
# fast instead of opening a real upstream session.
|
||||
if resolved_by_server_name_only:
|
||||
tool_known = (
|
||||
name in self.tool_name_to_mcp_server_name_mapping
|
||||
or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping
|
||||
)
|
||||
if not tool_known and self._mapping_has_tools_for_server(mcp_server):
|
||||
raise ValueError(f"Tool {name} not found")
|
||||
mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name)
|
||||
|
||||
#########################################################
|
||||
# Pre MCP Tool Call Hook
|
||||
|
|
@ -2923,36 +2979,9 @@ class MCPServerManager:
|
|||
)
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
# For per-user OAuth servers: if the client didn't supply a token in
|
||||
# oauth2_headers, look up the stored token from Redis / DB. This is the
|
||||
# call_tool equivalent of _get_user_oauth_extra_headers_from_db used in
|
||||
# list_tools.
|
||||
if (
|
||||
mcp_server.needs_user_oauth_token
|
||||
and not oauth2_headers
|
||||
and user_api_key_auth is not None
|
||||
):
|
||||
user_id = getattr(user_api_key_auth, "user_id", None)
|
||||
if user_id:
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415
|
||||
_get_user_oauth_extra_headers_from_db,
|
||||
)
|
||||
|
||||
stored_headers = await _get_user_oauth_extra_headers_from_db(
|
||||
server=mcp_server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if stored_headers:
|
||||
oauth2_headers = stored_headers
|
||||
except Exception as _lookup_exc:
|
||||
verbose_logger.debug(
|
||||
"call_tool: per-user token lookup failed for "
|
||||
"user=%s server=%s: %s",
|
||||
user_id,
|
||||
mcp_server.server_id,
|
||||
_lookup_exc,
|
||||
)
|
||||
oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(
|
||||
mcp_server, oauth2_headers, user_api_key_auth
|
||||
)
|
||||
|
||||
# For OpenAPI servers, call the tool handler directly instead of via MCP client
|
||||
if mcp_server.spec_path:
|
||||
|
|
@ -2988,26 +3017,7 @@ class MCPServerManager:
|
|||
hook_extra_headers=hook_result.get("extra_headers"),
|
||||
)
|
||||
|
||||
# For OpenAPI tools, await outside the client context
|
||||
try:
|
||||
mcp_responses = await asyncio.gather(*tasks)
|
||||
|
||||
# If proxy_logging_obj is None, the tool call result is at index 0
|
||||
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
result = mcp_responses[result_index]
|
||||
|
||||
return cast(CallToolResult, result)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
# Re-raise guardrail exceptions to properly fail the MCP call
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
|
||||
|
||||
#########################################################
|
||||
# End of Methods that call the upstream MCP servers
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue