fix(mcp): pass keyed OAuth grants through explicit operation contexts

This commit is contained in:
Joshua Valluru 2026-09-21 16:00:31 -07:00
parent 77ffddbf70
commit 23ea0ee428
10 changed files with 267 additions and 69 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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):