Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/server.py: 32%
1063 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2LiteLLM MCP Server Routes
3"""
5# pyright: reportInvalidTypeForm=false, reportArgumentType=false, reportOptionalCall=false
7import asyncio
8import contextlib
9import contextvars
10import hashlib
11import json
12import os
13import time
14import types
15from collections import Counter
16from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence
17from typing import TYPE_CHECKING, Final, NoReturn, Protocol
19import httpx
20from fastapi import FastAPI, HTTPException
21from pydantic import ConfigDict, TypeAdapter, ValidationError
22from starlette.requests import Request as StarletteRequest
23from starlette.responses import JSONResponse
24from starlette.routing import Route
25from starlette.types import Message, Receive, Scope, Send
27from litellm._logging import verbose_logger
28from litellm.constants import (
29 MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH,
30)
31from litellm.llms.custom_httpx.http_handler import (
32 get_async_httpx_client,
33 httpxSpecialProvider,
34)
35from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
36 MCPRequestHandler,
37 _is_mcp_admitted_user_subject,
38)
39from litellm.proxy._experimental.mcp_server.client_allowlist import (
40 MCPClientAllowlist,
41 check_mcp_client_allowed,
42 load_mcp_client_allowlist,
43)
44from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
45 get_request_base_url,
46)
47from litellm.proxy._experimental.mcp_server.exceptions import (
48 MCPUpstreamAuthError,
49)
50from litellm.proxy._experimental.mcp_server.mcp_context import (
51 _mcp_active_toolset_id,
52 _mcp_gateway_initialize_instructions,
53 _mcp_gateway_server_name,
54 _mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode
55 active_mcp_request_ctx_var,
56 get_active_mcp_request_ctx,
57)
58from litellm.proxy._experimental.mcp_server.mcp_debug import (
59 MCP_AUTH_DIAGNOSTICS_SCOPE_KEY,
60 MCPAuthDiagnostics,
61 MCPDebug,
62)
63from litellm.proxy._experimental.mcp_server.oauth_utils import (
64 _redact_mcp_resource_url,
65 get_passthrough_www_authenticate,
66 get_route_relative_request_path,
67 well_known_root_suffix,
68)
69from litellm.proxy._experimental.mcp_server.ui_session_utils import is_ui_session_credential
70from litellm.proxy._experimental.mcp_server.utils import (
71 LITELLM_MCP_SERVER_DESCRIPTION,
72 LITELLM_MCP_SERVER_NAME,
73 LITELLM_MCP_SERVER_VERSION,
74)
75from litellm.proxy._types import (
76 ProxyException,
77 SpecialMCPServerNames,
78 UserAPIKeyAuth,
79)
80from litellm.proxy.auth.ip_address_utils import IPAddressUtils
81from litellm.types.mcp import (
82 MCPAuth,
83 MCPGatewaySession,
84 MCPGatewaySessionGroupCount,
85 MCPGatewaySessionsResponse,
86 MCPGatewaySessionsTerminateResponse,
87 MCPSpecVersion,
88)
89from litellm.types.mcp_server.mcp_server_manager import MCPServer
91if TYPE_CHECKING: 91 ↛ 92line 91 didn't jump to line 92 because the condition on line 91 was never true
92 from mcp.server.session import ServerSession as _McpServerSession
95_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60
96# Upper bound on concurrent stateful sessions a single caller may hold. Each
97# `initialize` creates a session that survives until the idle timeout, so
98# without a cap an authenticated client could spam `initialize` and exhaust
99# memory. The caller's own oldest idle sessions are evicted to make room; if
100# the cap is still hit (every session in flight), the new `initialize` is
101# rejected with 429.
102_MAX_STATEFUL_SESSIONS_PER_OWNER: Final = 100
103# Maximum bytes to peek when sniffing the JSON-RPC method on a POST.
104# An `initialize` envelope is a few hundred bytes; capping the peek
105# prevents an authenticated client from forcing the proxy to buffer an
106# arbitrarily large body just to make a routing decision.
107_MCP_ROUTING_PEEK_MAX_BYTES: Final = 4096
108# ASGI scope keys carrying OTel request state into a stateful MCP message handler.
109_MCP_TRANSPORT_SPAN_SCOPE_KEY: Final = "litellm_otel_transport_span"
110_MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations"
111_MCP_PROTOCOL_VERSION_HEADER: Final = b"mcp-protocol-version"
114def reject_disallowed_mcp_origin(request: StarletteRequest) -> None:
115 from litellm.proxy.proxy_server import origins # noqa: PLC0415 # proxy imports this module during startup
117 if "*" not in origins and any(origin not in origins for origin in request.headers.getlist("origin")): 117 ↛ 118line 117 didn't jump to line 118 because the condition on line 117 was never true
118 raise HTTPException(status_code=403, detail="Invalid Origin header")
121def unsupported_protocol_version(scope: Scope) -> str | None:
122 """Return the unsupported ``MCP-Protocol-Version`` header value, if any.
124 SDK 2's ``StreamableHTTPSessionManager`` routes any version outside
125 ``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
126 bypasses litellm's session/auth model, so the ASGI entry rejects it.
127 """
128 headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
129 values: Final = tuple(
130 raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
131 )
132 for value in values: 132 ↛ 133line 132 didn't jump to line 133 because the loop on line 132 never started
133 if value and value not in HANDSHAKE_PROTOCOL_VERSIONS:
134 return value
135 return None
138# Check if MCP is available
139# "mcp" requires python 3.10 or higher, but several litellm users use python 3.8
140# We're making this conditional import to avoid breaking users who use python 3.8.
141# TODO: Make this a util function for litellm client usage
142MCP_AVAILABLE: bool = True
143try:
144 import weakref
146 from mcp import ReadResourceResult, Resource
147 from mcp.server import Server
148 from mcp.server.session import ServerSession as _McpServerSession
149 from mcp.types import (
150 BlobResourceContents,
151 GetPromptResult,
152 ResourceTemplate,
153 TextResourceContents,
154 )
156 # Robust auth lookup keyed by session_object.
157 _session_obj_auth_storage: "weakref.WeakKeyDictionary[object, MCPAuthenticatedUser]" = weakref.WeakKeyDictionary()
158except ImportError as e:
159 verbose_logger.debug("MCP module not found: %s", e)
160 MCP_AVAILABLE = False
161 # When MCP is not available, we set these to None at module level
162 # All code using these types is inside `if MCP_AVAILABLE:` blocks
163 # so they will never be accessed at runtime
164 BlobResourceContents = None
165 GetPromptResult = None
166 ReadResourceResult = None
167 Resource = None
168 ResourceTemplate = None
169 Server = None
170 TextResourceContents = None
172active_mcp_session_var: Final[contextvars.ContextVar["_McpServerSession | None"]] = contextvars.ContextVar(
173 "active_mcp_session", default=None
174)
177# Global variables to track initialization
178_SESSION_MANAGERS_INITIALIZED = False
179_INITIALIZATION_LOCK: Final = asyncio.Lock()
182def _jsonrpc_text_has_top_level_method(text: str) -> bool:
183 """Whether a (possibly truncated) JSON-RPC envelope has a ``method`` key at
184 the root object's top level.
186 Used to tell a request/notification (carries ``method``) apart from a
187 response (carries ``result``/``error`` and no top-level ``method``). A
188 response payload can itself nest a ``method`` field, so only keys at the
189 root object's depth are inspected rather than searching the whole string.
190 Returns ``True`` only when a top-level ``method`` key is positively found;
191 truncation that hides it yields ``False``.
192 """
193 depth = 0
194 in_string = False
195 escaped = False
196 in_object: Final[list[bool]] = []
197 reading_key = False
198 expect_key = False
199 key_chars: list[str] = []
200 for ch in text:
201 if in_string:
202 if escaped:
203 escaped = False
204 elif ch == "\\":
205 escaped = True
206 elif ch == '"':
207 in_string = False
208 if reading_key and depth == 1 and "".join(key_chars) == "method":
209 return True
210 elif reading_key:
211 key_chars.append(ch)
212 continue
213 if ch == '"':
214 in_string = True
215 reading_key = expect_key and depth >= 1 and in_object[-1]
216 key_chars = []
217 expect_key = False
218 elif ch == "{" or ch == "[":
219 depth += 1
220 in_object.append(ch == "{")
221 expect_key = ch == "{"
222 elif ch == "}" or ch == "]":
223 if in_object:
224 in_object.pop()
225 depth -= 1
226 if depth <= 0:
227 break
228 expect_key = False
229 elif ch == ",":
230 expect_key = bool(in_object) and in_object[-1]
231 elif ch == ":":
232 expect_key = False
233 return False
236def _mcp_meta_trace_carrier(req_ctx: object) -> dict[str, str] | None:
237 """The W3C trace context (``traceparent``/``tracestate``) the MCP client
238 propagated in the request's ``params._meta`` (SEP-414), or ``None``.
240 When present, the MCP span records this propagated context as a span *link*,
241 never the parent — a remote parent would root the span in a trace whose root
242 never reaches the gateway's tracing backend. The span itself nests under the
243 transport span of the request carrying this specific message, so a
244 streamable-HTTP session that multiplexes many messages still does not glue
245 every message under the session's first request;
246 see ``resolve_mcp_span_context``. The client's W3C Baggage is
247 deliberately excluded: it is caller-controlled, and the otel baggage processor
248 stamps allowlisted baggage keys (``litellm.team.id``, ``litellm.metadata.*``,
249 ...) onto the span, so honoring remote baggage would let a client spoof a
250 span's identity attribution.
251 """
252 meta: Final = getattr(req_ctx, "meta", None)
253 extra: Final = meta if isinstance(meta, Mapping) else getattr(meta, "model_extra", None)
254 if not isinstance(extra, Mapping):
255 return None
256 carrier: Final = {key: extra[key] for key in ("traceparent", "tracestate") if isinstance(extra.get(key), str)}
257 return carrier or None
260def _otel_set_mcp_trace_carrier(carrier: dict[str, str] | None) -> object:
261 """Stash ``carrier`` for the otel_v2 MCP span and return a reset token, or
262 ``None`` when otel_v2 is unavailable. Lazily imported so opentelemetry stays an
263 optional dependency."""
264 try:
265 from litellm.integrations.otel.plumbing.context import (
266 set_mcp_message_trace_carrier,
267 )
269 return set_mcp_message_trace_carrier(carrier)
270 except ImportError:
271 return None
274def _otel_reset_mcp_trace_carrier(token: object) -> None:
275 """Clear the per-message trace carrier so it never leaks to the next message on
276 the same session task. Paired with ``_otel_set_mcp_trace_carrier``."""
277 if token is None:
278 return
279 try:
280 from litellm.integrations.otel.plumbing.context import (
281 reset_mcp_message_trace_carrier,
282 )
284 reset_mcp_message_trace_carrier(token)
285 except ImportError:
286 return
289def _otel_publish_transport_span_on_scope(scope: Scope) -> None:
290 """Record this request's tracing span on its own ASGI scope.
292 Resolved on the ASGI request task, where the proxy's server span is anchored,
293 and read back by the MCP message handler through ``req_ctx.request`` — the
294 ``Request`` the transport attaches to each message. A stateful streamable-HTTP
295 session handles every message on the task spawned by its ``initialize`` POST, so
296 the handler's own task cannot see later requests' spans.
298 The scope, not the shared session auth context: a JSON-RPC *response* POST
299 deliberately skips the per-session lock (it can arrive while the tool call that
300 awaits it is still in flight), so a field on that shared object would be
301 overwritten mid-call and the tool call would attribute itself to the response's
302 request. A scope belongs to exactly one request and dies with it, which also
303 keeps a finished span from being retained by an idle session.
305 The live span, not just its context: a failed tool call stamps ``error.*`` on it,
306 which needs a span still open for writes. Lazily imported so opentelemetry stays
307 an optional dependency; a no-op when otel_v2 is unavailable or no request span is
308 anchored."""
309 try:
310 from litellm.integrations.otel.plumbing.context import (
311 request_root_span,
312 )
314 span: Final = request_root_span()
315 except ImportError:
316 return
317 if span is not None: 317 ↛ 318line 317 didn't jump to line 318 because the condition on line 317 was never true
318 scope[_MCP_TRANSPORT_SPAN_SCOPE_KEY] = span
321def _otel_value_from_message_scope(req_ctx: object, key: str) -> object:
322 request: Final = getattr(req_ctx, "request", None)
323 scope: Final = getattr(request, "scope", None)
324 if not isinstance(scope, Mapping):
325 return None
326 return scope.get(key)
329def _otel_transport_span_from_message(req_ctx: object) -> object:
330 """The tracing span of the HTTP request that carried this MCP message."""
331 return _otel_value_from_message_scope(req_ctx, _MCP_TRANSPORT_SPAN_SCOPE_KEY)
334def _otel_set_mcp_transport_span(span: object) -> object:
335 """Publish the current message's transport span, which the otel_v2 MCP span
336 attaches to and a failed tool call stamps its error on. Returns a reset token,
337 or ``None`` when otel_v2 is unavailable."""
338 if span is None:
339 return None
340 try:
341 from litellm.integrations.otel.plumbing.context import (
342 set_mcp_message_transport_span,
343 )
345 return set_mcp_message_transport_span(span)
346 except ImportError:
347 return None
350def _otel_reset_mcp_transport_span(token: object) -> None:
351 """Paired with ``_otel_set_mcp_transport_span``."""
352 if token is None:
353 return
354 try:
355 from litellm.integrations.otel.plumbing.context import (
356 reset_mcp_message_transport_span,
357 )
359 reset_mcp_message_transport_span(token)
360 except ImportError:
361 return
364def _otel_publish_request_destinations_on_scope(scope: Scope) -> None:
365 try:
366 from litellm.integrations.otel.plumbing.context import request_destinations
368 scope[_MCP_DESTINATIONS_SCOPE_KEY] = request_destinations()
369 except ImportError:
370 return
373def _otel_set_mcp_request_destinations(req_ctx: object) -> object:
374 destinations: Final = _otel_value_from_message_scope(req_ctx, _MCP_DESTINATIONS_SCOPE_KEY)
375 if not isinstance(destinations, tuple):
376 return None
377 try:
378 from litellm.integrations.otel.model.destination import OtelDestination
379 from litellm.integrations.otel.plumbing.context import set_request_destinations
381 destination_adapter: Final[TypeAdapter[tuple[OtelDestination, ...]]] = TypeAdapter(
382 tuple[OtelDestination, ...],
383 config=ConfigDict(revalidate_instances="always"),
384 )
385 validated_destinations: Final = destination_adapter.validate_python(destinations, strict=True)
386 return set_request_destinations(validated_destinations)
387 except (ImportError, ValidationError):
388 return None
391def _otel_reset_mcp_request_destinations(token: object) -> None:
392 if token is None:
393 return
394 try:
395 from litellm.integrations.otel.plumbing.context import reset_request_destinations
397 reset_request_destinations(token)
398 except ImportError:
399 return
402def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
403 """Map a ``ProxyException`` to an ``HTTPException`` that preserves its real
404 status code and headers.
406 ``user_api_key_auth`` raises ``ProxyException`` (not ``HTTPException``) on
407 auth failures. The MCP ASGI handlers re-raise ``HTTPException`` to keep the
408 status and any ``WWW-Authenticate`` challenge, but a ``ProxyException`` would
409 otherwise fall through to their generic handler and be flattened to a 500 —
410 dropping the 401 + challenge an OAuth client needs to re-authenticate, so the
411 tool call surfaces as a cancelled/terminated session instead.
412 """
413 try:
414 status_code = int(exc.code)
415 except (TypeError, ValueError):
416 status_code = 500
417 return HTTPException(
418 status_code=status_code,
419 detail=exc.message,
420 headers=exc.headers or None,
421 )
424if MCP_AVAILABLE: 424 ↛ 2730line 424 didn't jump to line 2730 because the condition on line 424 was always true
425 __all__ = (
426 "_MCP_CREDENTIAL_REQUEST_FIELDS",
427 "BlobResourceContents",
428 "ListMCPToolsRestAPIResponseObject",
429 "ResourceTemplate",
430 "TextResourceContents",
431 "_McpDeniedDetail",
432 "_aggregate_server_key",
433 "_build_virtual_call_logging_obj",
434 "_check_byok_credential",
435 "_client_has_passthrough_authorization",
436 "_client_has_per_server_auth_header",
437 "_dispatch_virtual_mcp_tool",
438 "_fire_mcp_tool_call_logging",
439 "_get_allowed_mcp_servers",
440 "_get_allowed_mcp_servers_from_mcp_server_names",
441 "_get_byok_credential",
442 "_get_prompts_from_mcp_servers",
443 "_get_resource_templates_from_mcp_servers",
444 "_get_resources_from_mcp_servers",
445 "_get_standard_logging_mcp_tool_call",
446 "_get_tools_from_mcp_servers",
447 "_get_user_oauth_extra_headers_from_db",
448 "_handle_local_mcp_tool",
449 "_handle_managed_mcp_tool",
450 "_http_detail_message",
451 "_invalidate_byok_cred_cache",
452 "_list_mcp_prompts",
453 "_list_mcp_resource_templates",
454 "_list_mcp_resources",
455 "_list_mcp_tools",
456 "_list_tools_before_first_call",
457 "_mcp_session_id_from_headers",
458 "_merge_gateway_initialize_instructions",
459 "_prefetch_oauth_creds_for_user",
460 "_prepare_mcp_server_headers",
461 "_raise_if_initialize_grants_no_mcp_servers",
462 "_redact_mcp_resource_url",
463 "_resolve_display_name_to_original",
464 "_run_post_mcp_call_guardrails",
465 "_server_answers_to",
466 "_tool_name_matches",
467 "apply_tool_overrides",
468 "call_mcp_tool",
469 "execute_mcp_tool",
470 "filter_tools_by_allowed_tools",
471 "filter_tools_by_key_team_permissions",
472 "fire_mcp_tool_call_failure_logging",
473 "global_mcp_server_manager",
474 "mcp_get_prompt",
475 "mcp_read_resource",
476 "raise_denied_scoped_mcp_access",
477 )
478 from mcp.server import Server
480 # Import auth context variables and middleware
481 from mcp.server.auth.middleware.auth_context import (
482 AuthContextMiddleware,
483 auth_context_var,
484 )
485 from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
486 from mcp.server.auth.provider import AccessToken
487 from mcp.server.context import ServerRequestContext
488 from mcp.server.lowlevel.server import NotificationOptions
489 from mcp.server.models import InitializationOptions
490 from mcp.shared.exceptions import MCPError
491 from mcp.types import (
492 CallToolRequest,
493 GetPromptRequest,
494 ListPromptsRequest,
495 ListResourcesRequest,
496 ListResourceTemplatesRequest,
497 ListToolsRequest,
498 ReadResourceRequest,
499 )
501 from litellm.proxy._experimental.mcp_server import operations
502 from litellm.proxy._experimental.mcp_server.contracts import OperationContext
503 from litellm.proxy._experimental.mcp_server.operations import (
504 _invalidate_byok_cred_cache,
505 _mcp_session_id_from_headers,
506 )
508 try:
509 from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
510 except ImportError:
511 StreamableHTTPSessionManager = None
512 from mcp.types import (
513 INVALID_REQUEST,
514 CallToolRequestParams,
515 CallToolResult,
516 GetPromptRequestParams,
517 Implementation,
518 InitializeRequest,
519 ListPromptsResult,
520 ListResourcesResult,
521 ListResourceTemplatesResult,
522 ListToolsResult,
523 PaginatedRequestParams,
524 ReadResourceRequestParams,
525 )
526 from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
528 from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
529 MCPAuthenticatedUser,
530 )
531 from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
532 MCPServerManager,
533 global_mcp_server_manager,
534 )
536 ######################################################
537 ############ MCP Tools List REST API Response Object #
538 # Defined here because we don't want to add `mcp` as a
539 # required dependency for `litellm` pip package
540 ######################################################
541 from litellm.proxy._experimental.mcp_server.operations import (
542 ListMCPToolsRestAPIResponseObject,
543 )
544 from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport
546 def _gateway_create_initialization_options(
547 self,
548 notification_options: NotificationOptions | None = None,
549 experimental_capabilities: dict[str, dict[str, object]] | None = None,
550 extensions: dict[str, dict[str, object]] | None = None,
551 ) -> InitializationOptions:
552 base_options: Final = Server.create_initialization_options(
553 self,
554 notification_options=notification_options,
555 experimental_capabilities=experimental_capabilities or {},
556 extensions=extensions,
557 )
558 opts: Final = (
559 base_options.model_copy(
560 update={ # mutable-ok: Pydantic update payload
561 "capabilities": base_options.capabilities.model_copy(
562 update={"prompts": None, "resources": None} # mutable-ok: Pydantic update payload
563 )
564 }
565 )
566 if _mcp_proxy_mode.get()
567 else base_options
568 )
569 updates: Final[dict[str, str]] = {}
570 merged: Final = _mcp_gateway_initialize_instructions.get()
571 if merged is not None:
572 updates["instructions"] = merged
573 scoped_server_name: Final = _mcp_gateway_server_name.get()
574 if scoped_server_name is not None:
575 updates["server_name"] = scoped_server_name
576 return opts.model_copy(update=updates) if updates else opts
578 ########################################################
579 ############ Initialize the MCP Server #################
580 ########################################################
581 server: Final[Server] = Server(
582 name=LITELLM_MCP_SERVER_NAME,
583 version=LITELLM_MCP_SERVER_VERSION,
584 )
585 server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
586 sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
588 # Create session managers
589 session_manager_stateless: Final = StreamableHTTPSessionManager(
590 app=server,
591 event_store=None,
592 json_response=False, # enables SSE streaming
593 stateless=True,
594 )
596 session_manager_stateful: Final = StreamableHTTPSessionManager(
597 app=server,
598 event_store=None, # TODO: Add EventStore for reconnection/event replay if needed
599 json_response=False, # enables SSE streaming
600 stateless=False,
601 )
602 _stateful_session_auth_contexts: Final[dict[str, MCPAuthenticatedUser]] = {}
603 _stateful_session_auth_context_last_seen: Final[dict[str, float]] = {}
604 # Maps session_id -> owner identifier (hashed API key/token) so we can
605 # reject requests that supply a session_id created by a different caller.
606 # Without this, a leaked mcp-session-id could be driven (or terminated)
607 # by any other authenticated proxy user.
608 _stateful_session_owners: Final[dict[str, str]] = {}
609 # Per-session lock that serializes ``handle_request`` for the same
610 # mcp-session-id. The stored ``MCPAuthenticatedUser`` is mutated in place
611 # by ``_update_auth_context`` each request; without this lock, two
612 # concurrent requests on the same session would clobber each other's
613 # auth headers / mcp_servers / oauth state while in-flight callbacks are
614 # still reading the shared object.
615 _stateful_session_locks: Final[dict[str, asyncio.Lock]] = {}
616 _stateful_session_active_request_counts: Final[dict[str, int]] = {}
617 _stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown
618 _admin_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: admin-closed id -> last replay
620 class _TerminableTransport(Protocol):
621 async def terminate(self) -> None: ... 621 ↛ exitline 621 didn't return from function 'terminate' because
623 class _TransportRegistry(Protocol):
624 def __contains__(self, session_id: object, /) -> bool: ... 624 ↛ exitline 624 didn't return from function '__contains__' because
626 def pop(self, session_id: str, default: None, /) -> "_TerminableTransport | None": ... 626 ↛ exitline 626 didn't return from function 'pop' because
628 def _stateful_server_instances() -> _TransportRegistry:
629 return getattr(session_manager_stateful, "_server_instances", {})
631 def _remove_stateful_session_tracking(session_id: str) -> None:
632 _stateful_session_auth_contexts.pop(session_id, None)
633 _stateful_session_auth_context_last_seen.pop(session_id, None)
634 _stateful_session_owners.pop(session_id, None)
635 _stateful_session_locks.pop(session_id, None)
636 _stateful_session_active_request_counts.pop(session_id, None)
637 _stateful_session_client_info.pop(session_id, None)
639 # Keep this alias so existing references to session_manager still work
640 session_manager: Final = session_manager_stateless
642 # Context managers for proper lifecycle management
643 _session_manager_cm = None
644 _session_manager_stateful_cm = None
645 _stateful_auth_context_cleanup_task: asyncio.Task | None = None
647 async def _purge_expired_stateful_session_auth_contexts(
648 now: float | None = None,
649 ) -> None:
650 """Terminate expired stateful sessions and drop their auth contexts."""
651 now = time.monotonic() if now is None else now
652 server_instances: Final = _stateful_server_instances()
653 expired_session_ids: Final[list[str]] = []
654 for session_id, last_seen in _stateful_session_auth_context_last_seen.items(): 654 ↛ 655line 654 didn't jump to line 655 because the loop on line 654 never started
655 if _stateful_session_active_request_counts.get(session_id, 0) > 0:
656 continue
657 if now - last_seen >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS or session_id not in server_instances:
658 expired_session_ids.append(session_id)
660 for session_id in expired_session_ids: 660 ↛ 666line 660 didn't jump to line 666 because the loop on line 660 never started
661 # Re-check the active-request count immediately before tearing
662 # the session down. ``await transport.terminate()`` yields to
663 # the event loop, so a request that started after the first
664 # collection pass could otherwise observe its transport being
665 # ripped out from under it mid-flight.
666 if _stateful_session_active_request_counts.get(session_id, 0) > 0:
667 continue
668 # Pop transport + terminate BEFORE removing owner/auth tracking.
669 # Reversing the order avoids a window where ``_stateful_session_owners``
670 # is empty but ``server_instances`` still serves the session — a
671 # concurrent request in that window would observe ``expected_owner
672 # is None`` and bypass the owner-binding check.
673 transport = server_instances.pop(session_id, None)
674 if transport is not None:
675 await transport.terminate()
676 _remove_stateful_session_tracking(session_id)
678 for session_id in list(_stateful_session_auth_context_last_seen): 678 ↛ 679line 678 didn't jump to line 679 because the loop on line 678 never started
679 if session_id not in _stateful_session_auth_contexts:
680 _remove_stateful_session_tracking(session_id)
681 _forget_expired_admin_terminated_session_ids(now)
683 async def _enforce_stateful_session_cap_for_owner(owner: str) -> bool:
684 """
685 Bound the number of concurrent stateful sessions a single caller holds
686 before routing a new ``initialize`` to the stateful manager.
688 Evicts the caller's *own* oldest idle sessions (no in-flight requests)
689 to make room, so a busy-but-legitimate client keeps its newest sessions
690 and other callers are never affected. Returns ``True`` if the new
691 session may proceed, or ``False`` when the caller is already at the cap
692 with every session in flight (the new ``initialize`` should be rejected).
693 """
694 server_instances: Final = _stateful_server_instances()
696 def _owned_live_session_ids() -> list[str]:
697 return [
698 session_id
699 for session_id, session_owner in _stateful_session_owners.items()
700 if session_owner == owner and session_id in server_instances
701 ]
703 owned: Final = _owned_live_session_ids()
704 if len(owned) < _MAX_STATEFUL_SESSIONS_PER_OWNER:
705 return True
707 for session_id in sorted(
708 owned,
709 key=lambda sid: _stateful_session_auth_context_last_seen.get(sid, 0.0),
710 ):
711 if len(_owned_live_session_ids()) < _MAX_STATEFUL_SESSIONS_PER_OWNER:
712 break
713 if _stateful_session_active_request_counts.get(session_id, 0) > 0:
714 continue
715 transport = server_instances.pop(session_id, None)
716 if transport is not None:
717 await transport.terminate()
718 _remove_stateful_session_tracking(session_id)
720 return len(_owned_live_session_ids()) < _MAX_STATEFUL_SESSIONS_PER_OWNER
722 async def _cleanup_expired_stateful_session_auth_contexts() -> None:
723 while True:
724 await asyncio.sleep(_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS)
725 try:
726 await _purge_expired_stateful_session_auth_contexts()
727 except Exception as e:
728 verbose_logger.exception("Error cleaning up expired MCP stateful sessions: %s", e)
730 async def initialize_session_managers():
731 """Initialize the session managers. Can be called from main app lifespan."""
732 global \
733 _SESSION_MANAGERS_INITIALIZED, \
734 _session_manager_cm, \
735 _session_manager_stateful_cm, \
736 _stateful_auth_context_cleanup_task
738 # Use async lock to prevent concurrent initialization
739 async with _INITIALIZATION_LOCK:
740 if _SESSION_MANAGERS_INITIALIZED: 740 ↛ 741line 740 didn't jump to line 741 because the condition on line 740 was never true
741 return
743 verbose_logger.info("Initializing MCP session managers...")
745 # Start the session managers with context managers
746 _session_manager_cm = session_manager_stateless.run()
747 _session_manager_stateful_cm = session_manager_stateful.run()
749 # Enter the context managers
750 await _session_manager_cm.__aenter__()
751 await _session_manager_stateful_cm.__aenter__()
752 _stateful_auth_context_cleanup_task = asyncio.create_task(_cleanup_expired_stateful_session_auth_contexts())
754 _SESSION_MANAGERS_INITIALIZED = True
755 verbose_logger.info("MCP Server started with StreamableHTTP and SSE session managers!")
757 async def shutdown_session_managers():
758 """Shutdown the session managers."""
759 global \
760 _SESSION_MANAGERS_INITIALIZED, \
761 _session_manager_cm, \
762 _session_manager_stateful_cm, \
763 _stateful_auth_context_cleanup_task
765 if _SESSION_MANAGERS_INITIALIZED:
766 verbose_logger.info("Shutting down MCP session managers...")
768 try:
769 if _stateful_auth_context_cleanup_task:
770 _stateful_auth_context_cleanup_task.cancel()
771 with contextlib.suppress(asyncio.CancelledError):
772 await _stateful_auth_context_cleanup_task
773 if _session_manager_stateful_cm:
774 await _session_manager_stateful_cm.__aexit__(None, None, None)
775 if _session_manager_cm:
776 await _session_manager_cm.__aexit__(None, None, None)
777 except Exception as e:
778 verbose_logger.exception("Error during session manager shutdown: %s", e)
780 _session_manager_cm = None
781 _session_manager_stateful_cm = None
782 _stateful_auth_context_cleanup_task = None
783 _SESSION_MANAGERS_INITIALIZED = False
785 @contextlib.asynccontextmanager
786 async def lifespan(app) -> AsyncIterator[None]:
787 """Application lifespan context manager."""
788 await initialize_session_managers()
789 try:
790 yield
791 finally:
792 await shutdown_session_managers()
794 ########################################################
795 ############### MCP Server Routes #######################
796 ########################################################
798 @contextlib.asynccontextmanager
799 async def _legacy_operation_context(ctx: ServerRequestContext, *, trace: bool) -> AsyncGenerator[OperationContext]:
800 with contextlib.ExitStack() as cleanup:
801 cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx))
802 cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session))
803 if trace:
804 cleanup.callback(
805 _otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx))
806 )
807 cleanup.callback(
808 _otel_reset_mcp_transport_span, _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx))
809 )
810 cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx))
811 (
812 auth,
813 token,
814 servers,
815 server_headers,
816 oauth_headers,
817 headers,
818 client_ip,
819 ) = await get_or_extract_auth_context()
820 yield operations.prepare_context(
821 auth, token, servers, server_headers, oauth_headers, headers, client_ip, _mcp_proxy_mode.get()
822 )
824 async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
825 try:
826 async with _legacy_operation_context(ctx, trace=True) as context:
827 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
828 ListToolsRequest(params=params), context
829 )
830 except MCPError:
831 raise
832 except HTTPException as exc:
833 raise MCPError(code=INVALID_REQUEST, message=operations._http_detail_message(exc.detail)) from exc
834 except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
835 verbose_logger.exception("Error in list_tools endpoint: %s", exc)
836 return ListToolsResult(tools=[])
838 def _capture_host_progress_callback(ctx: ServerRequestContext) -> Callable | None:
839 """Return a progress-forwarding callback bound to the host MCP session.
841 Returns ``None`` when the host did not supply a progress token.
842 """
843 host_ctx: Final = ctx
845 if not (host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta):
846 return None
847 host_token: Final = host_ctx.meta.get("progress_token")
848 if host_token is None or not (hasattr(host_ctx, "session") and host_ctx.session):
849 return None
850 host_session: Final = host_ctx.session
852 async def forward_progress(progress: float, total: float | None):
853 """Forward progress notifications from external MCP to Host"""
854 try:
855 await host_session.send_progress_notification(
856 progress_token=host_token,
857 progress=progress,
858 total=total,
859 )
860 verbose_logger.debug("Forwarded progress %s/%s to Host", progress, total)
861 except Exception as e:
862 verbose_logger.error("Failed to forward progress to Host: %s", e)
864 verbose_logger.debug("Host progressToken captured: %s...", str(host_token)[:8])
865 return forward_progress
867 def _reject_mcp_proxy_operation() -> NoReturn:
868 from mcp.shared.exceptions import MCPError
869 from mcp.types import METHOD_NOT_FOUND
871 raise MCPError(code=METHOD_NOT_FOUND, message="Operation unavailable on /mcp/proxy")
873 from litellm.proxy._experimental.mcp_server.operations import (
874 _build_virtual_call_logging_obj,
875 _dispatch_virtual_mcp_tool,
876 )
878 async def mcp_server_tool_call(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
879 async with _legacy_operation_context(ctx, trace=True) as context:
880 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
881 CallToolRequest(params=params), context
882 )
884 async def list_prompts(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListPromptsResult:
885 if _mcp_proxy_mode.get():
886 _reject_mcp_proxy_operation()
887 try:
888 async with _legacy_operation_context(ctx, trace=False) as context:
889 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
890 ListPromptsRequest(params=params), context
891 )
892 except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
893 verbose_logger.exception("Error in list_prompts endpoint: %s", exc)
894 return ListPromptsResult(prompts=[])
896 async def get_prompt(ctx: ServerRequestContext, params: GetPromptRequestParams) -> GetPromptResult:
897 if _mcp_proxy_mode.get():
898 _reject_mcp_proxy_operation()
899 async with _legacy_operation_context(ctx, trace=False) as context:
900 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
901 GetPromptRequest(params=params), context
902 )
904 async def list_resources(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListResourcesResult:
905 if _mcp_proxy_mode.get():
906 _reject_mcp_proxy_operation()
907 try:
908 async with _legacy_operation_context(ctx, trace=False) as context:
909 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
910 ListResourcesRequest(params=params), context
911 )
912 except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
913 verbose_logger.exception("Error in list_resources endpoint: %s", exc)
914 return ListResourcesResult(resources=[])
916 async def list_resource_templates(
917 ctx: ServerRequestContext, params: PaginatedRequestParams
918 ) -> ListResourceTemplatesResult:
919 if _mcp_proxy_mode.get():
920 _reject_mcp_proxy_operation()
921 try:
922 async with _legacy_operation_context(ctx, trace=False) as context:
923 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
924 ListResourceTemplatesRequest(params=params), context
925 )
926 except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
927 verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc)
928 return ListResourceTemplatesResult(resource_templates=[])
930 async def read_resource(ctx: ServerRequestContext, params: ReadResourceRequestParams) -> ReadResourceResult:
931 if _mcp_proxy_mode.get():
932 _reject_mcp_proxy_operation()
933 async with _legacy_operation_context(ctx, trace=False) as context:
934 return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
935 ReadResourceRequest(params=params), context
936 )
938 server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools)
939 server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call)
940 server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts)
941 server.add_request_handler("prompts/get", GetPromptRequestParams, get_prompt)
942 server.add_request_handler("resources/list", PaginatedRequestParams, list_resources)
943 server.add_request_handler("resources/templates/list", PaginatedRequestParams, list_resource_templates)
944 server.add_request_handler("resources/read", ReadResourceRequestParams, read_resource)
946 ########################################################
947 ############ End of MCP Server Routes ##################
948 ########################################################
950 ########################################################
951 ############ Helper Functions ##########################
952 ########################################################
954 from litellm.proxy._experimental.mcp_server.operations import (
955 _client_has_passthrough_authorization,
956 _client_has_per_server_auth_header,
957 _get_allowed_mcp_servers,
958 _get_allowed_mcp_servers_from_mcp_server_names,
959 _get_user_oauth_extra_headers_from_db,
960 _http_detail_message,
961 _McpDeniedDetail,
962 _merge_gateway_initialize_instructions,
963 _prefetch_oauth_creds_for_user,
964 _prepare_mcp_server_headers,
965 _raise_if_initialize_grants_no_mcp_servers,
966 _server_answers_to,
967 _tool_name_matches,
968 apply_tool_overrides,
969 filter_tools_by_allowed_tools,
970 raise_denied_scoped_mcp_access,
971 )
973 @contextlib.asynccontextmanager
974 async def _gateway_initialize_instructions_request_scope(
975 user_api_key_auth: UserAPIKeyAuth | None,
976 mcp_servers: list[str] | None,
977 client_ip: str | None,
978 scoped_server_endpoint: bool = False,
979 is_initialize: bool = False,
980 ) -> AsyncIterator[None]:
981 allowed: Final = await operations._get_allowed_mcp_servers(
982 user_api_key_auth=user_api_key_auth,
983 mcp_servers=mcp_servers,
984 client_ip=client_ip,
985 )
986 if is_initialize: 986 ↛ 987line 986 didn't jump to line 987 because the condition on line 986 was never true
987 await operations._raise_if_initialize_grants_no_mcp_servers(
988 allowed, user_api_key_auth, mcp_servers, client_ip
989 )
990 if allowed:
991 # return_exceptions=True: a per-server probe failure (incl. CancelledError
992 # bubbled from anyio task group teardown on connection refused) must not
993 # cancel sibling probes or 500 the gateway initialize request.
994 await asyncio.gather(
995 *[
996 operations.global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(s)
997 for s in allowed
998 if s is not None
999 ],
1000 return_exceptions=True,
1001 )
1002 merged: Final = operations._merge_gateway_initialize_instructions(allowed_mcp_servers=allowed)
1003 scoped_server_name = None
1004 if scoped_server_endpoint and len(allowed) == 1: 1004 ↛ 1005line 1004 didn't jump to line 1005 because the condition on line 1004 was never true
1005 scoped_server: Final = allowed[0]
1006 scoped_server_name = (
1007 scoped_server.alias or scoped_server.server_name or scoped_server.name or scoped_server.server_id
1008 )
1009 instructions_token: Final = _mcp_gateway_initialize_instructions.set(merged)
1010 server_name_token: Final = _mcp_gateway_server_name.set(scoped_server_name)
1011 try:
1012 yield
1013 finally:
1014 _mcp_gateway_initialize_instructions.reset(instructions_token)
1015 _mcp_gateway_server_name.reset(server_name_token)
1017 from litellm.proxy._experimental.mcp_server.operations import (
1018 _MCP_CREDENTIAL_REQUEST_FIELDS,
1019 _aggregate_server_key,
1020 _check_byok_credential,
1021 _fire_mcp_tool_call_logging,
1022 _get_byok_credential,
1023 _get_prompts_from_mcp_servers,
1024 _get_resource_templates_from_mcp_servers,
1025 _get_resources_from_mcp_servers,
1026 _get_standard_logging_mcp_tool_call,
1027 _get_tools_from_mcp_servers,
1028 _handle_local_mcp_tool,
1029 _handle_managed_mcp_tool,
1030 _list_mcp_prompts,
1031 _list_mcp_resource_templates,
1032 _list_mcp_resources,
1033 _list_mcp_tools,
1034 _list_tools_before_first_call,
1035 _resolve_display_name_to_original,
1036 _run_post_mcp_call_guardrails,
1037 call_mcp_tool,
1038 execute_mcp_tool,
1039 filter_tools_by_key_team_permissions,
1040 fire_mcp_tool_call_failure_logging,
1041 mcp_get_prompt,
1042 mcp_read_resource,
1043 )
1045 def _get_mcp_servers_in_path(path: str) -> list[str] | None:
1046 """
1047 Get the MCP servers from the path
1048 """
1049 import re
1051 if path.rstrip("/") in ("/mcp/sse", "/mcp/sse/messages"): 1051 ↛ 1052line 1051 didn't jump to line 1052 because the condition on line 1051 was never true
1052 return None
1053 mcp_servers_from_path: list[str] | None = None
1054 segments: Final = [s for s in path.split("/") if s]
1055 if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp": 1055 ↛ 1056line 1055 didn't jump to line 1056 because the condition on line 1055 was never true
1056 return [segments[0]]
1058 # Match /mcp/<servers_and_maybe_path>
1059 # Where servers can be comma-separated list of server names
1060 # Server names can contain slashes (e.g., "custom_solutions/user_123")
1061 mcp_path_match: Final = re.match(r"^/mcp/([^?#]+)(?:\?.*)?(?:#.*)?$", path)
1062 if mcp_path_match: 1062 ↛ 1063line 1062 didn't jump to line 1063 because the condition on line 1062 was never true
1063 servers_and_path: Final = mcp_path_match.group(1)
1065 if servers_and_path:
1066 # Check if it contains commas (comma-separated servers)
1067 if "," in servers_and_path:
1068 # For comma-separated, look for a path at the end
1069 # Common patterns: /tools, /chat/completions, etc.
1070 path_match: Final = re.search(r"/([^/,]+(?:/[^/,]+)*)$", servers_and_path)
1071 if path_match:
1072 # Path found at the end, remove it from servers
1073 path_part: Final = "/" + path_match.group(1)
1074 servers_part: Final = servers_and_path[: -len(path_part)]
1075 mcp_servers_from_path = [s.strip() for s in servers_part.split(",") if s.strip()]
1076 else:
1077 # No path, just comma-separated servers
1078 mcp_servers_from_path = [s.strip() for s in servers_and_path.split(",") if s.strip()]
1079 else:
1080 # Single server case - use regex approach for server/path separation
1081 # This handles cases like "custom_solutions/user_123/chat/completions"
1082 # where we want to extract "custom_solutions/user_123" as the server name
1083 single_server_match: Final = re.match(r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", servers_and_path)
1084 if single_server_match:
1085 server_name: Final = single_server_match.group(1)
1086 mcp_servers_from_path = [server_name]
1087 else:
1088 mcp_servers_from_path = [servers_and_path]
1089 return mcp_servers_from_path
1091 def _load_mcp_client_allowlist() -> MCPClientAllowlist | None:
1092 from litellm.proxy.proxy_server import general_settings
1094 return load_mcp_client_allowlist(general_settings)
1096 def reject_disallowed_mcp_client(headers: Mapping[str, str], user_api_key_auth: UserAPIKeyAuth | None) -> None:
1097 """Gate every MCP tool surface on ``mcp_allowed_clients``; the dashboard's own session is not a client app."""
1098 if user_api_key_auth is not None and is_ui_session_credential(user_api_key_auth): 1098 ↛ 1099line 1098 didn't jump to line 1099 because the condition on line 1098 was never true
1099 return
1100 rejection: Final = check_mcp_client_allowed(
1101 allowlist=_load_mcp_client_allowlist(),
1102 jwt_claims=user_api_key_auth.jwt_claims if user_api_key_auth is not None else None,
1103 headers=headers,
1104 )
1105 if rejection is None: 1105 ↛ 1107line 1105 didn't jump to line 1107 because the condition on line 1105 was always true
1106 return
1107 verbose_logger.warning("Rejected MCP request from a disallowed client application: %s", rejection.details)
1108 raise HTTPException(status_code=403, detail=rejection.response_body)
1110 async def extract_mcp_auth_context(scope, path):
1111 """
1112 Extracts mcp_servers from the path and processes the MCP request for auth context.
1113 Returns: (user_api_key_auth, mcp_auth_header, mcp_servers, mcp_server_auth_headers)
1114 """
1115 mcp_servers_from_path: Final = _get_mcp_servers_in_path(path)
1116 if mcp_servers_from_path is not None: 1116 ↛ 1117line 1116 didn't jump to line 1117 because the condition on line 1116 was never true
1117 (
1118 user_api_key_auth,
1119 mcp_auth_header,
1120 _,
1121 mcp_server_auth_headers,
1122 oauth2_headers,
1123 raw_headers,
1124 ) = await MCPRequestHandler.process_mcp_request(scope)
1125 mcp_servers = mcp_servers_from_path
1126 else:
1127 (
1128 user_api_key_auth,
1129 mcp_auth_header,
1130 mcp_servers,
1131 mcp_server_auth_headers,
1132 oauth2_headers,
1133 raw_headers,
1134 ) = await MCPRequestHandler.process_mcp_request(scope)
1135 return (
1136 user_api_key_auth,
1137 mcp_auth_header,
1138 mcp_servers,
1139 mcp_server_auth_headers,
1140 oauth2_headers,
1141 raw_headers,
1142 )
1144 def _get_session_id_from_scope(scope: Scope) -> str | None:
1145 """
1146 Extract mcp-session-id from ASGI scope headers.
1147 Returns None if not present.
1148 """
1149 scope_headers: Final[Sequence[tuple[bytes | str, bytes | str]]] = scope.get("headers", [])
1150 for header_name, header_value in scope_headers:
1151 name = header_name if isinstance(header_name, bytes) else header_name.encode()
1152 if name.lower() == b"mcp-session-id": 1152 ↛ 1153line 1152 didn't jump to line 1153 because the condition on line 1152 was never true
1153 return header_value.decode() if isinstance(header_value, bytes) else str(header_value)
1154 return None
1156 def _owner_fingerprint_for(
1157 user_api_key_auth: UserAPIKeyAuth | None,
1158 oauth2_headers: dict[str, str] | None = None,
1159 client_ip: str | None = None,
1160 ) -> str:
1161 """
1162 Stable, non-reversible identifier for the caller used to bind an
1163 mcp-session-id to its creator. Hash the resolved credential before
1164 using it so custom key formats are never stored in cleartext.
1166 For OAuth2 passthrough (``UserAPIKeyAuth()`` with no key/user_id),
1167 the caller's identity is the upstream OAuth bearer; hash it so two
1168 OAuth callers with different tokens don't both fingerprint to
1169 ``anonymous`` and end up sharing a session.
1171 When no caller-identifying credentials are available at all
1172 (e.g. proxy running without master key, or an unauthenticated
1173 passthrough path), fall back to the client IP so two unrelated
1174 anonymous callers from different sources do not collapse to a
1175 single ``anonymous`` owner and end up able to drive each other's
1176 stateful sessions. Note: when even client IP is unavailable
1177 (exotic deployments without trusted X-Forwarded-For and direct
1178 socket info), the fingerprint degrades to the ``anonymous``
1179 sentinel and cannot meaningfully protect against another
1180 unauthenticated caller who learns the session id — owner-binding
1181 is best-effort in that mode.
1182 """
1184 def _bytes_for_hash(value: object) -> bytes | None:
1185 """Only hash str/bytes secrets; skip mocks and other unexpected types."""
1186 if value is None:
1187 return None
1188 if isinstance(value, (bytes, bytearray)):
1189 return bytes(value)
1190 if isinstance(value, str):
1191 return value.encode("utf-8")
1192 return None
1194 if user_api_key_auth is not None:
1195 key_material: Final = _bytes_for_hash(getattr(user_api_key_auth, "api_key", None))
1196 if key_material:
1197 api_key_hash: Final = hashlib.sha256(key_material).hexdigest()
1198 return f"key:{api_key_hash}"
1199 uid_material: Final = _bytes_for_hash(getattr(user_api_key_auth, "user_id", None))
1200 if uid_material:
1201 user_id_hash: Final = hashlib.sha256(uid_material).hexdigest()
1202 return f"user:{user_id_hash}"
1203 if oauth2_headers:
1204 authz: Final = oauth2_headers.get("Authorization") or oauth2_headers.get("authorization")
1205 authz_bytes: Final = _bytes_for_hash(authz)
1206 if authz_bytes:
1207 return f"oauth:{hashlib.sha256(authz_bytes).hexdigest()}"
1208 if client_ip and isinstance(client_ip, str):
1209 return f"ip:{hashlib.sha256(client_ip.encode('utf-8')).hexdigest()}"
1210 return "anonymous"
1212 def _is_initialize_request(body: bytes) -> bool:
1213 """
1214 Check if the request body is a JSON-RPC initialize method.
1215 Returns True if method is "initialize", False otherwise or on parse error.
1216 """
1217 if not body: 1217 ↛ 1219line 1217 didn't jump to line 1219 because the condition on line 1217 was always true
1218 return False
1219 try:
1220 data: Final = json.loads(body)
1221 return isinstance(data, dict) and data.get("method") == "initialize"
1222 except (json.JSONDecodeError, UnicodeDecodeError, TypeError):
1223 return False
1225 def _extract_initialize_client_info(body: bytes) -> Implementation | None:
1226 try:
1227 return InitializeRequest.model_validate_json(body, by_name=False).params.client_info
1228 except ValidationError:
1229 return None
1231 def _group_session_counts(
1232 sessions: Sequence[MCPGatewaySession],
1233 label_for: Callable[[MCPGatewaySession], str | None],
1234 ) -> tuple[MCPGatewaySessionGroupCount, ...]:
1235 counts: Final = types.MappingProxyType(Counter(label_for(session) for session in sessions))
1236 return tuple(
1237 sorted(
1238 (MCPGatewaySessionGroupCount(label=label, count=count) for label, count in counts.items()),
1239 key=lambda group: (-group.count, group.label is None, group.label or ""),
1240 )
1241 )
1243 def _gateway_session_for(session_id: str, auth_user: MCPAuthenticatedUser, now: float) -> MCPGatewaySession:
1244 client_info: Final = _stateful_session_client_info.get(session_id)
1245 key_auth: Final = auth_user.user_api_key_auth
1246 return MCPGatewaySession(
1247 session_id_prefix=session_id[:MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH],
1248 client_name=client_info.name if client_info is not None else None,
1249 client_version=client_info.version if client_info is not None else None,
1250 user_id=key_auth.user_id if key_auth is not None else None,
1251 user_email=key_auth.user_email if key_auth is not None else None,
1252 key_alias=key_auth.key_alias if key_auth is not None else None,
1253 team_id=key_auth.team_id if key_auth is not None else None,
1254 team_alias=key_auth.team_alias if key_auth is not None else None,
1255 client_ip=auth_user.client_ip,
1256 idle_seconds=max(0.0, now - _stateful_session_auth_context_last_seen.get(session_id, now)),
1257 in_flight_requests=_stateful_session_active_request_counts.get(session_id, 0),
1258 )
1260 def get_mcp_gateway_sessions_report(now: float | None = None) -> MCPGatewaySessionsResponse:
1261 """Live stateful Streamable HTTP sessions held by this worker process.
1263 Only sessions whose transport is still registered with the stateful
1264 session manager are reported; SSE and stateless requests hold no
1265 session and are never counted.
1266 """
1267 report_time: Final = time.monotonic() if now is None else now
1268 live_session_ids: Final = frozenset(_stateful_server_instances())
1269 sessions: Final = tuple(
1270 _gateway_session_for(session_id, auth_user, report_time)
1271 for session_id, auth_user in tuple(_stateful_session_auth_contexts.items())
1272 if session_id in live_session_ids
1273 )
1274 return MCPGatewaySessionsResponse(
1275 worker_pid=os.getpid(),
1276 total_sessions=len(sessions),
1277 by_client=_group_session_counts(sessions, lambda session: session.client_name),
1278 by_user=_group_session_counts(sessions, lambda session: session.user_id),
1279 sessions=sessions,
1280 )
1282 def _session_matches_admin_selector(
1283 session_id: str,
1284 auth_user: MCPAuthenticatedUser,
1285 session_id_prefix: str | None,
1286 user_id: str | None,
1287 ) -> bool:
1288 if session_id_prefix is not None and not session_id.startswith(session_id_prefix):
1289 return False
1290 if user_id is None:
1291 return True
1292 key_auth: Final = auth_user.user_api_key_auth
1293 return key_auth is not None and key_auth.user_id == user_id
1295 def _forget_expired_admin_terminated_session_ids(now: float) -> None:
1296 for session_id in [ 1296 ↛ 1301line 1296 didn't jump to line 1301 because the loop on line 1296 never started
1297 session_id
1298 for session_id, last_replayed in _admin_terminated_session_ids.items()
1299 if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
1300 ]:
1301 del _admin_terminated_session_ids[session_id]
1303 def _is_admin_terminated_session_id(session_id: str, now: float) -> bool:
1304 last_replayed: Final = _admin_terminated_session_ids.get(session_id)
1305 if last_replayed is None:
1306 return False
1307 if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS:
1308 del _admin_terminated_session_ids[session_id]
1309 return False
1310 _admin_terminated_session_ids[session_id] = now
1311 return True
1313 async def terminate_mcp_gateway_sessions(
1314 *,
1315 session_id_prefix: str | None = None,
1316 user_id: str | None = None,
1317 ) -> MCPGatewaySessionsTerminateResponse:
1318 """Force-close every live stateful session on this worker matching the selector.
1320 The transport is terminated (open streams close), all per-session
1321 tracking is dropped, and the id is remembered so a client that keeps
1322 sending it receives 404 and has to ``initialize`` again, which re-runs
1323 admission. Only sessions held by this worker process are affected.
1324 """
1325 now: Final = time.monotonic()
1326 _forget_expired_admin_terminated_session_ids(now)
1327 server_instances: Final = _stateful_server_instances()
1328 targets: Final = tuple(
1329 (session_id, auth_user)
1330 for session_id, auth_user in tuple(_stateful_session_auth_contexts.items())
1331 if session_id in server_instances
1332 and _session_matches_admin_selector(session_id, auth_user, session_id_prefix, user_id)
1333 )
1334 terminated: Final = tuple(_gateway_session_for(session_id, auth_user, now) for session_id, auth_user in targets)
1335 for session_id, _ in targets: 1335 ↛ 1336line 1335 didn't jump to line 1336 because the loop on line 1335 never started
1336 _admin_terminated_session_ids[session_id] = now
1337 transport = server_instances.pop(session_id, None)
1338 _remove_stateful_session_tracking(session_id)
1339 if transport is not None:
1340 await transport.terminate()
1341 verbose_logger.warning("MCP session '%s' terminated by an administrator.", session_id)
1342 return MCPGatewaySessionsTerminateResponse(
1343 worker_pid=os.getpid(),
1344 terminated_sessions=len(terminated),
1345 sessions=terminated,
1346 )
1348 async def _read_request_body_for_routing(
1349 receive: Receive,
1350 ) -> tuple[list[Message], bytes]:
1351 """
1352 Read just enough of the request body to decide whether this is a
1353 JSON-RPC ``initialize`` call. Returns the consumed ASGI messages so
1354 the caller can replay them faithfully to the downstream handler, and
1355 the peeked body bytes (capped at ``_MCP_ROUTING_PEEK_MAX_BYTES``).
1357 Stops reading from the wire as soon as either (a) we have peeked
1358 ``_MCP_ROUTING_PEEK_MAX_BYTES`` of body, or (b) the body is complete.
1359 The remainder of an oversized body is streamed lazily through
1360 ``wrapped_receive`` in the caller — so an authenticated client cannot
1361 force the proxy to buffer an arbitrarily large payload just to make a
1362 routing decision.
1363 """
1364 consumed_messages: Final[list[Message]] = []
1365 body_chunks: Final[list[bytes]] = []
1366 peeked_bytes = 0
1368 while True:
1369 message = await receive()
1370 consumed_messages.append(message)
1372 if message.get("type") != "http.request": 1372 ↛ 1373line 1372 didn't jump to line 1373 because the condition on line 1372 was never true
1373 break
1375 body: bytes = message.get("body", b"") or b""
1376 if body: 1376 ↛ 1383line 1376 didn't jump to line 1383 because the condition on line 1376 was never true
1377 # Only retain up to the remaining peek budget for sniffing.
1378 # The full ``message`` is already in memory (delivered by
1379 # the ASGI server) and must round-trip to the downstream
1380 # handler via ``consumed_messages``, but ``body_chunks`` is
1381 # purely for the JSON-RPC method check — there is no reason
1382 # to copy a large body frame into a second buffer.
1383 remaining = _MCP_ROUTING_PEEK_MAX_BYTES - peeked_bytes
1384 if remaining > 0:
1385 body_chunks.append(body[:remaining])
1386 peeked_bytes += min(len(body), remaining)
1388 if not message.get("more_body", False): 1388 ↛ 1391line 1388 didn't jump to line 1391 because the condition on line 1388 was always true
1389 break
1391 if peeked_bytes >= _MCP_ROUTING_PEEK_MAX_BYTES:
1392 # Stop draining; downstream replay will pull remaining chunks
1393 # directly from the original `receive` via wrapped_receive.
1394 break
1396 return consumed_messages, b"".join(body_chunks)
1398 async def _handle_stale_mcp_session(
1399 scope: Scope,
1400 receive: Receive,
1401 send: Send,
1402 mgr: "StreamableHTTPSessionManager",
1403 ) -> bool:
1404 """
1405 Inspect the incoming ``mcp-session-id`` header **before** the
1406 request reaches the MCP SDK. If the session is stale (not known
1407 to this worker), strip the header so the SDK creates a fresh
1408 stateless session instead of returning a 400.
1410 Returns:
1411 True if the request was fully handled (e.g. DELETE on
1412 non-existent session). False if the request should continue
1413 to the session manager.
1415 Fixes https://github.com/BerriAI/litellm/issues/20992
1416 """
1417 _mcp_session_header: Final = b"mcp-session-id"
1418 _headers: Final[Sequence[tuple[bytes | str, bytes | str]]] = scope.get("headers", [])
1420 def _normalize_header_name(header_name: object) -> bytes | None:
1421 if isinstance(header_name, bytes):
1422 return header_name.lower()
1423 if isinstance(header_name, str):
1424 return header_name.lower().encode("utf-8", errors="replace")
1425 return None
1427 _session_id: str | None = None
1428 for header_name, header_value in _headers:
1429 if _normalize_header_name(header_name) == _mcp_session_header:
1430 if isinstance(header_value, bytes):
1431 _session_id = header_value.decode("utf-8", errors="replace")
1432 else:
1433 _session_id = str(header_value)
1434 break
1436 if _session_id is None:
1437 return False
1439 # Check in-memory session tracking
1440 known_sessions: Final = getattr(mgr, "_server_instances", None)
1441 # If we cannot inspect known_sessions, let the manager handle it
1442 if known_sessions is None:
1443 return False
1445 # If session exists in this worker's memory, let the manager handle it
1446 try:
1447 if _session_id in known_sessions:
1448 return False
1449 except Exception:
1450 verbose_logger.debug(
1451 "Unable to inspect active MCP sessions for '%s'. Deferring to session manager.",
1452 _session_id,
1453 )
1454 return False
1456 # --- Session not in this worker's memory ---
1457 method: Final = scope.get("method", "").upper()
1459 if method == "DELETE":
1460 _remove_stateful_session_tracking(_session_id)
1461 verbose_logger.info(
1462 "DELETE request for non-existent MCP session '%s'. Returning success (idempotent DELETE).",
1463 _session_id,
1464 )
1465 success_response: Final = JSONResponse(
1466 status_code=200,
1467 content={"message": "Session terminated successfully"},
1468 )
1469 await success_response(scope, receive, send)
1470 return True
1472 if _is_admin_terminated_session_id(_session_id, time.monotonic()):
1473 terminated_response: Final = JSONResponse(
1474 status_code=404,
1475 content={ # mutable-ok: JSONResponse content must be a plain dict
1476 "error": "Not Found",
1477 "details": "mcp-session-id was terminated by an administrator. Send initialize to start a new session.",
1478 },
1479 )
1480 await terminated_response(scope, receive, send)
1481 return True
1483 # Non-DELETE: strip stale session ID to allow new session creation
1484 verbose_logger.warning(
1485 "MCP session ID '%s' not found in this worker's memory. "
1486 "Stripping stale header to force new session creation.",
1487 _session_id,
1488 )
1489 scope["headers"] = [(k, v) for k, v in _headers if _normalize_header_name(k) != _mcp_session_header]
1490 return False
1492 async def _apply_toolset_scope(
1493 user_api_key_auth: UserAPIKeyAuth,
1494 toolset_id: str,
1495 ) -> UserAPIKeyAuth:
1496 """
1497 Restrict a key's MCP permissions to a single toolset.
1499 When a request arrives via /toolset/{name}/mcp we override the key's
1500 object_permission so that only the toolset's tools are visible.
1502 Raises HTTPException(403) if the key has an explicit toolset grant list
1503 that does not include toolset_id (i.e. mcp_toolsets is set but empty,
1504 or set to a list that omits this toolset). Admin keys always pass.
1505 """
1506 from litellm.proxy._types import LiteLLM_ObjectPermissionTable
1507 from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
1509 # A key scoped to no MCP servers opts out of every MCP path. Enforce it
1510 # here too, since toolset scoping replaces mcp_servers and would otherwise
1511 # drop the sentinel. Checked before the admin branch, mirroring
1512 # get_allowed_mcp_servers.
1513 original_op: Final = user_api_key_auth.object_permission
1514 if original_op is not None and SpecialMCPServerNames.no_mcp_servers.value in (original_op.mcp_servers or []): 1514 ↛ 1515line 1514 didn't jump to line 1515 because the condition on line 1514 was never true
1515 raise HTTPException(
1516 status_code=403,
1517 detail="API key is scoped to no MCP servers; toolset access is denied.",
1518 )
1520 # Access control: non-admin keys must have this toolset in their grant list.
1521 # Use _user_has_admin_view so that PROXY_ADMIN_VIEW_ONLY is also treated as admin.
1522 is_admin: Final = _user_has_admin_view(user_api_key_auth)
1523 if not is_admin: 1523 ↛ 1524line 1523 didn't jump to line 1524 because the condition on line 1523 was never true
1524 op: Final = user_api_key_auth.object_permission
1525 granted: Final = getattr(op, "mcp_toolsets", None) if op else None
1526 # granted=None → key has no explicit toolset grants → deny (same semantics as
1527 # fetch_mcp_toolsets which returns [] for non-admin keys with no grants configured).
1528 # granted=[] or list without toolset_id → also deny.
1529 if granted is None or toolset_id not in granted:
1530 raise HTTPException(
1531 status_code=403,
1532 detail=f"API key does not have access to toolset '{toolset_id}'.",
1533 )
1535 tool_permissions = await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(
1536 toolset_ids=[toolset_id]
1537 )
1538 server_ids: Final = list(tool_permissions.keys())
1539 existing_op: Final = user_api_key_auth.object_permission
1540 if existing_op is not None: 1540 ↛ 1541line 1540 didn't jump to line 1541 because the condition on line 1540 was never true
1541 updated_op = existing_op.model_copy(
1542 update={
1543 "mcp_servers": server_ids,
1544 "mcp_tool_permissions": tool_permissions,
1545 "mcp_toolsets": [],
1546 # mcp_access_groups is preserved: a key's access-group grants
1547 # remain valid even when the request is scoped to a single toolset.
1548 }
1549 )
1550 else:
1551 updated_op = LiteLLM_ObjectPermissionTable(
1552 object_permission_id="toolset-scope",
1553 mcp_servers=server_ids,
1554 mcp_tool_permissions=tool_permissions,
1555 )
1556 return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
1558 async def _raise_preemptive_401_for_unauthenticated_servers(
1559 scope: Scope,
1560 mcp_servers: list[str] | None,
1561 oauth2_headers: dict[str, str] | None,
1562 mcp_server_auth_headers: dict[str, dict[str, str]] | None,
1563 user_api_key_auth: UserAPIKeyAuth | None,
1564 client_ip: str | None,
1565 allowed_server_ids: set[str] | None = None,
1566 raw_headers: Mapping[str, str] | None = None,
1567 ) -> None:
1568 """Fail fast with HTTP 401 for MCP servers that need user auth but
1569 didn't receive it on this request. Covers both gateway-managed OAuth2
1570 (points clients at the gateway AS metadata) and pass-through OAuth
1571 (points clients at the upstream resource-metadata via our well-known).
1573 ``allowed_server_ids`` may be passed by callers that have already
1574 narrowed the authorized server set (e.g. toolset scoping); servers
1575 not in that set are skipped so a client targeting a toolset that
1576 excludes a passthrough server is not pushed into an OAuth flow for
1577 a server it will be 403'd on immediately after authentication.
1578 """
1579 for server_name in mcp_servers or []: 1579 ↛ 1580line 1579 didn't jump to line 1580 because the loop on line 1579 never started
1580 server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
1581 if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
1582 # Caller's narrowed scope excludes this server — skip the
1583 # preemptive challenge and let downstream authorization
1584 # return 403.
1585 continue
1586 if server is not None and server.auth_type == MCPAuth.oauth2 and server.oauth2_flow == "client_credentials":
1587 # Stamped M2M: the challenge decision below never reads discovered
1588 # metadata, so deferred-discovery failures must not 503 this loop.
1589 # Unstamped rows stay on the discover-first path because filling
1590 # authorization_url/token_url can change their inferred flow.
1591 continue
1592 if server is not None:
1593 server = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(server)
1594 if server and server.auth_type == MCPAuth.oauth2:
1595 # The challenge decision is per oauth2 sub-mode, not per header:
1596 # gateway-managed modes (M2M and interactive authorization_code)
1597 # never receive a client-supplied upstream token, so a bearer in
1598 # Authorization is a LiteLLM key (surfaced here as oauth2_headers)
1599 # and must not suppress the challenge. Only the delegate mode
1600 # treats a present bearer as the upstream token. The sub-mode is
1601 # resolved the same way egress resolves it, via
1602 # effective_oauth2_flow: an unstamped (null oauth2_flow) row with
1603 # the M2M shape resolves to client_credentials, so the bare
1604 # has_client_credentials column is never trusted here.
1605 if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
1606 # M2M: the gateway mints its own token at egress from the
1607 # stored client credentials, so there is nothing to challenge.
1608 continue
1610 if getattr(server, "delegate_auth_to_upstream", False) is not True:
1611 # Gateway-managed interactive (authorization_code): the only
1612 # thing that authorizes egress is a stored per-user token, so
1613 # challenge whenever one is absent, regardless of any bearer.
1614 # The v2 resolver owns the existence check, so every
1615 # authorization_code resolution (egress and this discovery
1616 # challenge) runs through it. A keyless admitted subject is
1617 # challenged with the per-server resource_metadata (whose
1618 # authorization server is the gateway itself, vaulting via the
1619 # authorize interlude); the per-server relay advertised below
1620 # cannot vault without a litellm key on its token request.
1621 if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
1622 continue
1624 if _is_mcp_admitted_user_subject(user_api_key_auth):
1625 raise HTTPException(
1626 status_code=401,
1627 detail="Unauthorized",
1628 headers={
1629 "www-authenticate": get_passthrough_www_authenticate(
1630 scope=scope,
1631 server_name=server_name,
1632 )
1633 },
1634 )
1636 request = StarletteRequest(scope)
1637 base_url = get_request_base_url(request)
1638 _path = get_route_relative_request_path(scope)
1640 # Pick the well-known AS-metadata form that matches the inbound route
1641 # so strict RFC 9728 §3.2 clients can resolve it correctly.
1642 as_metadata_root = f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}"
1643 if _path.startswith(f"/mcp/{server_name}"):
1644 _as_url = f"{as_metadata_root}/mcp/{server_name}"
1645 else:
1646 _as_url = f"{as_metadata_root}/{server_name}"
1647 authorization_uri = f'Bearer authorization_uri="{_as_url}"'
1649 raise HTTPException(
1650 status_code=401,
1651 detail="Unauthorized",
1652 headers={"www-authenticate": authorization_uri},
1653 )
1655 if not oauth2_headers:
1656 # Delegate-auth servers run upstream PKCE: a present bearer is
1657 # the upstream token, so only challenge when it is absent, with
1658 # the proxied resource_metadata (RFC 9728), not the gateway
1659 # authorization_uri above which would authorize against the
1660 # gateway instead of the upstream IdP.
1661 www_authenticate = get_passthrough_www_authenticate(
1662 scope=scope,
1663 server_name=server_name,
1664 )
1665 raise HTTPException(
1666 status_code=401,
1667 detail="Unauthorized",
1668 headers={"www-authenticate": www_authenticate},
1669 )
1670 # Delegate server with a bearer present: it is the upstream token,
1671 # so admit the session and move to the next target. Every oauth2
1672 # sub-mode is terminal here (continue or raise) so no oauth2 server
1673 # reaches the token_exchange / pass-through blocks below.
1674 continue
1676 # token_exchange (OBO): the caller supplied no subject token. Challenge at connect
1677 # (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata
1678 # so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM
1679 # then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the
1680 # header lost, so the discovery flow needs this pre-emptive challenge.
1681 if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
1682 from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
1683 raise_token_exchange_challenge,
1684 )
1685 from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
1686 get_request_root_path,
1687 )
1689 raise_token_exchange_challenge(server, root_path=get_request_root_path())
1691 # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
1692 # the exchange here at the transport edge, so a rejected subject raises the RFC 9728
1693 # challenge and any other failure its public status, instead of the session opening and
1694 # list_tools masking it as an empty tool list. The manager owns which modes pre-flight
1695 # and what each mints from. Gated to single-server routes the key may reach; the
1696 # multi-server aggregate keeps absorbing per-server auth failures so one bad server
1697 # cannot 401 the whole connect.
1698 if (
1699 server
1700 and len(mcp_servers or []) == 1
1701 and server.server_id
1702 in frozenset(
1703 allowed.server_id
1704 for allowed in await operations._get_allowed_mcp_servers(
1705 user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
1706 )
1707 )
1708 ):
1709 await operations.global_mcp_server_manager.preflight_token_exchange(
1710 server=server,
1711 oauth2_headers=oauth2_headers,
1712 user_api_key_auth=user_api_key_auth,
1713 raw_headers=raw_headers,
1714 )
1716 # Pass-through OAuth: when the admin has opted a server into
1717 # forwarding the client's bearer token (is_oauth_passthrough) and
1718 # the client hasn't supplied one, fail fast with 401 and point
1719 # them at the gateway's oauth-protected-resource well-known URL.
1720 # That endpoint proxies the upstream's metadata so the client
1721 # kicks off OAuth against the real upstream IdP, not the gateway.
1722 if (
1723 server
1724 and server.is_oauth_passthrough
1725 and not operations._client_has_passthrough_authorization(
1726 server, oauth2_headers, mcp_server_auth_headers
1727 )
1728 ):
1729 www_authenticate = get_passthrough_www_authenticate(
1730 scope=scope,
1731 server_name=server_name,
1732 )
1733 raise HTTPException(
1734 status_code=401,
1735 detail="Unauthorized",
1736 headers={"www-authenticate": www_authenticate},
1737 )
1739 if (
1740 server
1741 and server.is_oauth_delegate
1742 and len(mcp_servers or []) == 1
1743 and _get_forwarded_auth_from_scope(scope) is None
1744 and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
1745 ):
1746 www_authenticate = get_passthrough_www_authenticate(
1747 scope=scope,
1748 server_name=server_name,
1749 )
1750 raise HTTPException(
1751 status_code=401,
1752 detail="Unauthorized",
1753 headers={"www-authenticate": www_authenticate},
1754 )
1756 if (
1757 server
1758 and server.is_true_passthrough
1759 and len(mcp_servers or []) == 1
1760 and not _scope_has_authorization_header(scope)
1761 and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
1762 ):
1763 if server.is_dcr_bridge:
1764 raise HTTPException(
1765 status_code=401,
1766 detail="Unauthorized",
1767 headers={
1768 "www-authenticate": get_passthrough_www_authenticate(
1769 scope=scope,
1770 server_name=server_name,
1771 )
1772 },
1773 )
1774 upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "")
1775 if upstream_status == 401 and upstream_www_authenticate:
1776 raise HTTPException(
1777 status_code=401,
1778 detail="Unauthorized",
1779 headers={"www-authenticate": upstream_www_authenticate},
1780 )
1782 def _get_authorization_header_from_scope(scope: Scope) -> str | None:
1783 """First ``Authorization`` header value in the ASGI scope, or None."""
1784 scope_headers: Final[Sequence[tuple[bytes, bytes]]] = scope.get("headers", [])
1785 for key, value in scope_headers:
1786 if key.lower() == b"authorization": 1786 ↛ 1787line 1786 didn't jump to line 1787 because the condition on line 1786 was never true
1787 return value.decode("latin-1")
1788 return None
1790 def _scope_has_authorization_header(scope: Scope) -> bool:
1791 return _get_authorization_header_from_scope(scope) is not None
1793 def _get_forwarded_auth_from_scope(scope: Scope) -> str | None:
1794 """Return the upstream-bound ``Authorization`` header value, or None.
1796 Only returns the ``Authorization`` header when ``x-litellm-api-key`` is
1797 also present. In that case ``Authorization`` is unambiguously the
1798 upstream token the caller wants forwarded to the MCP server. When
1799 ``x-litellm-api-key`` is absent the ``Authorization`` header may itself
1800 be the LiteLLM proxy API key (backward-compat path in
1801 ``MCPRequestHandler.process_mcp_request``), and forwarding it upstream
1802 would leak the proxy key to a third-party MCP server.
1803 """
1804 scope_headers: Final[Sequence[tuple[bytes, bytes]]] = scope.get("headers", [])
1805 has_litellm_key_header: Final = any(key.lower() == b"x-litellm-api-key" for key, _ in scope_headers)
1806 if not has_litellm_key_header: 1806 ↛ 1807line 1806 didn't jump to line 1807 because the condition on line 1806 was never true
1807 return None
1808 return _get_authorization_header_from_scope(scope)
1810 async def _probe_upstream_auth(
1811 url: str,
1812 auth_header: str,
1813 timeout: float = 5.0,
1814 ) -> tuple[int, str | None]:
1815 """JSON-RPC initialize-probe the upstream URL to check whether the token is accepted.
1817 Uses POST so StreamableHTTP MCP servers run the same auth path as a
1818 real client request. Returns (status_code, www_authenticate).
1819 Fails-open with (200, None) on network errors so a transient hiccup
1820 does not block valid requests.
1822 Uses the public ``AsyncHTTPHandler.post()`` interface and catches
1823 ``httpx.HTTPStatusError`` separately so the 401/403 we want to surface
1824 is not swallowed by the broad fail-open ``except Exception`` below.
1825 """
1826 client: Final = get_async_httpx_client(
1827 llm_provider=httpxSpecialProvider.MCP,
1828 params={"timeout": timeout},
1829 )
1830 probe_payload: Final = {
1831 "jsonrpc": "2.0",
1832 "id": "litellm-mcp-auth-probe",
1833 "method": "initialize",
1834 "params": {
1835 "protocolVersion": MCPSpecVersion.jun_2025.value,
1836 "capabilities": {},
1837 "clientInfo": {
1838 "name": "litellm-mcp-auth-probe",
1839 "version": "1.0.0",
1840 },
1841 },
1842 }
1843 probe_headers: Final = {
1844 "Accept": "application/json, text/event-stream",
1845 **({"Authorization": auth_header} if auth_header else {}),
1846 }
1847 try:
1848 resp: Final = await client.post(
1849 url=url,
1850 headers=probe_headers,
1851 json=probe_payload,
1852 timeout=timeout,
1853 )
1854 return resp.status_code, resp.headers.get("www-authenticate")
1855 except httpx.HTTPStatusError as exc:
1856 # AsyncHTTPHandler.post() calls raise_for_status(); a 401/403 from
1857 # upstream lands here. Return its status so the caller can map it
1858 # to the appropriate response.
1859 return exc.response.status_code, exc.response.headers.get("www-authenticate")
1860 except Exception as exc:
1861 verbose_logger.debug("_probe_upstream_auth: probe to %s failed (%s), allowing request through", url, exc)
1862 return 200, None
1864 async def _check_passthrough_upstream_auth(
1865 scope: Scope,
1866 user_api_key_auth: UserAPIKeyAuth | None,
1867 mcp_servers: list[str] | None,
1868 client_ip: str | None,
1869 ) -> None:
1870 """Probe pass-through upstream servers in parallel before the MCP session starts.
1872 Only servers the caller's key is already authorized to reach are probed —
1873 the list is derived from _get_allowed_mcp_servers so that a user cannot
1874 trigger an upstream probe against a server their key is not permitted for.
1876 The MCP SDK commits HTTP 200 headers before invoking handlers, so a 401
1877 can only be returned before that point. This function raises HTTPException(401)
1878 with a WWW-Authenticate header if any upstream rejects the client token, or 403
1879 if the upstream accepts it but forbids the caller.
1880 Fails-open: network errors are logged and the request is allowed through.
1882 """
1883 forwarded_auth: Final = _get_forwarded_auth_from_scope(scope)
1884 if not forwarded_auth: 1884 ↛ 1889line 1884 didn't jump to line 1889 because the condition on line 1884 was always true
1885 return
1887 # Use the authorized server set, not the raw user-supplied names, so that
1888 # a caller cannot force a probe to a server their key is not allowed to use.
1889 allowed_servers: Final = await operations._get_allowed_mcp_servers(
1890 user_api_key_auth=user_api_key_auth,
1891 mcp_servers=mcp_servers,
1892 client_ip=client_ip,
1893 )
1894 passthrough_targets: Final[tuple[tuple[MCPServer, str, str], ...]] = tuple(
1895 (srv, forwarded_auth, srv.name)
1896 for srv in allowed_servers
1897 # Restrict to genuine OAuth pass-through servers (auth_type none +
1898 # Authorization in extra_headers). Gateway-managed OAuth2 servers
1899 # must not receive the ``resource_metadata=`` challenge emitted
1900 # below — they require ``authorization_uri=`` pointing at the
1901 # gateway AS metadata. ``is_oauth_passthrough`` already requires
1902 # ``auth_type in (None, MCPAuth.none)``, which is mutually
1903 # exclusive with ``has_client_credentials`` (oauth2 + M2M flow),
1904 # so M2M servers are implicitly excluded here.
1905 if srv.is_oauth_passthrough
1906 )
1907 probe_targets: Final = passthrough_targets
1908 if not probe_targets:
1909 return
1911 probe_results: Final = await asyncio.gather(
1912 *[_probe_upstream_auth(srv.url or "", auth_header) for srv, auth_header, _ in probe_targets]
1913 )
1914 for (srv, _, challenge_server_name), (probe_status, _) in zip(probe_targets, probe_results):
1915 if probe_status == 401:
1916 # Token is missing or expired: keep pass-through clients on the
1917 # protected-resource discovery flow so they re-authorize against
1918 # the upstream IdP metadata proxied by LiteLLM.
1919 www_authenticate = get_passthrough_www_authenticate(
1920 scope=scope,
1921 server_name=challenge_server_name,
1922 invalid_token=True,
1923 )
1924 raise HTTPException(
1925 status_code=401,
1926 detail="Unauthorized",
1927 headers={"www-authenticate": www_authenticate},
1928 )
1929 if probe_status == 403:
1930 # Token is valid but the caller lacks permission — do not hint
1931 # at re-authorization (RFC 9110: a fresh token with the same
1932 # scopes would just hit 403 again and loop indefinitely).
1933 raise HTTPException(
1934 status_code=403,
1935 detail="Forbidden",
1936 )
1938 async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None:
1939 """Handle MCP requests through StreamableHTTP."""
1940 try:
1941 reject_disallowed_mcp_origin(StarletteRequest(scope))
1942 bad_version: Final = unsupported_protocol_version(scope)
1943 if bad_version is not None: 1943 ↛ 1944line 1943 didn't jump to line 1944 because the condition on line 1943 was never true
1944 supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
1945 await JSONResponse(
1946 status_code=400,
1947 content={ # mutable-ok: JSON-RPC error payload
1948 "jsonrpc": "2.0",
1949 "id": None,
1950 "error": {
1951 "code": INVALID_REQUEST,
1952 "message": f"Unsupported MCP-Protocol-Version {bad_version}; supported: {supported}",
1953 },
1954 },
1955 )(scope, receive, send)
1956 return
1957 path: Final[str] = scope.get("path", "")
1958 (
1959 user_api_key_auth,
1960 mcp_auth_header,
1961 mcp_servers,
1962 mcp_server_auth_headers,
1963 oauth2_headers,
1964 raw_headers,
1965 ) = await extract_mcp_auth_context(scope, path)
1966 reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth)
1967 scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1
1969 # Extract client IP for MCP access control
1970 _client_ip: Final = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope))
1972 verbose_logger.debug("MCP request mcp_servers (header/path): %s", mcp_servers)
1973 verbose_logger.debug(
1974 "MCP server auth headers: %s", list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None
1975 )
1977 # Strip any client-supplied x-mcp-toolset-id to prevent forgery.
1978 scope_headers: Final[Sequence[tuple[bytes, bytes]]] = scope.get("headers", [])
1979 scope["headers"] = [(k, v) for k, v in scope_headers if k.lower() != b"x-mcp-toolset-id"]
1981 # Apply toolset scope if set server-side via ContextVar (set by
1982 # /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py).
1983 active_toolset_id: Final = _mcp_active_toolset_id.get()
1984 toolset_allowed_server_ids: set[str] | None = None
1985 if active_toolset_id and user_api_key_auth is not None:
1986 user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
1987 op: Final = user_api_key_auth.object_permission
1988 toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set()
1990 # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
1991 # Must run after toolset scoping so the challenge set is derived
1992 # from the fully-authorized server set: a passthrough server that
1993 # the active toolset excludes should not trigger an OAuth flow
1994 # for a server the caller will be 403'd on after authentication.
1995 await _raise_preemptive_401_for_unauthenticated_servers(
1996 scope=scope,
1997 mcp_servers=mcp_servers,
1998 oauth2_headers=oauth2_headers,
1999 mcp_server_auth_headers=mcp_server_auth_headers,
2000 user_api_key_auth=user_api_key_auth,
2001 client_ip=_client_ip,
2002 allowed_server_ids=toolset_allowed_server_ids,
2003 raw_headers=raw_headers,
2004 )
2006 # Pre-flight auth check for pass-through servers. Must run after
2007 # toolset scoping so the probe list is derived from the fully-authorized
2008 # server set, not the raw user-supplied names.
2009 await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _client_ip)
2011 # Inject masked debug headers when client sends x-litellm-mcp-debug: true
2012 _debug_headers: Final = MCPDebug.maybe_build_debug_headers(
2013 raw_headers=raw_headers,
2014 scope=dict(scope),
2015 mcp_servers=mcp_servers,
2016 oauth2_headers=oauth2_headers,
2017 client_ip=_client_ip,
2018 )
2019 diagnostics: Final = MCPAuthDiagnostics() if _debug_headers else None
2020 if diagnostics is not None: 2020 ↛ 2021line 2020 didn't jump to line 2021 because the condition on line 2020 was never true
2021 scope[MCP_AUTH_DIAGNOSTICS_SCOPE_KEY] = diagnostics
2022 send = MCPDebug.wrap_send_with_debug_headers(
2023 send, _debug_headers, diagnostics.headers, request_method=scope.get("method")
2024 )
2026 # Ensure session managers are initialized
2027 if not _SESSION_MANAGERS_INITIALIZED:
2028 await initialize_session_managers()
2029 # Give it a moment to start up
2030 await asyncio.sleep(0.1)
2032 # Route based on mcp-session-id and request method:
2033 # - Has session ID → stateful (Claude Code, Cursor, VSCode)
2034 # - No session ID + initialize → stateful (so client gets mcp-session-id)
2035 # - No session ID + other → stateless (curl, Inspector, Notion)
2036 session_id = _get_session_id_from_scope(scope)
2037 is_initialize = False
2038 consumed_messages: list[Message] = []
2040 # Owner-binding: a live stateful session may only be driven by the
2041 # caller that created it. Reject mismatches with 403 so a leaked
2042 # mcp-session-id cannot be hijacked by another authenticated user.
2043 #
2044 # Run before ``_handle_stale_mcp_session`` so a non-owner cannot
2045 # force-clean another caller's residual tracking entries via a
2046 # stale DELETE, and before peeking the request body so the 403
2047 # response sees a pristine ``receive`` channel.
2048 if session_id: 2048 ↛ 2049line 2048 didn't jump to line 2049 because the condition on line 2048 was never true
2049 expected_owner: Final = _stateful_session_owners.get(session_id)
2050 request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
2051 if expected_owner is not None and expected_owner != request_owner:
2052 verbose_logger.warning(
2053 "Rejecting MCP request: session '%s' owner mismatch.",
2054 session_id,
2055 )
2056 forbidden_response: Final = JSONResponse(
2057 status_code=403,
2058 content={
2059 "error": "Forbidden",
2060 "details": "mcp-session-id is bound to a different caller.",
2061 },
2062 )
2063 await forbidden_response(scope, receive, send)
2064 return
2066 # Handle stale session IDs before choosing a target manager. Stale
2067 # non-DELETE requests have their session header stripped and should
2068 # be routed as no-session requests.
2069 if session_id: 2069 ↛ 2070line 2069 didn't jump to line 2070 because the condition on line 2069 was never true
2070 handled: Final = await _handle_stale_mcp_session(scope, receive, send, session_manager_stateful)
2071 if handled:
2072 # Request was fully handled (e.g., DELETE on non-existent session)
2073 return
2074 session_id = _get_session_id_from_scope(scope)
2076 body = b""
2077 if scope.get("method") == "POST":
2078 consumed_messages, body = await _read_request_body_for_routing(receive)
2079 is_initialize = _is_initialize_request(body)
2081 use_stateful: Final = bool(session_id or is_initialize)
2082 target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless
2084 verbose_logger.debug(
2085 f"MCP routing to {'stateful' if use_stateful else 'stateless'} manager"
2086 + (f" (session={session_id[:8]}...)" if session_id else "")
2087 + (" (initialize)" if is_initialize else "")
2088 )
2090 # A new `initialize` (no session id) is about to create a stateful
2091 # session. Cap how many a single caller can hold so an authenticated
2092 # client cannot spam `initialize` and exhaust memory.
2093 if is_initialize and not session_id: 2093 ↛ 2094line 2093 didn't jump to line 2094 because the condition on line 2093 was never true
2094 request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
2095 if not await _enforce_stateful_session_cap_for_owner(request_owner):
2096 verbose_logger.warning(
2097 "Rejecting MCP initialize: caller already holds the maximum number of active stateful sessions."
2098 )
2099 too_many_response: Final = JSONResponse(
2100 status_code=429,
2101 content={
2102 "error": "Too Many Requests",
2103 "details": "Too many active MCP sessions for this caller.",
2104 },
2105 )
2106 await too_many_response(scope, receive, send)
2107 return
2109 # Replay body messages if we consumed them for peeking
2110 original_receive: Final = receive
2111 if consumed_messages:
2113 async def wrapped_receive():
2114 if consumed_messages: 2114 ↛ 2116line 2114 didn't jump to line 2116 because the condition on line 2114 was always true
2115 return consumed_messages.pop(0)
2116 return await original_receive()
2118 receive = wrapped_receive
2120 # Serialize requests on the same stateful session so concurrent
2121 # callers don't clobber each other's auth context mid-flight.
2122 #
2123 # Skip the lock for streaming GETs (SSE channels held open for the
2124 # life of the session): holding a per-session lock for a long-lived
2125 # stream would block every subsequent POST on the same session.
2126 # POST/DELETE are the methods that actually mutate the shared
2127 # auth context, so serializing those is sufficient for the
2128 # clobbering race between concurrent JSON-RPC calls.
2129 #
2130 # Also skip the lock for JSON-RPC *responses* (POSTs that carry
2131 # a ``result`` or ``error`` but no ``method``). These are replies
2132 # to server-initiated requests such as ``elicitation/create`` or
2133 # ``sampling/createMessage``. The in-flight tool-call POST that
2134 # triggered the server request already holds the session lock, so
2135 # trying to acquire it again for the response POST would deadlock.
2136 is_jsonrpc_response = False
2137 request_method: Final = (scope.get("method") or "").upper()
2138 if body and request_method == "POST": 2138 ↛ 2139line 2138 didn't jump to line 2139 because the condition on line 2138 was never true
2139 try:
2140 _peeked: Final = json.loads(body)
2141 if (
2142 isinstance(_peeked, dict)
2143 and _peeked.get("jsonrpc") == "2.0"
2144 and "id" in _peeked
2145 and "method" not in _peeked
2146 and ("result" in _peeked or "error" in _peeked)
2147 ):
2148 is_jsonrpc_response = True
2149 verbose_logger.debug(
2150 "MCP: detected JSON-RPC response POST (id=%s), skipping session lock to avoid deadlock",
2151 _peeked.get("id"),
2152 )
2153 except (json.JSONDecodeError, UnicodeDecodeError, TypeError):
2154 # Peek cap truncated the body, so it can't be fully parsed.
2155 # Scan the top-level keys (depth-aware) instead of a flat
2156 # substring search: a response's result payload may nest a
2157 # "method" field, and misreading that would acquire the lock
2158 # and deadlock the in-flight tool call awaiting this
2159 # response. A false skip is harmless; a false acquire is not.
2160 _body_str: Final = body.decode("utf-8", errors="replace")
2161 if (
2162 '"jsonrpc"' in _body_str
2163 and ('"result"' in _body_str or '"error"' in _body_str)
2164 and not _jsonrpc_text_has_top_level_method(_body_str)
2165 ):
2166 is_jsonrpc_response = True
2167 verbose_logger.debug(
2168 "MCP: detected truncated JSON-RPC response POST via "
2169 "top-level key scan, skipping session lock to avoid deadlock"
2170 )
2172 session_lock: asyncio.Lock | None = None
2173 if use_stateful and session_id and request_method in ("POST", "DELETE") and not is_jsonrpc_response: 2173 ↛ 2174line 2173 didn't jump to line 2174 because the condition on line 2173 was never true
2174 session_lock = _stateful_session_locks.setdefault(session_id, asyncio.Lock())
2176 active_request_session_ids: Final[list[str]] = []
2178 def _increment_active_request_session(session_id_to_track: str) -> None:
2179 if session_id_to_track in active_request_session_ids:
2180 return
2181 active_request_session_ids.append(session_id_to_track)
2182 _stateful_session_active_request_counts[session_id_to_track] = (
2183 _stateful_session_active_request_counts.get(session_id_to_track, 0) + 1
2184 )
2186 if use_stateful and session_id: 2186 ↛ 2187line 2186 didn't jump to line 2187 because the condition on line 2186 was never true
2187 _increment_active_request_session(session_id)
2189 def _track_initialized_stateful_session(
2190 initialized_session_id: str,
2191 ) -> None:
2192 _increment_active_request_session(initialized_session_id)
2194 async def _dispatch() -> None:
2195 _otel_publish_transport_span_on_scope(scope)
2196 _otel_publish_request_destinations_on_scope(scope)
2197 auth_user: Final = _set_or_update_auth_context(
2198 user_api_key_auth=user_api_key_auth,
2199 mcp_auth_header=mcp_auth_header,
2200 mcp_servers=mcp_servers,
2201 mcp_server_auth_headers=mcp_server_auth_headers,
2202 oauth2_headers=oauth2_headers,
2203 raw_headers=raw_headers,
2204 client_ip=_client_ip,
2205 session_id=session_id if use_stateful else None,
2206 touch_last_seen=(scope.get("method") or "").upper() != "DELETE",
2207 copy_existing_session_auth_context=is_initialize,
2208 )
2209 local_send = send
2210 if use_stateful and is_initialize: 2210 ↛ 2211line 2210 didn't jump to line 2211 because the condition on line 2210 was never true
2211 local_send = _wrap_send_with_stateful_session_auth_context(
2212 local_send,
2213 auth_user,
2214 _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip),
2215 _track_initialized_stateful_session,
2216 client_info=_extract_initialize_client_info(body),
2217 )
2219 async with _gateway_initialize_instructions_request_scope(
2220 user_api_key_auth,
2221 mcp_servers,
2222 _client_ip,
2223 scoped_server_endpoint=scoped_server_endpoint,
2224 is_initialize=is_initialize,
2225 ):
2226 await target_manager.handle_request(scope, receive, local_send)
2227 if use_stateful and session_id and scope.get("method") == "DELETE": 2227 ↛ 2228line 2227 didn't jump to line 2228 because the condition on line 2227 was never true
2228 _remove_stateful_session_tracking(session_id)
2230 try:
2231 if session_lock is not None: 2231 ↛ 2232line 2231 didn't jump to line 2232 because the condition on line 2231 was never true
2232 async with session_lock:
2233 await _dispatch()
2234 else:
2235 await _dispatch()
2236 finally:
2237 for active_request_session_id in active_request_session_ids: 2237 ↛ 2238line 2237 didn't jump to line 2238 because the loop on line 2237 never started
2238 active_request_count = _stateful_session_active_request_counts.get(active_request_session_id, 0) - 1
2239 if active_request_count > 0:
2240 _stateful_session_active_request_counts[active_request_session_id] = active_request_count
2241 else:
2242 _stateful_session_active_request_counts.pop(active_request_session_id, None)
2244 if scope.get("method") != "DELETE" and active_request_session_id in _stateful_session_auth_contexts:
2245 _stateful_session_auth_context_last_seen[active_request_session_id] = time.monotonic()
2247 # Periodic cleanup iterates _stateful_session_auth_context_last_seen,
2248 # so locks for untracked sessions must be dropped here.
2249 if active_request_count <= 0 and active_request_session_id not in _stateful_session_auth_contexts:
2250 _stateful_session_locks.pop(active_request_session_id, None)
2251 except MCPUpstreamAuthError as e:
2252 # Upstream delegated auth returned 401; surface it to the client so
2253 # standards-compliant MCP clients trigger the upstream OAuth flow.
2254 raise e.to_http_exception(
2255 base_url=get_request_base_url(StarletteRequest(scope)),
2256 request_path=scope.get("_original_path") or scope.get("path"),
2257 )
2258 except HTTPException:
2259 # Re-raise HTTP exceptions to preserve status codes and details
2260 raise
2261 except ProxyException as e:
2262 # Auth failures from user_api_key_auth arrive as ProxyException, not
2263 # HTTPException. Preserve the real status (e.g. 401 + WWW-Authenticate)
2264 # so OAuth clients can re-authenticate instead of receiving a generic
2265 # 500 that surfaces as a cancelled tool call.
2266 raise _proxy_exception_to_http_exception(e)
2267 except Exception as e:
2268 verbose_logger.exception("Error handling MCP request: %s", e)
2269 # Try to send a graceful error response for non-HTTP exceptions
2270 try:
2271 from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR
2273 error_response: Final = JSONResponse(
2274 status_code=HTTP_500_INTERNAL_SERVER_ERROR,
2275 content={"error": "MCP request failed", "details": str(e)},
2276 )
2277 await error_response(scope, receive, send)
2278 except Exception as response_error:
2279 verbose_logger.exception("Failed to send error response: %s", response_error)
2280 # If we can't send a proper response, re-raise the original error
2281 raise e
2283 async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None:
2284 """Handle MCP requests through SSE."""
2285 try:
2286 reject_disallowed_mcp_origin(StarletteRequest(scope))
2287 bad_version: Final = unsupported_protocol_version(scope)
2288 if bad_version is not None:
2289 supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
2290 await JSONResponse(
2291 status_code=400,
2292 content={ # mutable-ok: JSON-RPC error payload
2293 "jsonrpc": "2.0",
2294 "id": None,
2295 "error": {
2296 "code": INVALID_REQUEST,
2297 "message": f"Unsupported MCP-Protocol-Version {bad_version}; supported: {supported}",
2298 },
2299 },
2300 )(scope, receive, send)
2301 return
2302 from litellm.proxy.auth.auth_utils import get_request_route
2304 path: Final = get_request_route(StarletteRequest(scope))
2305 (
2306 user_api_key_auth,
2307 mcp_auth_header,
2308 mcp_servers,
2309 mcp_server_auth_headers,
2310 oauth2_headers,
2311 raw_headers,
2312 ) = await extract_mcp_auth_context(scope, path)
2313 reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth)
2314 scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1
2316 # Extract client IP for MCP access control
2317 _sse_client_ip: Final = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope))
2319 verbose_logger.debug("MCP request mcp_servers (header/path): %s", mcp_servers)
2320 verbose_logger.debug(
2321 "MCP server auth headers: %s", list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None
2322 )
2324 # Strip any client-supplied x-mcp-toolset-id to prevent forgery.
2325 scope_headers: Final[Sequence[tuple[bytes, bytes]]] = scope.get("headers", [])
2326 scope["headers"] = [(k, v) for k, v in scope_headers if k.lower() != b"x-mcp-toolset-id"]
2328 # Apply toolset scope if set server-side via ContextVar so the
2329 # downstream probe list matches the fully-authorized server set
2330 # (mirrors the streamable HTTP handler).
2331 active_toolset_id: Final = _mcp_active_toolset_id.get()
2332 toolset_allowed_server_ids: set[str] | None = None
2333 if active_toolset_id and user_api_key_auth is not None:
2334 user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
2335 op: Final = user_api_key_auth.object_permission
2336 toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set()
2338 # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
2339 # Must run after toolset scoping so the challenge set is derived
2340 # from the fully-authorized server set: a passthrough server that
2341 # the active toolset excludes should not trigger an OAuth flow
2342 # for a server the caller will be 403'd on after authentication.
2343 await _raise_preemptive_401_for_unauthenticated_servers(
2344 scope=scope,
2345 mcp_servers=mcp_servers,
2346 oauth2_headers=oauth2_headers,
2347 mcp_server_auth_headers=mcp_server_auth_headers,
2348 user_api_key_auth=user_api_key_auth,
2349 client_ip=_sse_client_ip,
2350 allowed_server_ids=toolset_allowed_server_ids,
2351 raw_headers=raw_headers,
2352 )
2354 # Pre-flight auth check for pass-through servers: surface upstream
2355 # 401/403 as a proper challenge before the SSE session commits 200
2356 # headers, so clients can refresh their OAuth token instead of
2357 # being stuck with a silently empty tool list. Must run after
2358 # toolset scoping so the probe list is derived from the fully-
2359 # authorized server set, not the raw user-supplied names.
2360 await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _sse_client_ip)
2361 set_auth_context(
2362 user_api_key_auth=user_api_key_auth,
2363 mcp_auth_header=mcp_auth_header,
2364 mcp_servers=mcp_servers,
2365 mcp_server_auth_headers=mcp_server_auth_headers,
2366 oauth2_headers=oauth2_headers,
2367 raw_headers=raw_headers,
2368 client_ip=_sse_client_ip,
2369 )
2371 owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _sse_client_ip)
2372 transport_scope: Final[Scope] = {
2373 **scope,
2374 "user": AuthenticatedUser(AccessToken(token=owner, client_id=owner, scopes=[])),
2375 }
2376 if scope["method"] == "POST":
2377 await sse.handle_post_message(transport_scope, receive, send)
2378 return
2380 async with _gateway_initialize_instructions_request_scope(
2381 user_api_key_auth,
2382 mcp_servers,
2383 _sse_client_ip,
2384 scoped_server_endpoint=scoped_server_endpoint,
2385 is_initialize=scope.get("method") == "GET",
2386 ):
2387 async with sse.connect_sse(transport_scope, receive, send) as (read_stream, write_stream):
2388 await server.run(read_stream, write_stream, server.create_initialization_options())
2389 except MCPUpstreamAuthError as e:
2390 # Upstream delegated auth returned 401; surface it to the client so
2391 # standards-compliant MCP clients trigger the upstream OAuth flow.
2392 raise e.to_http_exception(
2393 base_url=get_request_base_url(StarletteRequest(scope)),
2394 request_path=scope.get("_original_path") or scope.get("path"),
2395 )
2396 except HTTPException:
2397 # Re-raise HTTP exceptions to preserve status codes and details
2398 # (e.g. 401 + WWW-Authenticate challenges from OAuth pass-through).
2399 raise
2400 except ProxyException as e:
2401 # Auth failures from user_api_key_auth arrive as ProxyException, not
2402 # HTTPException. Preserve the real status (e.g. 401 + WWW-Authenticate)
2403 # so OAuth clients can re-authenticate instead of receiving a generic
2404 # 500 that surfaces as a cancelled tool call.
2405 raise _proxy_exception_to_http_exception(e)
2406 except Exception as e:
2407 verbose_logger.exception("Error handling MCP request: %s", e)
2408 # Try to send a graceful error response for non-HTTP exceptions
2409 try:
2410 # Send a proper HTTP error response instead of letting the exception bubble up
2411 from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR
2413 error_response: Final = JSONResponse(
2414 status_code=HTTP_500_INTERNAL_SERVER_ERROR,
2415 content={"error": "MCP request failed", "details": str(e)},
2416 )
2417 await error_response(scope, receive, send)
2418 except Exception as response_error:
2419 verbose_logger.exception("Failed to send error response: %s", response_error)
2420 # If we can't send a proper response, re-raise the original error
2421 raise e
2423 app = FastAPI(
2424 title=LITELLM_MCP_SERVER_NAME,
2425 description=LITELLM_MCP_SERVER_DESCRIPTION,
2426 version=LITELLM_MCP_SERVER_VERSION,
2427 lifespan=lifespan,
2428 )
2430 # Routes
2431 @app.get(
2432 "/enabled",
2433 description="Returns if the MCP server is enabled",
2434 )
2435 def get_mcp_server_enabled() -> dict[str, bool]:
2436 """
2437 Returns if the MCP server is enabled
2438 """
2439 return {"enabled": MCP_AVAILABLE}
2441 class _LegacySseEndpoint:
2442 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
2443 await handle_sse_mcp(scope, receive, send)
2445 for sse_path, sse_method in (
2446 ("/sse", "GET"),
2447 ("/sse/", "GET"),
2448 ("/sse/messages", "POST"),
2449 ("/sse/messages/", "POST"),
2450 ):
2451 app.router.routes.append(Route(sse_path, endpoint=_LegacySseEndpoint(), methods=[sse_method]))
2453 # Mount the MCP handlers
2454 app.mount("/", handle_streamable_http_mcp)
2455 app.mount("/mcp", handle_streamable_http_mcp)
2456 app.mount("/{mcp_server_name}/mcp", handle_streamable_http_mcp)
2457 app.add_middleware(AuthContextMiddleware)
2459 ########################################################
2460 ############ Auth Context Functions ####################
2461 ########################################################
2463 def _update_auth_context(
2464 auth_user: MCPAuthenticatedUser,
2465 user_api_key_auth: UserAPIKeyAuth | None,
2466 mcp_auth_header: str | None = None,
2467 mcp_servers: list[str] | None = None,
2468 mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
2469 oauth2_headers: dict[str, str] | None = None,
2470 raw_headers: dict[str, str] | None = None,
2471 client_ip: str | None = None,
2472 ) -> None:
2473 auth_user.user_api_key_auth = user_api_key_auth
2474 auth_user.mcp_auth_header = mcp_auth_header
2475 auth_user.mcp_servers = mcp_servers
2476 auth_user.mcp_server_auth_headers = mcp_server_auth_headers or {}
2477 auth_user.oauth2_headers = oauth2_headers
2478 auth_user.raw_headers = raw_headers
2479 auth_user.client_ip = client_ip
2481 def set_auth_context(
2482 user_api_key_auth: UserAPIKeyAuth | None,
2483 mcp_auth_header: str | None = None,
2484 mcp_servers: list[str] | None = None,
2485 mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
2486 oauth2_headers: dict[str, str] | None = None,
2487 raw_headers: dict[str, str] | None = None,
2488 client_ip: str | None = None,
2489 ) -> MCPAuthenticatedUser:
2490 """
2491 Set the UserAPIKeyAuth in the auth context variable.
2493 Args:
2494 user_api_key_auth: UserAPIKeyAuth object
2495 mcp_auth_header: MCP auth header to be passed to the MCP server (deprecated)
2496 mcp_servers: Optional list of server names and access groups to filter by
2497 mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
2498 client_ip: Client IP address for MCP access control
2499 """
2500 auth_user: Final = MCPAuthenticatedUser(
2501 user_api_key_auth=user_api_key_auth,
2502 mcp_auth_header=mcp_auth_header,
2503 mcp_servers=mcp_servers,
2504 mcp_server_auth_headers=mcp_server_auth_headers,
2505 oauth2_headers=oauth2_headers,
2506 raw_headers=raw_headers,
2507 client_ip=client_ip,
2508 )
2509 auth_context_var.set(auth_user)
2510 return auth_user
2512 def _set_or_update_auth_context(
2513 user_api_key_auth: UserAPIKeyAuth | None,
2514 mcp_auth_header: str | None = None,
2515 mcp_servers: list[str] | None = None,
2516 mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
2517 oauth2_headers: dict[str, str] | None = None,
2518 raw_headers: dict[str, str] | None = None,
2519 client_ip: str | None = None,
2520 session_id: str | None = None,
2521 touch_last_seen: bool = True,
2522 copy_existing_session_auth_context: bool = False,
2523 ) -> MCPAuthenticatedUser:
2524 auth_user: Final = _stateful_session_auth_contexts.get(session_id) if session_id else None
2525 if auth_user is not None and session_id is not None: 2525 ↛ 2526line 2525 didn't jump to line 2526 because the condition on line 2525 was never true
2526 if touch_last_seen:
2527 _stateful_session_auth_context_last_seen[session_id] = time.monotonic()
2528 if copy_existing_session_auth_context:
2529 return set_auth_context(
2530 user_api_key_auth=user_api_key_auth,
2531 mcp_auth_header=mcp_auth_header,
2532 mcp_servers=mcp_servers,
2533 mcp_server_auth_headers=mcp_server_auth_headers,
2534 oauth2_headers=oauth2_headers,
2535 raw_headers=raw_headers,
2536 client_ip=client_ip,
2537 )
2538 _update_auth_context(
2539 auth_user=auth_user,
2540 user_api_key_auth=user_api_key_auth,
2541 mcp_auth_header=mcp_auth_header,
2542 mcp_servers=mcp_servers,
2543 mcp_server_auth_headers=mcp_server_auth_headers,
2544 oauth2_headers=oauth2_headers,
2545 raw_headers=raw_headers,
2546 client_ip=client_ip,
2547 )
2548 auth_context_var.set(auth_user)
2549 return auth_user
2550 return set_auth_context(
2551 user_api_key_auth=user_api_key_auth,
2552 mcp_auth_header=mcp_auth_header,
2553 mcp_servers=mcp_servers,
2554 mcp_server_auth_headers=mcp_server_auth_headers,
2555 oauth2_headers=oauth2_headers,
2556 raw_headers=raw_headers,
2557 client_ip=client_ip,
2558 )
2560 def _wrap_send_with_stateful_session_auth_context(
2561 send: Send,
2562 auth_user: MCPAuthenticatedUser,
2563 owner_fingerprint: str,
2564 on_session_registered: Callable[[str], None] | None = None,
2565 client_info: Implementation | None = None,
2566 ) -> Send:
2567 async def wrapped_send(message: Message) -> None:
2568 if message.get("type") == "http.response.start":
2569 response_headers: Final[Sequence[tuple[bytes | str, bytes | str]]] = message.get("headers", [])
2570 for key, value in response_headers:
2571 header_name = key if isinstance(key, bytes) else str(key).encode()
2572 if header_name.lower() == b"mcp-session-id":
2573 session_id = value.decode() if isinstance(value, bytes) else str(value)
2574 if on_session_registered is not None:
2575 on_session_registered(session_id)
2576 auth_context_var.set(auth_user)
2577 _stateful_session_auth_contexts[session_id] = auth_user
2578 _stateful_session_auth_context_last_seen[session_id] = time.monotonic()
2579 _stateful_session_owners[session_id] = owner_fingerprint
2580 if client_info is not None:
2581 _stateful_session_client_info[session_id] = client_info
2582 break
2583 await send(message)
2585 return wrapped_send
2587 def get_auth_context() -> tuple[
2588 UserAPIKeyAuth | None,
2589 str | None,
2590 list[str] | None,
2591 dict[str, dict[str, str]] | None,
2592 dict[str, str] | None,
2593 dict[str, str] | None,
2594 str | None,
2595 ]:
2596 """
2597 Get the UserAPIKeyAuth from the auth context variable.
2599 Returns:
2600 Tuple containing: UserAPIKeyAuth, MCP auth header (deprecated),
2601 MCP servers, server-specific auth headers, OAuth2 headers, raw headers, client IP
2602 """
2603 auth_user: Final = auth_context_var.get()
2604 if auth_user and isinstance(auth_user, MCPAuthenticatedUser):
2605 return (
2606 auth_user.user_api_key_auth,
2607 auth_user.mcp_auth_header,
2608 auth_user.mcp_servers,
2609 auth_user.mcp_server_auth_headers,
2610 auth_user.oauth2_headers,
2611 auth_user.raw_headers,
2612 auth_user.client_ip,
2613 )
2614 return None, None, None, None, None, None, None
2616 def _get_current_session():
2617 ctx: Final = get_active_mcp_request_ctx()
2618 return ctx.session if ctx is not None else None
2620 def _cache_auth_context_lazily():
2621 session: Final = _get_current_session()
2622 if session is None:
2623 return
2624 try:
2625 if session in _session_obj_auth_storage:
2626 return
2627 except TypeError:
2628 verbose_logger.debug(
2629 "_cache_auth_context_lazily: session object is unhashable (type=%s), cannot cache auth context",
2630 type(session).__name__,
2631 )
2632 return
2634 auth: Final = auth_context_var.get()
2635 if auth and isinstance(auth, MCPAuthenticatedUser):
2636 try:
2637 _session_obj_auth_storage[session] = auth
2638 except TypeError:
2639 verbose_logger.debug(
2640 "_cache_auth_context_lazily: could not store auth via "
2641 "session identity — session object is unhashable"
2642 )
2644 def _recover_auth_from_session() -> MCPAuthenticatedUser | None:
2645 session: Final = _get_current_session()
2646 if session is None:
2647 return None
2649 stored: MCPAuthenticatedUser | None = None
2650 try:
2651 stored = _session_obj_auth_storage.get(session)
2652 except TypeError:
2653 verbose_logger.debug(
2654 "_recover_auth_from_session: session object is unhashable "
2655 "(type=%s), skipping _session_obj_auth_storage lookup",
2656 type(session).__name__,
2657 )
2659 return stored
2661 async def get_or_extract_auth_context() -> tuple[
2662 UserAPIKeyAuth | None,
2663 str | None,
2664 list[str] | None,
2665 dict[str, dict[str, str]] | None,
2666 dict[str, str] | None,
2667 dict[str, str] | None,
2668 str | None,
2669 ]:
2670 """
2671 Get auth context from ContextVar first, then fall back to session
2672 storage (which survives cross-task boundaries in the MCP SDK).
2673 """
2674 (
2675 user_api_key_auth,
2676 mcp_auth_header,
2677 mcp_servers,
2678 mcp_server_auth_headers,
2679 oauth2_headers,
2680 raw_headers,
2681 _client_ip,
2682 ) = get_auth_context()
2684 if user_api_key_auth is not None:
2685 _cache_auth_context_lazily()
2686 else:
2687 stored: Final = _recover_auth_from_session()
2689 if stored:
2690 user_api_key_auth = stored.user_api_key_auth
2691 mcp_auth_header = stored.mcp_auth_header
2692 mcp_servers = stored.mcp_servers
2693 mcp_server_auth_headers = stored.mcp_server_auth_headers
2694 oauth2_headers = stored.oauth2_headers
2695 raw_headers = stored.raw_headers
2696 _client_ip = stored.client_ip
2697 return (
2698 user_api_key_auth,
2699 mcp_auth_header,
2700 mcp_servers,
2701 mcp_server_auth_headers,
2702 oauth2_headers,
2703 raw_headers,
2704 _client_ip,
2705 )
2707 def get_active_mcp_session() -> _McpServerSession | None:
2708 """Return the active MCP session captured during handler execution."""
2709 session: Final = active_mcp_session_var.get()
2710 if session is not None:
2711 return session
2712 return _get_current_session()
2714 def get_active_auth_context() -> MCPAuthenticatedUser | None:
2715 """Return auth context from ContextVar or session storage."""
2716 auth: Final = auth_context_var.get()
2717 if auth and isinstance(auth, MCPAuthenticatedUser):
2718 return auth
2720 stored: Final = _recover_auth_from_session()
2721 if stored is not None:
2722 return stored
2723 return None
2725 ########################################################
2726 ############ End of Auth Context Functions #############
2727 ########################################################
2729else:
2730 app = FastAPI()