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

1""" 

2LiteLLM MCP Server Routes 

3""" 

4 

5# pyright: reportInvalidTypeForm=false, reportArgumentType=false, reportOptionalCall=false 

6 

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 

18 

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 

26 

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 

90 

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 

93 

94 

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" 

112 

113 

114def reject_disallowed_mcp_origin(request: StarletteRequest) -> None: 

115 from litellm.proxy.proxy_server import origins # noqa: PLC0415 # proxy imports this module during startup 

116 

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

119 

120 

121def unsupported_protocol_version(scope: Scope) -> str | None: 

122 """Return the unsupported ``MCP-Protocol-Version`` header value, if any. 

123 

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 

136 

137 

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 

145 

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 ) 

155 

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 

171 

172active_mcp_session_var: Final[contextvars.ContextVar["_McpServerSession | None"]] = contextvars.ContextVar( 

173 "active_mcp_session", default=None 

174) 

175 

176 

177# Global variables to track initialization 

178_SESSION_MANAGERS_INITIALIZED = False 

179_INITIALIZATION_LOCK: Final = asyncio.Lock() 

180 

181 

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. 

185 

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 

234 

235 

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``. 

239 

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 

258 

259 

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 ) 

268 

269 return set_mcp_message_trace_carrier(carrier) 

270 except ImportError: 

271 return None 

272 

273 

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 ) 

283 

284 reset_mcp_message_trace_carrier(token) 

285 except ImportError: 

286 return 

287 

288 

289def _otel_publish_transport_span_on_scope(scope: Scope) -> None: 

290 """Record this request's tracing span on its own ASGI scope. 

291 

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. 

297 

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. 

304 

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 ) 

313 

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 

319 

320 

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) 

327 

328 

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) 

332 

333 

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 ) 

344 

345 return set_mcp_message_transport_span(span) 

346 except ImportError: 

347 return None 

348 

349 

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 ) 

358 

359 reset_mcp_message_transport_span(token) 

360 except ImportError: 

361 return 

362 

363 

364def _otel_publish_request_destinations_on_scope(scope: Scope) -> None: 

365 try: 

366 from litellm.integrations.otel.plumbing.context import request_destinations 

367 

368 scope[_MCP_DESTINATIONS_SCOPE_KEY] = request_destinations() 

369 except ImportError: 

370 return 

371 

372 

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 

380 

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 

389 

390 

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 

396 

397 reset_request_destinations(token) 

398 except ImportError: 

399 return 

400 

401 

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. 

405 

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 ) 

422 

423 

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 

479 

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 ) 

500 

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 ) 

507 

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 

527 

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 ) 

535 

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 

545 

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 

577 

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

587 

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 ) 

595 

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 

619 

620 class _TerminableTransport(Protocol): 

621 async def terminate(self) -> None: ... 621 ↛ exitline 621 didn't return from function 'terminate' because

622 

623 class _TransportRegistry(Protocol): 

624 def __contains__(self, session_id: object, /) -> bool: ... 624 ↛ exitline 624 didn't return from function '__contains__' because

625 

626 def pop(self, session_id: str, default: None, /) -> "_TerminableTransport | None": ... 626 ↛ exitline 626 didn't return from function 'pop' because

627 

628 def _stateful_server_instances() -> _TransportRegistry: 

629 return getattr(session_manager_stateful, "_server_instances", {}) 

630 

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) 

638 

639 # Keep this alias so existing references to session_manager still work 

640 session_manager: Final = session_manager_stateless 

641 

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 

646 

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) 

659 

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) 

677 

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) 

682 

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. 

687 

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

695 

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 ] 

702 

703 owned: Final = _owned_live_session_ids() 

704 if len(owned) < _MAX_STATEFUL_SESSIONS_PER_OWNER: 

705 return True 

706 

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) 

719 

720 return len(_owned_live_session_ids()) < _MAX_STATEFUL_SESSIONS_PER_OWNER 

721 

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) 

729 

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 

737 

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 

742 

743 verbose_logger.info("Initializing MCP session managers...") 

744 

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

748 

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

753 

754 _SESSION_MANAGERS_INITIALIZED = True 

755 verbose_logger.info("MCP Server started with StreamableHTTP and SSE session managers!") 

756 

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 

764 

765 if _SESSION_MANAGERS_INITIALIZED: 

766 verbose_logger.info("Shutting down MCP session managers...") 

767 

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) 

779 

780 _session_manager_cm = None 

781 _session_manager_stateful_cm = None 

782 _stateful_auth_context_cleanup_task = None 

783 _SESSION_MANAGERS_INITIALIZED = False 

784 

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

793 

794 ######################################################## 

795 ############### MCP Server Routes ####################### 

796 ######################################################## 

797 

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 ) 

823 

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=[]) 

837 

838 def _capture_host_progress_callback(ctx: ServerRequestContext) -> Callable | None: 

839 """Return a progress-forwarding callback bound to the host MCP session. 

840 

841 Returns ``None`` when the host did not supply a progress token. 

842 """ 

843 host_ctx: Final = ctx 

844 

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 

851 

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) 

863 

864 verbose_logger.debug("Host progressToken captured: %s...", str(host_token)[:8]) 

865 return forward_progress 

866 

867 def _reject_mcp_proxy_operation() -> NoReturn: 

868 from mcp.shared.exceptions import MCPError 

869 from mcp.types import METHOD_NOT_FOUND 

870 

871 raise MCPError(code=METHOD_NOT_FOUND, message="Operation unavailable on /mcp/proxy") 

872 

873 from litellm.proxy._experimental.mcp_server.operations import ( 

874 _build_virtual_call_logging_obj, 

875 _dispatch_virtual_mcp_tool, 

876 ) 

877 

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 ) 

883 

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=[]) 

895 

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 ) 

903 

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=[]) 

915 

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=[]) 

929 

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 ) 

937 

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) 

945 

946 ######################################################## 

947 ############ End of MCP Server Routes ################## 

948 ######################################################## 

949 

950 ######################################################## 

951 ############ Helper Functions ########################## 

952 ######################################################## 

953 

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 ) 

972 

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) 

1016 

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 ) 

1044 

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 

1050 

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

1057 

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) 

1064 

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 

1090 

1091 def _load_mcp_client_allowlist() -> MCPClientAllowlist | None: 

1092 from litellm.proxy.proxy_server import general_settings 

1093 

1094 return load_mcp_client_allowlist(general_settings) 

1095 

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) 

1109 

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 ) 

1143 

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 

1155 

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. 

1165 

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. 

1170 

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

1183 

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 

1193 

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" 

1211 

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 

1224 

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 

1230 

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 ) 

1242 

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 ) 

1259 

1260 def get_mcp_gateway_sessions_report(now: float | None = None) -> MCPGatewaySessionsResponse: 

1261 """Live stateful Streamable HTTP sessions held by this worker process. 

1262 

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 ) 

1281 

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 

1294 

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] 

1302 

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 

1312 

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. 

1319 

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 ) 

1347 

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``). 

1356 

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 

1367 

1368 while True: 

1369 message = await receive() 

1370 consumed_messages.append(message) 

1371 

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 

1374 

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) 

1387 

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 

1390 

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 

1395 

1396 return consumed_messages, b"".join(body_chunks) 

1397 

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. 

1409 

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. 

1414 

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", []) 

1419 

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 

1426 

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 

1435 

1436 if _session_id is None: 

1437 return False 

1438 

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 

1444 

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 

1455 

1456 # --- Session not in this worker's memory --- 

1457 method: Final = scope.get("method", "").upper() 

1458 

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 

1471 

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 

1482 

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 

1491 

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. 

1498 

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. 

1501 

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 

1508 

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 ) 

1519 

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 ) 

1534 

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

1557 

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

1572 

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 

1609 

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 

1623 

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 ) 

1635 

1636 request = StarletteRequest(scope) 

1637 base_url = get_request_base_url(request) 

1638 _path = get_route_relative_request_path(scope) 

1639 

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}"' 

1648 

1649 raise HTTPException( 

1650 status_code=401, 

1651 detail="Unauthorized", 

1652 headers={"www-authenticate": authorization_uri}, 

1653 ) 

1654 

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 

1675 

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 ) 

1688 

1689 raise_token_exchange_challenge(server, root_path=get_request_root_path()) 

1690 

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 ) 

1715 

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 ) 

1738 

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 ) 

1755 

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 ) 

1781 

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 

1789 

1790 def _scope_has_authorization_header(scope: Scope) -> bool: 

1791 return _get_authorization_header_from_scope(scope) is not None 

1792 

1793 def _get_forwarded_auth_from_scope(scope: Scope) -> str | None: 

1794 """Return the upstream-bound ``Authorization`` header value, or None. 

1795 

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) 

1809 

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. 

1816 

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. 

1821 

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 

1863 

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. 

1871 

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. 

1875 

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. 

1881 

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 

1886 

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 

1910 

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 ) 

1937 

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 

1968 

1969 # Extract client IP for MCP access control 

1970 _client_ip: Final = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) 

1971 

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 ) 

1976 

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"] 

1980 

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

1989 

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 ) 

2005 

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) 

2010 

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 ) 

2025 

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) 

2031 

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] = [] 

2039 

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 

2065 

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) 

2075 

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) 

2080 

2081 use_stateful: Final = bool(session_id or is_initialize) 

2082 target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless 

2083 

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 ) 

2089 

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 

2108 

2109 # Replay body messages if we consumed them for peeking 

2110 original_receive: Final = receive 

2111 if consumed_messages: 

2112 

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

2117 

2118 receive = wrapped_receive 

2119 

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 ) 

2171 

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

2175 

2176 active_request_session_ids: Final[list[str]] = [] 

2177 

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 ) 

2185 

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) 

2188 

2189 def _track_initialized_stateful_session( 

2190 initialized_session_id: str, 

2191 ) -> None: 

2192 _increment_active_request_session(initialized_session_id) 

2193 

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 ) 

2218 

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) 

2229 

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) 

2243 

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

2246 

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 

2272 

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 

2282 

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 

2303 

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 

2315 

2316 # Extract client IP for MCP access control 

2317 _sse_client_ip: Final = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) 

2318 

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 ) 

2323 

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"] 

2327 

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

2337 

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 ) 

2353 

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 ) 

2370 

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 

2379 

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 

2412 

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 

2422 

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 ) 

2429 

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} 

2440 

2441 class _LegacySseEndpoint: 

2442 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: 

2443 await handle_sse_mcp(scope, receive, send) 

2444 

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

2452 

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) 

2458 

2459 ######################################################## 

2460 ############ Auth Context Functions #################### 

2461 ######################################################## 

2462 

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 

2480 

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. 

2492 

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 

2511 

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 ) 

2559 

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) 

2584 

2585 return wrapped_send 

2586 

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. 

2598 

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 

2615 

2616 def _get_current_session(): 

2617 ctx: Final = get_active_mcp_request_ctx() 

2618 return ctx.session if ctx is not None else None 

2619 

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 

2633 

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 ) 

2643 

2644 def _recover_auth_from_session() -> MCPAuthenticatedUser | None: 

2645 session: Final = _get_current_session() 

2646 if session is None: 

2647 return None 

2648 

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 ) 

2658 

2659 return stored 

2660 

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

2683 

2684 if user_api_key_auth is not None: 

2685 _cache_auth_context_lazily() 

2686 else: 

2687 stored: Final = _recover_auth_from_session() 

2688 

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 ) 

2706 

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

2713 

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 

2719 

2720 stored: Final = _recover_auth_from_session() 

2721 if stored is not None: 

2722 return stored 

2723 return None 

2724 

2725 ######################################################## 

2726 ############ End of Auth Context Functions ############# 

2727 ######################################################## 

2728 

2729else: 

2730 app = FastAPI()