mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(mcp): pass keyed OAuth grants through explicit operation contexts
This commit is contained in:
parent
77ffddbf70
commit
23ea0ee428
10 changed files with 267 additions and 69 deletions
|
|
@ -35,6 +35,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
ConnectionBinding,
|
||||
ConnectionCredential,
|
||||
EnvelopeIdentity,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
|
||||
|
|
@ -552,49 +553,15 @@ class MCPRequestHandler:
|
|||
scope.pop(CONNECTION_SCOPE_KEY, None)
|
||||
connection_header: Final = headers.get("authorization")
|
||||
if is_connection_credential(connection_header):
|
||||
if not has_explicit_litellm_key or request_route != "/mcp":
|
||||
raise HTTPException(status_code=401, detail="A connection credential requires the original MCP key")
|
||||
targets: Final = MCPRequestHandler._resolve_target_server_names(request_route, mcp_servers)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
target: Final = (
|
||||
global_mcp_server_manager.get_mcp_server_by_name(
|
||||
targets[0], client_ip=IPAddressUtils.get_mcp_client_ip(request)
|
||||
)
|
||||
if len(targets) == 1
|
||||
else None
|
||||
scope[CONNECTION_SCOPE_KEY] = await MCPRequestHandler._admit_connection_credential(
|
||||
request=request,
|
||||
request_route=request_route,
|
||||
connection_header=connection_header or "",
|
||||
litellm_api_key=litellm_api_key,
|
||||
mcp_servers=mcp_servers,
|
||||
validated_user_api_key_auth=validated_user_api_key_auth,
|
||||
has_explicit_litellm_key=has_explicit_litellm_key,
|
||||
)
|
||||
allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(validated_user_api_key_auth)
|
||||
if (
|
||||
target is None
|
||||
or target.server_id not in allowed
|
||||
or not target.is_gateway_managed_oauth2
|
||||
or not target.needs_user_oauth_token
|
||||
or target.oauth_identity_binding is not None
|
||||
):
|
||||
raise HTTPException(status_code=403, detail="Connection credential does not authorize this MCP server")
|
||||
expected_binding: Final = ConnectionBinding(
|
||||
key_hash=hash_token(_get_bearer_token_or_received_api_key(litellm_api_key)),
|
||||
server_id=target.server_id,
|
||||
resource=f"{get_request_base_url(request)}/mcp",
|
||||
)
|
||||
connection: Final = open_connection_credential(connection_header or "")
|
||||
if connection is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid or expired MCP connection credential",
|
||||
headers=MappingProxyType(
|
||||
{
|
||||
"www-authenticate": connection_challenge(request, expected_binding),
|
||||
"Cache-Control": "no-store",
|
||||
}
|
||||
),
|
||||
)
|
||||
if connection.binding != expected_binding:
|
||||
raise HTTPException(
|
||||
status_code=401, detail="Connection credential belongs to a different key or resource"
|
||||
)
|
||||
scope[CONNECTION_SCOPE_KEY] = connection
|
||||
|
||||
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
|
||||
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
|
||||
|
|
@ -623,6 +590,58 @@ class MCPRequestHandler:
|
|||
raw_headers,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _admit_connection_credential(
|
||||
request: Request,
|
||||
request_route: str,
|
||||
connection_header: str,
|
||||
litellm_api_key: str,
|
||||
mcp_servers: list[str] | None,
|
||||
validated_user_api_key_auth: UserAPIKeyAuth,
|
||||
has_explicit_litellm_key: bool,
|
||||
) -> ConnectionCredential:
|
||||
if not has_explicit_litellm_key or request_route != "/mcp":
|
||||
raise HTTPException(status_code=401, detail="A connection credential requires the original MCP key")
|
||||
targets: Final = MCPRequestHandler._resolve_target_server_names(request_route, mcp_servers)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
target: Final = (
|
||||
global_mcp_server_manager.get_mcp_server_by_name(
|
||||
targets[0], client_ip=IPAddressUtils.get_mcp_client_ip(request)
|
||||
)
|
||||
if len(targets) == 1
|
||||
else None
|
||||
)
|
||||
allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(validated_user_api_key_auth)
|
||||
if (
|
||||
target is None
|
||||
or target.server_id not in allowed
|
||||
or not target.is_gateway_managed_oauth2
|
||||
or not target.needs_user_oauth_token
|
||||
or target.oauth_identity_binding is not None
|
||||
):
|
||||
raise HTTPException(status_code=403, detail="Connection credential does not authorize this MCP server")
|
||||
expected_binding: Final = ConnectionBinding(
|
||||
key_hash=hash_token(_get_bearer_token_or_received_api_key(litellm_api_key)),
|
||||
server_id=target.server_id,
|
||||
resource=f"{get_request_base_url(request)}/mcp",
|
||||
)
|
||||
connection: Final = open_connection_credential(connection_header)
|
||||
if connection is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid or expired MCP connection credential",
|
||||
headers=MappingProxyType(
|
||||
{
|
||||
"www-authenticate": connection_challenge(request, expected_binding),
|
||||
"Cache-Control": "no-store",
|
||||
}
|
||||
),
|
||||
)
|
||||
if connection.binding != expected_binding:
|
||||
raise HTTPException(status_code=401, detail="Connection credential belongs to a different key or resource")
|
||||
return connection
|
||||
|
||||
@staticmethod
|
||||
def _is_gateway_admission_credential(value: str | None) -> bool:
|
||||
"""True when a header value is a gateway admission credential — a session bearer or bridge
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from datetime import datetime
|
|||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
|
@ -26,6 +27,7 @@ class OperationContext:
|
|||
raw_headers: Mapping[str, str] | None = field(default=None, repr=False)
|
||||
client_ip: str | None = None
|
||||
mcp_proxy_mode: bool = False
|
||||
connection_credential: ConnectionCredential | None = field(default=None, repr=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "_caller", copy_caller(self._caller))
|
||||
|
|
|
|||
|
|
@ -11,8 +11,6 @@ from typing import TYPE_CHECKING, Final
|
|||
if TYPE_CHECKING:
|
||||
from mcp.server.context import ServerRequestContext
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
|
||||
|
||||
# The SDK 1.x ``mcp.server.lowlevel.server.request_ctx`` ContextVar was removed in
|
||||
# SDK 2, which hands each request handler a ``ServerRequestContext`` argument
|
||||
# instead. The handlers set this var so downstream helpers (session auth caching,
|
||||
|
|
@ -42,19 +40,3 @@ _mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gatew
|
|||
|
||||
# Set server-side by the /mcp/proxy route. Never populated from client-supplied headers.
|
||||
_mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False)
|
||||
|
||||
|
||||
def get_connection_credential(server_id: str) -> "ConnectionCredential | None":
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
|
||||
|
||||
context: Final = get_active_mcp_request_ctx()
|
||||
request: Final = context.request if context is not None else None
|
||||
if not isinstance(request, Request):
|
||||
return None
|
||||
value: Final = request.scope.get("litellm.mcp.connection_grant")
|
||||
if not isinstance(value, ConnectionCredential) or value.binding.server_id != server_id:
|
||||
return None
|
||||
return value
|
||||
|
|
|
|||
|
|
@ -111,6 +111,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
|
|||
to_server_spec,
|
||||
to_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
InvalidatableOAuthTokenStore,
|
||||
)
|
||||
|
|
@ -4108,6 +4109,7 @@ class MCPServerManager:
|
|||
cred_provider: UpstreamCredentialProvider | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
|
@ -4133,13 +4135,17 @@ class MCPServerManager:
|
|||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
transport: Final = resolved_server.transport or MCPTransport.sse
|
||||
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import get_connection_credential
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.presented_token_store import (
|
||||
PresentedOAuthTokenStore,
|
||||
)
|
||||
|
||||
connection: Final = get_connection_credential(resolved_server.server_id)
|
||||
connection: Final = (
|
||||
connection_credential
|
||||
if connection_credential is not None
|
||||
and connection_credential.binding.server_id == resolved_server.server_id
|
||||
else None
|
||||
)
|
||||
if connection is not None and connection.exp <= int(datetime.datetime.now(datetime.timezone.utc).timestamp()):
|
||||
raise HTTPException(status_code=401, detail="MCP connection credential expired; reconnect")
|
||||
provider: Final = (
|
||||
|
|
@ -4173,7 +4179,9 @@ class MCPServerManager:
|
|||
sampling_cb = (
|
||||
_create_sampling_callback(
|
||||
operation_context=OperationContext(
|
||||
_caller=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip
|
||||
_caller=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
)
|
||||
if resolved_server.allow_sampling
|
||||
|
|
@ -4323,6 +4331,7 @@ class MCPServerManager:
|
|||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
oauth2_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
|
@ -4414,6 +4423,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
|
||||
## HANDLE OPENAPI TOOLS
|
||||
|
|
@ -4525,6 +4535,7 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[Prompt]:
|
||||
try:
|
||||
headers: Final = (
|
||||
|
|
@ -4547,6 +4558,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
|
||||
key: Final = self._discovery_key(
|
||||
|
|
@ -4571,6 +4583,7 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[Resource]:
|
||||
try:
|
||||
headers: Final = (
|
||||
|
|
@ -4593,6 +4606,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
|
||||
key: Final = self._discovery_key(
|
||||
|
|
@ -4617,6 +4631,7 @@ class MCPServerManager:
|
|||
add_prefix: bool = True,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[ResourceTemplate]:
|
||||
try:
|
||||
headers: Final = (
|
||||
|
|
@ -4639,6 +4654,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
credential_fingerprint: Final = await client.discovery_auth_fingerprint()
|
||||
key: Final = self._discovery_key(
|
||||
|
|
@ -4663,6 +4679,7 @@ class MCPServerManager:
|
|||
extra_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> ReadResourceResult:
|
||||
"""Read resource contents from a specific MCP server."""
|
||||
|
||||
|
|
@ -4686,6 +4703,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
|
||||
return await client.read_resource(url)
|
||||
|
|
@ -4700,6 +4718,7 @@ class MCPServerManager:
|
|||
extra_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> GetPromptResult:
|
||||
"""Fetch a specific prompt definition from a single MCP server."""
|
||||
|
||||
|
|
@ -4723,6 +4742,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
|
||||
get_prompt_request_params: Final = GetPromptRequestParams(
|
||||
|
|
@ -5805,6 +5825,7 @@ class MCPServerManager:
|
|||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
"""Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry.
|
||||
|
||||
|
|
@ -5832,6 +5853,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback)
|
||||
|
||||
|
|
@ -5850,6 +5872,7 @@ class MCPServerManager:
|
|||
hook_extra_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call a regular MCP tool using the MCP client.
|
||||
|
|
@ -5996,6 +6019,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
|
||||
call_tool_params: Final = MCPCallToolRequestParams(
|
||||
|
|
@ -6021,6 +6045,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
|
||||
tool_call_coro = _obo_call_tool_limited()
|
||||
|
|
@ -6303,6 +6328,7 @@ class MCPServerManager:
|
|||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
|
@ -6434,6 +6460,7 @@ class MCPServerManager:
|
|||
host_progress_callback=host_progress_callback,
|
||||
hook_extra_headers=hook_result.get("extra_headers"),
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
connection_credential=connection_credential,
|
||||
)
|
||||
|
||||
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
|
||||
|
|
|
|||
|
|
@ -85,6 +85,7 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
|||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
|
@ -250,6 +251,7 @@ async def _dispatch_virtual_mcp_tool(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult | None:
|
||||
"""Handle the mcp_tool_search / mcp_tool_call virtual tools.
|
||||
|
||||
|
|
@ -297,6 +299,7 @@ async def _dispatch_virtual_mcp_tool(
|
|||
)
|
||||
try:
|
||||
proxy_result: Final = await handle_mcp_proxy_tool(
|
||||
connection_credential=connection_credential,
|
||||
name=name,
|
||||
arguments=arguments or {}, # mutable-ok: proxy handler payload
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
|
|
@ -364,6 +367,7 @@ async def _dispatch_virtual_mcp_tool(
|
|||
args: Final = arguments or {}
|
||||
if name == MCP_TOOL_SEARCH_TOOL_NAME:
|
||||
return await handle_mcp_tool_search(
|
||||
connection_credential=connection_credential,
|
||||
query=TypeAdapter(str).validate_python(args.get("query", "")),
|
||||
top_k=coerce_top_k(args.get("top_k", 5)),
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
|
|
@ -399,6 +403,7 @@ async def _dispatch_virtual_mcp_tool(
|
|||
types.MappingProxyType({"name": args.get("tool_name", ""), "arguments": args.get("arguments") or {}})
|
||||
)
|
||||
return await handle_mcp_tool_call(
|
||||
connection_credential=connection_credential,
|
||||
tool_name=tool_request.name,
|
||||
arguments=tool_request.arguments or {},
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
|
|
@ -943,6 +948,7 @@ async def _get_tools_from_mcp_servers(
|
|||
request_tags: list[str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -1105,6 +1111,7 @@ async def _get_tools_from_mcp_servers(
|
|||
|
||||
try:
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
connection_credential=connection_credential,
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
|
|
@ -1230,6 +1237,7 @@ async def _get_prompts_from_mcp_servers(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[Prompt]:
|
||||
"""
|
||||
Helper method to fetch prompt from MCP servers based on server filtering criteria.
|
||||
|
|
@ -1269,6 +1277,7 @@ async def _get_prompts_from_mcp_servers(
|
|||
|
||||
try:
|
||||
prompts = await global_mcp_server_manager.get_prompts_from_server(
|
||||
connection_credential=connection_credential,
|
||||
server=server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1298,6 +1307,7 @@ async def _get_resources_from_mcp_servers(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[Resource]:
|
||||
"""Fetch resources from allowed MCP servers."""
|
||||
|
||||
|
|
@ -1324,6 +1334,7 @@ async def _get_resources_from_mcp_servers(
|
|||
|
||||
try:
|
||||
resources = await global_mcp_server_manager.get_resources_from_server(
|
||||
connection_credential=connection_credential,
|
||||
server=server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1351,6 +1362,7 @@ async def _get_resource_templates_from_mcp_servers(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[ResourceTemplate]:
|
||||
"""Fetch resource templates from allowed MCP servers."""
|
||||
|
||||
|
|
@ -1377,6 +1389,7 @@ async def _get_resource_templates_from_mcp_servers(
|
|||
|
||||
try:
|
||||
resource_templates = await global_mcp_server_manager.get_resource_templates_from_server(
|
||||
connection_credential=connection_credential,
|
||||
server=server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1446,6 +1459,7 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source: str | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -1464,6 +1478,7 @@ async def _list_mcp_tools(
|
|||
|
||||
try:
|
||||
listing: Final = await _get_tools_from_mcp_servers(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -1493,6 +1508,7 @@ async def _list_mcp_prompts(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[Prompt]:
|
||||
"""
|
||||
List all available MCP prompts.
|
||||
|
|
@ -1510,6 +1526,7 @@ async def _list_mcp_prompts(
|
|||
managed_prompts = []
|
||||
try:
|
||||
managed_prompts = await _get_prompts_from_mcp_servers(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -1534,12 +1551,14 @@ async def _list_mcp_resources(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[Resource]:
|
||||
"""List all available MCP resources."""
|
||||
|
||||
managed_resources: list[Resource] = []
|
||||
try:
|
||||
managed_resources = await _get_resources_from_mcp_servers(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -1563,12 +1582,14 @@ async def _list_mcp_resource_templates(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> list[ResourceTemplate]:
|
||||
"""List all available MCP resource templates."""
|
||||
|
||||
managed_resource_templates: list[ResourceTemplate] = []
|
||||
try:
|
||||
managed_resource_templates = await _get_resource_templates_from_mcp_servers(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -1731,6 +1752,7 @@ async def _list_tools_before_first_call(
|
|||
oauth2_headers: dict[str, str] | None,
|
||||
raw_headers: dict[str, str] | None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> None:
|
||||
"""List ``server`` with the caller's own credentials when it does not yet expose ``tool_name`` here.
|
||||
|
||||
|
|
@ -1746,6 +1768,7 @@ async def _list_tools_before_first_call(
|
|||
return
|
||||
try:
|
||||
await _get_tools_from_mcp_servers(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=[server.server_id],
|
||||
|
|
@ -1771,9 +1794,11 @@ async def execute_mcp_tool(
|
|||
host_progress_callback: ProgressCallback | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
**kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract
|
||||
) -> CallToolResult:
|
||||
context: Final = prepare_context(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
|
|
@ -1806,6 +1831,7 @@ async def _execute_mcp_tool(
|
|||
host_progress_callback: ProgressCallback | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
|
|
@ -1865,6 +1891,7 @@ async def _execute_mcp_tool(
|
|||
else strip_known_server_prefix(name, first_call_target)
|
||||
)
|
||||
await _list_tools_before_first_call(
|
||||
connection_credential=connection_credential,
|
||||
server=first_call_target,
|
||||
tool_name=first_call_tool_name,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
|
|
@ -2059,6 +2086,7 @@ async def _execute_mcp_tool(
|
|||
#########################################################
|
||||
elif mcp_server:
|
||||
response = await _handle_managed_mcp_tool(
|
||||
connection_credential=connection_credential,
|
||||
server_name=server_name,
|
||||
name=original_tool_name,
|
||||
arguments=arguments,
|
||||
|
|
@ -2281,6 +2309,7 @@ async def call_mcp_tool(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
|
|
@ -2325,6 +2354,7 @@ async def call_mcp_tool(
|
|||
|
||||
# Delegate to execute_mcp_tool for execution
|
||||
response = await execute_mcp_tool(
|
||||
connection_credential=connection_credential,
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
|
|
@ -2363,6 +2393,7 @@ async def mcp_get_prompt(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> GetPromptResult:
|
||||
"""
|
||||
Fetch a specific MCP prompt, handling both prefixed and unprefixed names.
|
||||
|
|
@ -2399,6 +2430,7 @@ async def mcp_get_prompt(
|
|||
)
|
||||
|
||||
return await global_mcp_server_manager.get_prompt_from_server(
|
||||
connection_credential=connection_credential,
|
||||
server=server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
prompt_name=original_prompt_name,
|
||||
|
|
@ -2419,6 +2451,7 @@ async def mcp_read_resource(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> ReadResourceResult:
|
||||
"""Read resource contents from upstream MCP servers."""
|
||||
|
||||
|
|
@ -2452,6 +2485,7 @@ async def mcp_read_resource(
|
|||
)
|
||||
|
||||
return await global_mcp_server_manager.read_resource_from_server(
|
||||
connection_credential=connection_credential,
|
||||
server=server,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
url=url,
|
||||
|
|
@ -2506,12 +2540,14 @@ async def _handle_managed_mcp_tool(
|
|||
host_progress_callback: ProgressCallback | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
"""Handle tool execution for managed server tools"""
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
call_tool_result: Final = await global_mcp_server_manager.call_tool(
|
||||
connection_credential=connection_credential,
|
||||
server_name=server_name,
|
||||
name=name,
|
||||
arguments=arguments,
|
||||
|
|
@ -2576,6 +2612,7 @@ _MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset(
|
|||
"mcp_server_auth_headers",
|
||||
"oauth2_headers",
|
||||
"user_api_key_auth",
|
||||
"connection_credential",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -2622,6 +2659,7 @@ async def _execute_handle_list_tools(
|
|||
# Get mcp_servers from context variable
|
||||
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
|
||||
listing: Final = await _list_mcp_tools(
|
||||
connection_credential=context.connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -2681,6 +2719,7 @@ async def _execute_mcp_server_tool_call(
|
|||
# Inside this try so virtual-tool errors convert to isError
|
||||
# CallToolResult instead of raising out of the protocol handler.
|
||||
virtual_tool_result: Final = await _dispatch_virtual_mcp_tool(
|
||||
connection_credential=context.connection_credential,
|
||||
name=params.name,
|
||||
arguments=params.arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
|
|
@ -2729,6 +2768,7 @@ async def _execute_mcp_server_tool_call(
|
|||
data = body_data
|
||||
|
||||
response: Final = await call_mcp_tool(
|
||||
connection_credential=context.connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -2822,6 +2862,7 @@ async def _execute_list_prompts(
|
|||
# Get mcp_servers from context variable
|
||||
verbose_logger.debug("MCP list_prompts - Calling _list_prompts")
|
||||
prompts: Final = await _list_mcp_prompts(
|
||||
connection_credential=context.connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -2856,6 +2897,7 @@ async def _execute_get_prompt(
|
|||
|
||||
verbose_logger.debug("MCP mcp_server_tool_call - User API Key Auth from context: %s", user_api_key_auth)
|
||||
return await mcp_get_prompt(
|
||||
connection_credential=context.connection_credential,
|
||||
name=params.name,
|
||||
arguments=params.arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
|
|
@ -2891,6 +2933,7 @@ async def _execute_list_resources(
|
|||
)
|
||||
|
||||
resources: Final = await _list_mcp_resources(
|
||||
connection_credential=context.connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -2929,6 +2972,7 @@ async def _execute_list_resource_templates(
|
|||
)
|
||||
|
||||
resource_templates: Final = await _list_mcp_resource_templates(
|
||||
connection_credential=context.connection_credential,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
|
|
@ -2962,6 +3006,7 @@ async def _execute_read_resource(
|
|||
) = context.legacy_auth()
|
||||
|
||||
read_resource_result: Final = await mcp_read_resource(
|
||||
connection_credential=context.connection_credential,
|
||||
url=params.uri,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -2991,8 +3036,10 @@ def prepare_context(
|
|||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> OperationContext:
|
||||
return OperationContext(
|
||||
connection_credential=connection_credential,
|
||||
_caller=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=tuple(mcp_servers) if mcp_servers is not None else None,
|
||||
|
|
@ -3060,6 +3107,7 @@ class GatewayOperations:
|
|||
case AuthorizedToolCall():
|
||||
auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth()
|
||||
return await _execute_mcp_tool(
|
||||
connection_credential=context.connection_credential,
|
||||
name=operation.name,
|
||||
arguments=dict(operation.arguments), # mutable-ok: existing tool hooks own mutable argument data
|
||||
allowed_mcp_servers=list(
|
||||
|
|
|
|||
|
|
@ -828,8 +828,19 @@ if MCP_AVAILABLE:
|
|||
headers,
|
||||
client_ip,
|
||||
) = await get_or_extract_auth_context()
|
||||
credential: Final = (
|
||||
ctx.request.scope.get(CONNECTION_SCOPE_KEY) if isinstance(ctx.request, StarletteRequest) else None
|
||||
)
|
||||
yield operations.prepare_context(
|
||||
auth, token, servers, server_headers, oauth_headers, headers, client_ip, _mcp_proxy_mode.get()
|
||||
auth,
|
||||
token,
|
||||
servers,
|
||||
server_headers,
|
||||
oauth_headers,
|
||||
headers,
|
||||
client_ip,
|
||||
_mcp_proxy_mode.get(),
|
||||
connection_credential=credential if isinstance(credential, ConnectionCredential) else None,
|
||||
)
|
||||
|
||||
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ if TYPE_CHECKING:
|
|||
from mcp.types import CallToolResult, Tool
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
MCP_TOOL_SEARCH_SETTINGS_KEY: Final[str] = "mcp_tool_search"
|
||||
|
|
@ -462,6 +463,7 @@ async def handle_mcp_tool_search(
|
|||
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
|
||||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
from litellm.proxy._experimental.mcp_server.operations import (
|
||||
_list_mcp_tools,
|
||||
|
|
@ -488,6 +490,7 @@ async def handle_mcp_tool_search(
|
|||
else None
|
||||
)
|
||||
mcp_listing: Final = await _list_mcp_tools(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=client_ip,
|
||||
|
|
@ -513,6 +516,7 @@ async def handle_mcp_proxy_tool(
|
|||
oauth2_headers: dict[str, str] | None = None, # mutable-ok: preserve forwarded headers
|
||||
raw_headers: dict[str, str] | None = None, # mutable-ok: preserve request headers
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
from fastapi import HTTPException
|
||||
from jsonschema import ValidationError as JsonSchemaValidationError
|
||||
|
|
@ -524,6 +528,7 @@ async def handle_mcp_proxy_tool(
|
|||
)
|
||||
|
||||
listing: Final = await _list_mcp_tools(
|
||||
connection_credential=connection_credential,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=client_ip,
|
||||
|
|
@ -579,6 +584,7 @@ async def handle_mcp_proxy_tool(
|
|||
return _text_tool_result(f"Invalid arguments: {exc.message}", is_error=True)
|
||||
|
||||
return await handle_mcp_tool_call(
|
||||
connection_credential=connection_credential,
|
||||
tool_name=_mcp_proxy_identity(tool)["tool_name"],
|
||||
arguments=tool_arguments,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -606,6 +612,7 @@ async def handle_mcp_tool_call(
|
|||
litellm_logging_obj: LiteLLMLoggingObj | None = None,
|
||||
requested_server_id: str | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
connection_credential: ConnectionCredential | None = None,
|
||||
) -> CallToolResult:
|
||||
from litellm.proxy._experimental.mcp_server.operations import (
|
||||
_get_allowed_mcp_servers,
|
||||
|
|
@ -634,6 +641,7 @@ async def handle_mcp_tool_call(
|
|||
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
|
||||
|
||||
return await execute_mcp_tool(
|
||||
connection_credential=connection_credential,
|
||||
name=tool_name,
|
||||
arguments=arguments,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
|
|
|
|||
|
|
@ -1080,6 +1080,7 @@ async def test_mcp_get_prompt_success():
|
|||
extra_headers={"X-Test": "1"},
|
||||
raw_headers=None,
|
||||
client_ip=None,
|
||||
connection_credential=None,
|
||||
)
|
||||
assert result is prompt_result
|
||||
|
||||
|
|
@ -1143,6 +1144,7 @@ async def test_mcp_read_resource_success():
|
|||
extra_headers={"X-Test": "1"},
|
||||
raw_headers=None,
|
||||
client_ip=None,
|
||||
connection_credential=None,
|
||||
)
|
||||
assert result is read_result
|
||||
|
||||
|
|
@ -8986,6 +8988,7 @@ async def test_fire_mcp_tool_call_logging_strips_credentials_from_failure_hook()
|
|||
"mcp_auth_header": "upstream-secret",
|
||||
"mcp_server_auth_headers": {"srv": {"authorization": "Bearer srv-secret"}},
|
||||
"oauth2_headers": {"authorization": "Bearer oauth-secret"},
|
||||
"connection_credential": "connection-secret",
|
||||
"user_api_key_auth": user_auth,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -4070,6 +4070,7 @@ class TestMCPServerManager:
|
|||
user_api_key_auth=None,
|
||||
raw_headers=None,
|
||||
client_ip=None,
|
||||
connection_credential=None,
|
||||
)
|
||||
mock_client.list_resource_templates.assert_awaited_once()
|
||||
assert result == expected_templates
|
||||
|
|
@ -14513,7 +14514,7 @@ async def test_connection_grants_follow_current_message_and_never_leak_to_anothe
|
|||
)
|
||||
reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request))
|
||||
try:
|
||||
client = await manager._create_mcp_client(server)
|
||||
client = await manager._create_mcp_client(server, connection_credential=credential)
|
||||
sent = await client.prepare_request_auth()
|
||||
assert sent.headers["authorization"] == f"Bearer {token}"
|
||||
store.fetch.assert_not_awaited()
|
||||
|
|
@ -14522,10 +14523,14 @@ async def test_connection_grants_follow_current_message_and_never_leak_to_anothe
|
|||
|
||||
reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request))
|
||||
try:
|
||||
other_client = await manager._create_mcp_client(other)
|
||||
other_client = await manager._create_mcp_client(other, connection_credential=credential)
|
||||
other_sent = await other_client.prepare_request_auth()
|
||||
assert other_sent.headers["authorization"] == "Bearer saved-vault-token"
|
||||
assert store.fetch.call_args.args[1] == "other-target"
|
||||
explicit_client = await manager._create_mcp_client(
|
||||
server, user_api_key_auth=UserAPIKeyAuth(user_id="independent-caller")
|
||||
)
|
||||
assert (await explicit_client.prepare_request_auth()).headers["authorization"] == "Bearer saved-vault-token"
|
||||
finally:
|
||||
active_mcp_request_ctx_var.reset(reset)
|
||||
saved_client = await manager._create_mcp_client(server)
|
||||
|
|
@ -14570,7 +14575,7 @@ async def test_connection_expiry_between_admission_and_egress_never_uses_vault()
|
|||
reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request))
|
||||
try:
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await manager._create_mcp_client(server)
|
||||
await manager._create_mcp_client(server, connection_credential=credential)
|
||||
assert exc.value.status_code == 401
|
||||
assert "expired" in exc.value.detail
|
||||
store.fetch.assert_not_awaited()
|
||||
|
|
@ -14609,3 +14614,96 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
|
|||
assert captured["client_ip"] is None
|
||||
finally:
|
||||
auth_context_var.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", [
|
||||
"tools/list", "tools/call", "prompts/list", "prompts/get", "resources/list",
|
||||
"resources/templates/list", "resources/read", "virtual/search", "virtual/call",
|
||||
"proxy/search", "proxy/schema", "proxy/call",
|
||||
])
|
||||
async def test_native_operations_send_only_current_connection_credential(method):
|
||||
from types import SimpleNamespace
|
||||
from datetime import timezone
|
||||
from mcp import types
|
||||
from mcp.server.context import ServerRequestContext
|
||||
from pydantic import SecretStr
|
||||
from starlette.requests import Request
|
||||
from litellm.proxy._experimental.mcp_server import operations, server as ingress
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import CONNECTION_SCOPE_KEY
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import UpstreamCredentialProvider
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionBinding, ConnectionCredential
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
|
||||
|
||||
store = SimpleNamespace(fetch=AsyncMock(return_value=OAuthToken(access_token="saved-vault-token")))
|
||||
manager = MCPServerManager(cred_provider=UpstreamCredentialProvider(oauth_token_store=store))
|
||||
target = MCPServer(
|
||||
server_id="catalog", name="catalog", server_name="catalog", url="https://catalog.example/mcp",
|
||||
transport="http", auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code",
|
||||
)
|
||||
manager.registry = {target.server_id: target}
|
||||
upstream = _DiscoveryUpstream()
|
||||
tool = {"name": "example", "description": "Example tool", "inputSchema": {"type": "object", "properties": {}}}
|
||||
|
||||
async def respond(request):
|
||||
payload = _JSONRPC_ADAPTER.validate_json(request.content) if request.method == "POST" else None
|
||||
if isinstance(payload, types.JSONRPCRequest) and payload.method in ("tools/list", "tools/call", "prompts/get", "resources/read"):
|
||||
upstream.requests = (*upstream.requests, (payload.method, request.headers.get("authorization", "")))
|
||||
results = {
|
||||
"tools/list": {"tools": [tool]},
|
||||
"tools/call": {"content": [{"type": "text", "text": "executed"}], "isError": False},
|
||||
"prompts/get": {"messages": []},
|
||||
"resources/read": {"contents": [{"uri": "test://example", "text": "resource body"}]},
|
||||
}
|
||||
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": results[payload.method]})
|
||||
return await upstream.respond(request)
|
||||
|
||||
requests = {
|
||||
"tools/list": types.ListToolsRequest(),
|
||||
"tools/call": types.CallToolRequest(params=types.CallToolRequestParams(name="catalog-example", arguments={})),
|
||||
"prompts/list": types.ListPromptsRequest(),
|
||||
"prompts/get": types.GetPromptRequest(params=types.GetPromptRequestParams(name="catalog-example")),
|
||||
"resources/list": types.ListResourcesRequest(),
|
||||
"resources/templates/list": types.ListResourceTemplatesRequest(),
|
||||
"resources/read": types.ReadResourceRequest(params=types.ReadResourceRequestParams(uri="test://example")),
|
||||
"virtual/search": types.CallToolRequest(params=types.CallToolRequestParams(name="mcp_tool_search", arguments={"query": "example"})),
|
||||
"virtual/call": types.CallToolRequest(params=types.CallToolRequestParams(name="mcp_tool_call", arguments={"tool_name": "catalog-example", "arguments": {}})),
|
||||
"proxy/search": types.CallToolRequest(params=types.CallToolRequestParams(name="search_tools", arguments={"query": "example"})),
|
||||
"proxy/schema": types.CallToolRequest(params=types.CallToolRequestParams(name="get_tool_schema", arguments={"tool_id": "28a7a373ebe572627a98e19b5347b405"})),
|
||||
"proxy/call": types.CallToolRequest(params=types.CallToolRequestParams(name="call_tool", arguments={"tool_id": "28a7a373ebe572627a98e19b5347b405", "arguments": {}})),
|
||||
}
|
||||
caller = UserAPIKeyAuth(object_permission={"object_permission_id": "test", "mcp_tool_search_enabled": method.startswith("virtual/")})
|
||||
auth = (caller, None, ["catalog"], None, None, {}, None)
|
||||
proxy_reset = ingress._mcp_proxy_mode.set(method.startswith("proxy/"))
|
||||
try:
|
||||
with (
|
||||
_mcp_upstream(respond),
|
||||
patch.object(operations, "global_mcp_server_manager", manager),
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[target])),
|
||||
patch.object(manager, "get_allowed_mcp_servers", AsyncMock(return_value=["catalog"])),
|
||||
patch.object(ingress, "get_or_extract_auth_context", AsyncMock(return_value=auth)),
|
||||
):
|
||||
for value in ("first", "second", "unvalidated"):
|
||||
credential = ConnectionCredential(
|
||||
kind="connection_access", binding=ConnectionBinding(key_hash="key", server_id="catalog", resource="https://gateway.example/mcp"),
|
||||
client_id="client", token=SecretStr(value or "unused"), jti=value or "unused",
|
||||
exp=int(datetime.now(timezone.utc).timestamp()) + 300,
|
||||
) if value in ("first", "second") else value
|
||||
request = Request({"type": "http", "method": "POST", "path": "/mcp", "headers": [], CONNECTION_SCOPE_KEY: credential})
|
||||
ctx = ServerRequestContext(session=SimpleNamespace(), lifespan_context={}, protocol_version="2025-06-18", method=requests[method].method, request=request)
|
||||
start = len(upstream.requests)
|
||||
async with ingress._legacy_operation_context(ctx, trace=False) as context:
|
||||
result = await operations.GatewayOperations().execute(requests[method], context)
|
||||
sent = upstream.requests[start:]
|
||||
assert sent, result
|
||||
expected = f"Bearer {value}" if value in ("first", "second") else "Bearer saved-vault-token"
|
||||
assert {authorization for _, authorization in sent} == {expected}
|
||||
expected_method = "tools/call" if method in ("virtual/call", "proxy/call") else "tools/list" if method.startswith(("virtual/", "proxy/")) else method
|
||||
assert expected_method in {name for name, _ in sent}
|
||||
if isinstance(result, types.CallToolResult):
|
||||
assert result.is_error is False, result
|
||||
assert result.content
|
||||
if value in ("first", "second"):
|
||||
store.fetch.assert_not_awaited()
|
||||
finally:
|
||||
ingress._mcp_proxy_mode.reset(proxy_reset)
|
||||
|
|
|
|||
|
|
@ -67,7 +67,7 @@ async def test_legacy_adapter_cleans_context_after_cancelled_operation():
|
|||
|
||||
previous_session = server.active_mcp_session_var.get()
|
||||
previous_request = active_mcp_request_ctx_var.get()
|
||||
request = SimpleNamespace(session=object())
|
||||
request = SimpleNamespace(session=object(), request=None)
|
||||
auth = (None, None, None, None, None, None, None)
|
||||
|
||||
async def cancelled_operation():
|
||||
|
|
@ -93,7 +93,7 @@ async def test_legacy_adapter_cleans_context_when_trace_setup_fails():
|
|||
|
||||
previous_session = server.active_mcp_session_var.get()
|
||||
previous_request = active_mcp_request_ctx_var.get()
|
||||
request = SimpleNamespace(session=object())
|
||||
request = SimpleNamespace(session=object(), request=None)
|
||||
|
||||
async def enter_operation():
|
||||
async with server._legacy_operation_context(request, trace=True):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue