diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index bbf40f6e9ef..a72e8e34a49 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1212,11 +1212,17 @@ class MCPServerManager: return [] # Get server-specific auth header if available - server_auth_header = None - if mcp_server_auth_headers and server.alias: - server_auth_header = mcp_server_auth_headers.get(server.alias) - elif mcp_server_auth_headers and server.server_name: - server_auth_header = mcp_server_auth_headers.get(server.server_name) + server_auth_header: Optional[Union[str, Dict[str, str]]] = None + if mcp_server_auth_headers: + from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + ) + + server_auth_header = lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=server.alias, + server_name=server.server_name, + ) # Fall back to deprecated mcp_auth_header if no server-specific header found if server_auth_header is None: @@ -2707,16 +2713,15 @@ class MCPServerManager: server_auth_header: Optional[Union[Dict[str, str], str]] = None if mcp_server_auth_headers: # Normalize keys for case-insensitive lookup - normalized_headers = { - k.lower(): v for k, v in mcp_server_auth_headers.items() - } + from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + ) - if mcp_server.alias: - server_auth_header = normalized_headers.get(mcp_server.alias.lower()) - if server_auth_header is None and mcp_server.server_name: - server_auth_header = normalized_headers.get( - mcp_server.server_name.lower() - ) + server_auth_header = lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=mcp_server.alias, + server_name=mcp_server.server_name, + ) # Fall back to deprecated mcp_auth_header if no server-specific header found if server_auth_header is None: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 09176f7253a..e8b591c39cf 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -3,8 +3,8 @@ import os from ipaddress import ip_address -from typing import List, Optional -from urllib.parse import urlparse, urlunparse +from typing import Any, Dict, List, NoReturn, Optional +from urllib.parse import ParseResult, urlparse, urlunparse from fastapi import HTTPException, Request @@ -43,6 +43,33 @@ _DEFAULT_NATIVE_REDIRECT_URIS: List[str] = [ _warned_invalid_proxy_base_url: Optional[str] = None +def _oauth_invalid_request( + error_description: str, + *, + hint: Optional[str] = None, + **extra: Any, +) -> NoReturn: + """Raise ``invalid_request`` (RFC 6749) with a debuggable description. + + FastAPI serializes ``detail`` as JSON. Callers still see ``error``: + ``invalid_request``; ``error_description`` and ``hint`` explain what + failed and how to fix it (e.g. reverse-proxy / PROXY_BASE_URL issues). + """ + detail: Dict[str, Any] = { + "error": "invalid_request", + "error_description": error_description, + } + if hint: + detail["hint"] = hint + detail.update(extra) + raise HTTPException(status_code=400, detail=detail) + + +def _origin_label(scheme: str, netloc: str) -> str: + """Human-readable origin for error messages (scheme + host[:port]).""" + return f"{scheme}://{netloc}" if netloc else f"{scheme}://" + + def _resolve_proxy_base_url_env() -> Optional[str]: global _warned_invalid_proxy_base_url configured = os.environ.get("PROXY_BASE_URL", "").strip() @@ -118,17 +145,15 @@ def validate_loopback_redirect_uri(redirect_uri: str) -> None: ``"127.0.0.1"`` alone would miss ``127.0.0.2`` and the full-form IPv6 loopback ``0:0:0:0:0:0:0:1``. """ - try: - parsed = urlparse(redirect_uri) - except ValueError: - raise HTTPException(status_code=400, detail="invalid_request") + parsed = _parse_redirect_uri_for_validation(redirect_uri) if parsed.scheme not in ("http", "https"): - raise HTTPException(status_code=400, detail="invalid_request") - # Fragments are not allowed in OAuth redirect URIs (RFC 6749 §3.1.2) - # — rejecting them prevents a ``http://127.0.0.1/cb#frag?code=...`` - # from silently eating the authorization code. + _oauth_invalid_request( + f"redirect_uri scheme {parsed.scheme!r} is not allowed; use http or https.", + ) if parsed.fragment: - raise HTTPException(status_code=400, detail="invalid_request") + _oauth_invalid_request( + "redirect_uri must not contain a URL fragment (#...).", + ) host = (parsed.hostname or "").lower() if host == "localhost": return @@ -139,7 +164,10 @@ def validate_loopback_redirect_uri(redirect_uri: str) -> None: # Unparseable host (malformed IPv6, etc.) — treat as invalid, # don't let it bubble up as a 500. pass - raise HTTPException(status_code=400, detail="invalid_request") + _oauth_invalid_request( + "redirect_uri must use a loopback host (localhost or 127.0.0.0/8).", + hint="Native MCP clients should register a callback on http://127.0.0.1:/...", + ) def _strip_default_port(scheme: str, netloc: str) -> str: @@ -293,6 +321,180 @@ def _matches_trusted_native_redirect_uri(parsed) -> bool: return False +def _parse_redirect_uri_for_validation(redirect_uri: str) -> ParseResult: + try: + return urlparse(redirect_uri) + except ValueError: + _oauth_invalid_request( + "redirect_uri is not a valid URL.", + hint="Use a full absolute URL for redirect_uri (e.g. https://your-host/ui/mcp/oauth/callback).", + ) + + +def _validate_trusted_http_redirect_shape(parsed: ParseResult) -> bool: + """Return True when ``parsed`` is an allowlisted native callback (caller may return).""" + if parsed.scheme not in ("http", "https"): + if _matches_trusted_native_redirect_uri(parsed): + return True + _oauth_invalid_request( + f"redirect_uri scheme {parsed.scheme!r} is not allowed; use http/https " + "or a registered native callback (e.g. cursor://).", + hint="Add the full URI to MCP_TRUSTED_NATIVE_REDIRECT_URIS for custom native clients.", + ) + if parsed.fragment: + _oauth_invalid_request( + "redirect_uri must not contain a URL fragment (#...).", + ) + if not parsed.netloc: + _oauth_invalid_request( + "redirect_uri must include a host (e.g. https://your-host/path).", + ) + if parsed.username is not None or parsed.password is not None: + _oauth_invalid_request( + "redirect_uri must not contain userinfo (user:pass@host).", + ) + if "\\" in parsed.netloc: + _oauth_invalid_request( + "redirect_uri host must not contain backslashes.", + ) + return False + + +def _resolve_proxy_base_for_redirect(request: Request) -> Optional[str]: + try: + return get_request_base_url(request) + except Exception as exc: + verbose_logger.warning( + "validate_trusted_redirect_uri: could not determine proxy origin, " + "falling back to loopback + allowlist. error=%s", + exc, + ) + return None + + +def _trusted_redirect_uri_is_allowed( + parsed: ParseResult, + redirect_netloc: str, + proxy_base: Optional[str], +) -> bool: + if proxy_base: + proxy_parsed = urlparse(proxy_base) + if ( + parsed.scheme == proxy_parsed.scheme + and redirect_netloc + == _strip_default_port(proxy_parsed.scheme, proxy_parsed.netloc) + ): + return True + + host = (parsed.hostname or "").lower() + if host == "localhost": + return True + try: + if ip_address(host).is_loopback: + return True + except ValueError: + pass + + if parsed.scheme == "https": + for entry in _parse_trusted_redirect_origins(): + if _matches_trusted_origin_entry(redirect_netloc, entry): + return True + return False + + +def _build_trusted_redirect_rejection_message( + redirect_uri: str, + parsed: ParseResult, + redirect_netloc: str, + proxy_base: Optional[str], +) -> str: + """Build a client-facing rejection message. + + Intentionally omits the proxy's resolved scheme / host / port to avoid + leaking internal network topology (e.g. ``http://litellm-internal:4000``) + through an unauthenticated endpoint. Full diagnostic detail — including + the computed proxy base — is logged server-side by the caller. + """ + redirect_origin = _origin_label(parsed.scheme, redirect_netloc) + proxy_parsed = urlparse(proxy_base) if proxy_base else None + proxy_netloc_norm = ( + _strip_default_port(proxy_parsed.scheme, proxy_parsed.netloc) + if proxy_parsed and proxy_parsed.netloc + else "" + ) + + mismatch_parts: List[str] = [] + if proxy_parsed and proxy_parsed.netloc: + if parsed.scheme != proxy_parsed.scheme: + mismatch_parts.append( + f"scheme: redirect_uri uses {parsed.scheme!r}, but the proxy " + "resolved a different scheme " + "(TLS often terminates at ingress — set PROXY_BASE_URL to https://… " + "or trust X-Forwarded-Proto from your ingress)" + ) + if redirect_netloc != proxy_netloc_norm: + mismatch_parts.append( + f"host/port: redirect_uri {redirect_netloc!r} does not match " + "the proxy origin" + ) + + if mismatch_parts: + return ( + f"redirect_uri origin ({redirect_origin}) does not match the proxy " + "origin. " + "; ".join(mismatch_parts) + ) + return ( + f"redirect_uri ({redirect_uri!r}) is not allowed: not same-origin with " + f"the proxy origin, not loopback, and not listed in " + f"{_TRUSTED_REDIRECT_ORIGINS_ENV}." + ) + + +def _raise_trusted_redirect_uri_rejected( + request: Request, + redirect_uri: str, + parsed: ParseResult, + redirect_netloc: str, + proxy_base: Optional[str], +) -> NoReturn: + description = _build_trusted_redirect_rejection_message( + redirect_uri, parsed, redirect_netloc, proxy_base + ) + + hint = ( + "Align the proxy public URL with the browser URL. Set PROXY_BASE_URL to your " + "HTTPS origin (e.g. https://litellm.example.com), or enable " + "general_settings.use_x_forwarded_for with mcp_trusted_proxy_ranges for your " + "ingress. Verify: curl https:///.well-known/oauth-authorization-server " + "| jq .issuer — issuer must match window.location.origin in the UI." + ) + + verbose_logger.warning( + "MCP OAuth: rejecting redirect_uri %r. %s " + "Computed proxy base=%r (PROXY_BASE_URL=%r). " + "Inbound headers: X-Forwarded-Proto=%r X-Forwarded-Host=%r " + "X-Forwarded-Port=%r Host=%r. " + "Trusted-redirect-origins env=%r. " + "Trusted-native-redirect-uris env=%r.", + redirect_uri, + description, + proxy_base, + os.environ.get("PROXY_BASE_URL"), + request.headers.get("X-Forwarded-Proto"), + request.headers.get("X-Forwarded-Host"), + request.headers.get("X-Forwarded-Port"), + request.headers.get("Host"), + os.environ.get(_TRUSTED_REDIRECT_ORIGINS_ENV), + os.environ.get(_TRUSTED_NATIVE_REDIRECT_URIS_ENV), + ) + + _oauth_invalid_request( + description, + hint=hint, + redirect_uri=redirect_uri, + ) + + def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None: """Accept ``redirect_uri`` when it is (a) same-origin with the proxy's own request origin, (b) loopback, (c) listed in the @@ -316,98 +518,13 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None: BYOK endpoints, which only serve native MCP clients, retain :func:`validate_loopback_redirect_uri`. """ - try: - parsed = urlparse(redirect_uri) - except ValueError: - raise HTTPException(status_code=400, detail="invalid_request") - if parsed.scheme not in ("http", "https"): - if _matches_trusted_native_redirect_uri(parsed): - return - raise HTTPException(status_code=400, detail="invalid_request") - if parsed.fragment: - raise HTTPException(status_code=400, detail="invalid_request") - if not parsed.netloc or parsed.username is not None or parsed.password is not None: - raise HTTPException(status_code=400, detail="invalid_request") - # Reject userinfo (``user:pass@host``) outright: OAuth redirect_uris - # have no legitimate reason to carry credentials, and allowing them - # opens a host-confusion attack where the netloc *looks* allowlisted - # (``app.example.com:443@attacker.example``) but the browser navigates - # to the post-``@`` host and hands the authorization code to the - # attacker. We compare against ``hostname`` after this, but defense in - # depth keeps malformed netloc strings from reaching the wildcard - # splitter. - if parsed.username is not None or parsed.password is not None: - raise HTTPException(status_code=400, detail="invalid_request") - # Reject backslash in netloc: urlparse keeps ``\`` as part of netloc, - # but browsers normalize ``\`` to ``/`` for http(s) URLs and treat it - # as the start of the path. An attacker can exploit that split by - # crafting ``https://attacker.net\app.example.com/cb`` — urlparse sees - # ``attacker.net\app.example.com`` (matches ``*.example.com``) while - # the browser navigates to ``attacker.net`` with the auth code. - if "\\" in parsed.netloc: - raise HTTPException(status_code=400, detail="invalid_request") - - redirect_netloc = _strip_default_port(parsed.scheme, parsed.netloc) - - # (a) Same-origin. Swallow ``get_request_base_url`` failures so the - # loopback + allowlist paths remain reachable when the origin can't - # be determined (e.g. request came from an untrusted proxy and - # ``get_request_base_url`` raised). - proxy_base: Optional[str] = None - try: - proxy_base = get_request_base_url(request) - except Exception as exc: - verbose_logger.warning( - "validate_trusted_redirect_uri: could not determine proxy origin, " - "falling back to loopback + allowlist. error=%s", - exc, - ) - proxy_base = None - if proxy_base: - proxy_parsed = urlparse(proxy_base) - if ( - parsed.scheme == proxy_parsed.scheme - and redirect_netloc - == _strip_default_port(proxy_parsed.scheme, proxy_parsed.netloc) - ): - return - - # (b) Loopback — same rule as validate_loopback_redirect_uri. - host = (parsed.hostname or "").lower() - if host == "localhost": + parsed = _parse_redirect_uri_for_validation(redirect_uri) + if _validate_trusted_http_redirect_shape(parsed): return - try: - if ip_address(host).is_loopback: - return - except ValueError: - pass - - # (c) Ops allowlist. https only. - if parsed.scheme == "https": - for entry in _parse_trusted_redirect_origins(): - if _matches_trusted_origin_entry(redirect_netloc, entry): - return - - verbose_logger.warning( - "MCP OAuth: rejecting redirect_uri %r as invalid_request. " - "Computed proxy base=%r (PROXY_BASE_URL=%r). " - "Inbound headers: X-Forwarded-Proto=%r X-Forwarded-Host=%r " - "X-Forwarded-Port=%r Host=%r. " - "Trusted-redirect-origins env=%r. " - "Trusted-native-redirect-uris env=%r. " - "If this should be accepted, either align ingress X-Forwarded-* " - "with the browser URL, set PROXY_BASE_URL to your public origin, " - "add the redirect_uri host to MCP_TRUSTED_REDIRECT_ORIGINS, or " - "for native MCP clients (cursor://, etc.) add the full redirect_uri " - "to MCP_TRUSTED_NATIVE_REDIRECT_URIS.", - redirect_uri, - proxy_base, - os.environ.get("PROXY_BASE_URL"), - request.headers.get("X-Forwarded-Proto"), - request.headers.get("X-Forwarded-Host"), - request.headers.get("X-Forwarded-Port"), - request.headers.get("Host"), - os.environ.get(_TRUSTED_REDIRECT_ORIGINS_ENV), - os.environ.get(_TRUSTED_NATIVE_REDIRECT_URIS_ENV), + redirect_netloc = _strip_default_port(parsed.scheme, parsed.netloc) + proxy_base = _resolve_proxy_base_for_redirect(request) + if _trusted_redirect_uri_is_allowed(parsed, redirect_netloc, proxy_base): + return + _raise_trusted_redirect_uri_rejected( + request, redirect_uri, parsed, redirect_netloc, proxy_base ) - raise HTTPException(status_code=400, detail="invalid_request") diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7150dee10cf..cec5224e183 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -62,20 +62,16 @@ if MCP_AVAILABLE: mcp_auth_header: Optional[str], ) -> Optional[Union[Dict[str, str], str]]: """Helper function to get server-specific auth header with case-insensitive matching.""" - if mcp_server_auth_headers and server.alias: - normalized_server_alias = server.alias.lower() - normalized_headers = { - k.lower(): v for k, v in mcp_server_auth_headers.items() - } - server_auth = normalized_headers.get(normalized_server_alias) - if server_auth is not None: - return server_auth - elif mcp_server_auth_headers and server.server_name: - normalized_server_name = server.server_name.lower() - normalized_headers = { - k.lower(): v for k, v in mcp_server_auth_headers.items() - } - server_auth = normalized_headers.get(normalized_server_name) + from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + ) + + if mcp_server_auth_headers: + server_auth = lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=getattr(server, "alias", None), + server_name=getattr(server, "server_name", None), + ) if server_auth is not None: return server_auth return mcp_auth_header diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 5676aaf0d22..5205426edf3 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1114,10 +1114,16 @@ if MCP_AVAILABLE: ) -> Tuple[Optional[Union[Dict[str, str], str]], Optional[Dict[str, str]]]: """Build auth and extra headers for a server.""" server_auth_header: Optional[Union[Dict[str, str], str]] = None - if mcp_server_auth_headers and server.alias is not None: - server_auth_header = mcp_server_auth_headers.get(server.alias) - elif mcp_server_auth_headers and server.server_name is not None: - server_auth_header = mcp_server_auth_headers.get(server.server_name) + if mcp_server_auth_headers: + from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + ) + + server_auth_header = lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=server.alias, + server_name=server.server_name, + ) extra_headers: Optional[Dict[str, str]] = None if server.auth_type == MCPAuth.oauth2: diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index df5705c3425..b8b9207555e 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -2,7 +2,8 @@ MCP Server Utilities """ -from typing import Any, Dict, Iterator, Mapping, Optional, Tuple +import re +from typing import Any, Dict, Iterator, Mapping, Optional, Tuple, Union import hashlib import importlib @@ -117,6 +118,50 @@ def normalize_server_name(server_name: str) -> str: return server_name.replace(" ", "_") +_MCP_ALIAS_HEADER_INVALID_RE = re.compile(r"[^a-z0-9_]") + + +def sanitize_mcp_alias_for_header(alias: str) -> str: + """ + Sanitize an MCP server alias for x-mcp-{alias}-{header} HTTP headers. + + Must stay in sync with ui/litellm-dashboard/src/utils/mcpHeaderUtils.ts. + """ + sanitized = _MCP_ALIAS_HEADER_INVALID_RE.sub("_", alias.lower().strip()) + sanitized = re.sub(r"_+", "_", sanitized) + return sanitized.strip("_") + + +def lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers: Mapping[str, Union[str, Dict[str, str]]], + *, + alias: Optional[str] = None, + server_name: Optional[str] = None, +) -> Optional[Union[str, Dict[str, str]]]: + """ + Resolve server-specific auth headers with case-insensitive matching. + + Tries the raw alias/server_name (lowercased) and the header-safe sanitized + alias so dashboard clients using sanitize_mcp_alias_for_header() still match. + """ + if not mcp_server_auth_headers: + return None + + normalized_headers = {k.lower(): v for k, v in mcp_server_auth_headers.items()} + + for identifier in (alias, server_name): + if not identifier: + continue + keys_to_try = [identifier.lower()] + sanitized = sanitize_mcp_alias_for_header(identifier) + if sanitized and sanitized not in keys_to_try: + keys_to_try.append(sanitized) + for key in keys_to_try: + if key in normalized_headers: + return normalized_headers[key] + return None + + def validate_and_normalize_mcp_server_payload(payload: Any) -> None: """ Validate and normalize MCP server payload fields (server_name and alias). diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 809b13aeea6..c20fb09eeba 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -1770,6 +1770,26 @@ def test_get_server_auth_header_fallback_to_default(): assert result == "Bearer default_token" +def test_get_server_auth_header_hyphenated_alias_sanitized_header_key(): + """Header keys use sanitized alias; lookup must match legacy hyphenated aliases.""" + from litellm.proxy._experimental.mcp_server.rest_endpoints import ( + _get_server_auth_header, + ) + + mock_server = MagicMock() + mock_server.alias = "GitHub-MCP" + mock_server.server_name = "github_mcp_server" + + mcp_server_auth_headers = { + "github_mcp": {"Authorization": "Bearer github-mcp-token"}, + } + + result = _get_server_auth_header( + mock_server, mcp_server_auth_headers, "Bearer default_token" + ) + assert result == {"Authorization": "Bearer github-mcp-token"} + + def test_get_server_auth_header_no_auth_headers(): """Test _get_server_auth_header function with no auth headers.""" from litellm.proxy._experimental.mcp_server.rest_endpoints import ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index b06cc7f0f12..c8789e0b0a6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1345,7 +1345,13 @@ def test_validate_trusted_redirect_uri_logs_diagnostic_on_rejection( "https://litellm.example.com/ui/mcp/oauth/callback", ) assert exc_info.value.status_code == 400 - assert exc_info.value.detail == "invalid_request" + detail = exc_info.value.detail + assert isinstance(detail, dict) + assert detail.get("error") == "invalid_request" + assert "error_description" in detail + assert "redirect_uri origin" in detail["error_description"] + assert "proxy origin" in detail["error_description"] + assert "hint" in detail matching = [r for r in caplog.records if "rejecting redirect_uri" in r.getMessage()] assert len(matching) == 1, ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_header_alias_utils.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_header_alias_utils.py new file mode 100644 index 00000000000..2627199570b --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_header_alias_utils.py @@ -0,0 +1,18 @@ +"""Tests for MCP header alias sanitization and auth header lookup.""" + +from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + sanitize_mcp_alias_for_header, +) + + +def test_sanitize_mcp_alias_for_header(): + assert sanitize_mcp_alias_for_header("My Server") == "my_server" + assert sanitize_mcp_alias_for_header("GitHub-MCP!") == "github_mcp" + assert sanitize_mcp_alias_for_header("github_mcp2") == "github_mcp2" + + +def test_lookup_mcp_server_auth_in_headers_sanitized_alias(): + headers = {"github_mcp": {"Authorization": "Bearer token"}} + result = lookup_mcp_server_auth_in_headers(headers, alias="GitHub-MCP") + assert result == {"Authorization": "Bearer token"} diff --git a/ui/litellm-dashboard/src/app/mcp/oauth/callback/page.tsx b/ui/litellm-dashboard/src/app/mcp/oauth/callback/page.tsx index 0539d6d8f19..3b3729c1ac9 100644 --- a/ui/litellm-dashboard/src/app/mcp/oauth/callback/page.tsx +++ b/ui/litellm-dashboard/src/app/mcp/oauth/callback/page.tsx @@ -4,11 +4,12 @@ import { Suspense, useEffect, useMemo } from "react"; import { useSearchParams } from "next/navigation"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; -// Written to sessionStorage so both the admin hook (useMcpOAuthFlow) and the -// user hook (useUserMcpOAuthFlow) can pick up the result. Each hook reads -// its own namespace to avoid cross-flow collisions. +// Written to sessionStorage so the admin hook (useMcpOAuthFlow), the user hook +// (useUserMcpOAuthFlow), and the tools re-auth hook (useToolsOAuthFlow) can each +// pick up the result. Each hook reads its own namespace to avoid cross-flow collisions. const ADMIN_RESULT_KEY = "litellm-mcp-oauth-result"; const USER_RESULT_KEY = "litellm-user-mcp-oauth-result"; +const TOOLS_RESULT_KEY = "litellm-tools-mcp-oauth-result"; const RETURN_URL_STORAGE_KEY = "litellm-mcp-oauth-return-url"; const resolveDefaultRedirect = () => { @@ -50,11 +51,12 @@ const McpOAuthCallbackContent = () => { } try { - // Write to both namespace keys (admin and user) so whichever hook is - // active can consume the result. sessionStorage only — no localStorage. + // Write to all namespace keys so whichever hook is active can consume + // the result. sessionStorage only — no localStorage. const serialized = JSON.stringify(payload); setSecureItem(ADMIN_RESULT_KEY, serialized); setSecureItem(USER_RESULT_KEY, serialized); + setSecureItem(TOOLS_RESULT_KEY, serialized); } catch (err) { // Silently ignore storage errors } diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index f8b0141b25d..108911bdbf1 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -3,6 +3,7 @@ import { Modal, Tooltip, Form, Select, Input, Switch, Collapse } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, TextInput } from "@tremor/react"; import { createMCPServer, registerMCPServer } from "../networking"; +import { setToken } from "@/utils/mcpTokenStore"; import { AUTH_TYPE, DiscoverableMCPServer, OAUTH_FLOW, MCPServer, MCPServerCostInfo, TRANSPORT } from "./types"; import OAuthFormFields from "./OAuthFormFields"; import MCPServerCostConfig from "./mcp_server_cost_config"; @@ -24,6 +25,7 @@ export const mcpLogoImg = `${asset_logos_folder}mcp_logo.png`; interface CreateMCPServerProps { userRole: string; + userID?: string | null; accessToken: string | null; onCreateSuccess: (newMcpServer: MCPServer) => void; isModalVisible: boolean; @@ -47,6 +49,7 @@ const reduceStaticHeaders = (list: unknown): Record => { }; const CreateMCPServer: React.FC = ({ + userID, userRole, accessToken, onCreateSuccess, @@ -409,6 +412,21 @@ const CreateMCPServer: React.FC = ({ ? await createMCPServer(accessToken, payload) : await registerMCPServer(accessToken, payload); + // Cache the OAuth token in sessionStorage so the Tools tab can use it + // immediately without re-authenticating. No backend DB write. + if (oauthTokenResponse?.access_token && response?.server_id) { + setToken( + response.server_id, + { + access_token: oauthTokenResponse.access_token, + expires_in: oauthTokenResponse.expires_in, + refresh_token: oauthTokenResponse.refresh_token, + token_type: oauthTokenResponse.token_type, + }, + userID, + ); + } + NotificationsManager.success( isAdmin ? "MCP Server created successfully" diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index 1f8f7f68d33..5a8035d4e0b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -174,6 +174,7 @@ export const MCPServerView: React.FC = ({ serverId={mcpServer.server_id} accessToken={accessToken} auth_type={mcpServer.auth_type} + tokenUrl={mcpServer.token_url} userRole={userRole} userID={userID} serverAlias={mcpServer.alias} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx index 72d5e4b5aa8..42583fdab07 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -287,6 +287,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) (null); const [toolError, setToolError] = useState(null); const [toolSearchTerm, setToolSearchTerm] = useState(""); - + // State for passthrough headers const [passthroughHeaders, setPassthroughHeaders] = useState>({}); const [showHeaderInput, setShowHeaderInput] = useState(false); + // OAuth session token (sessionStorage-backed, cleared on tab/browser close). + // Only the interactive (authorization_code/PKCE) flow needs a user-facing + // auth gate. M2M (client_credentials) servers are also `auth_type === "oauth2"`, + // but the backend fetches their token internally — gating tool listing on + // them would force users through a non-existent authorization endpoint. + // We detect M2M via the presence of `tokenUrl`, matching the heuristic in + // `mcp_server_edit.tsx`. + const isOAuth = auth_type === "oauth2" && !tokenUrl; + const [oauthToken, setOauthToken] = useState(() => + isOAuth && isTokenValid(serverId, userID) + ? (getToken(serverId, userID)?.access_token ?? null) + : null + ); + + // Re-sync token when serverId/userID changes (useState initializer only runs on mount). + useEffect(() => { + if (!isOAuth) { + setOauthToken(null); + return; + } + setOauthToken( + isTokenValid(serverId, userID) + ? (getToken(serverId, userID)?.access_token ?? null) + : null + ); + }, [serverId, userID, isOAuth]); + + const { startOAuthFlow, status: oauthStatus, error: oauthError } = useToolsOAuthFlow({ + accessToken: accessToken ?? "", + serverId, + serverAlias, + userId: userID, + onSuccess: setOauthToken, + }); + // Check if this server has extra headers configured const hasExtraHeaders = extraHeaders && extraHeaders.length > 0; // Build custom headers for MCP server requests const buildCustomHeaders = () => { - if (!serverAlias || !hasExtraHeaders) return undefined; - const customHeaders: Record = {}; - - // Add passthrough headers with server-specific prefix - Object.entries(passthroughHeaders).forEach(([headerName, headerValue]) => { - if (headerValue && headerValue.trim()) { - // Format: x-mcp-{alias}-{header_name} - const mcpHeaderName = `x-mcp-${serverAlias}-${headerName.toLowerCase()}`; - customHeaders[mcpHeaderName] = headerValue; + + // Include the session OAuth token using MCP-specific headers so it doesn't + // conflict with the Authorization header used by the LiteLLM proxy itself. + // The backend's _get_mcp_server_auth_headers_from_headers() picks up the + // x-mcp-{alias}-{header} pattern and forwards it to the upstream MCP server. + // When no alias is available, fall back to x-mcp-auth (legacy but still supported). + if (oauthToken) { + if (serverAlias) { + const safeAlias = sanitizeMcpAliasForHeader(serverAlias); + if (safeAlias) { + customHeaders[`x-mcp-${safeAlias}-authorization`] = `Bearer ${oauthToken}`; + } else { + customHeaders["x-mcp-auth"] = `Bearer ${oauthToken}`; + } + } else { + customHeaders["x-mcp-auth"] = `Bearer ${oauthToken}`; } - }); - + } + + // Add passthrough headers with server-specific prefix + if (serverAlias && hasExtraHeaders) { + const safeAlias = sanitizeMcpAliasForHeader(serverAlias); + if (safeAlias) { + Object.entries(passthroughHeaders).forEach(([headerName, headerValue]) => { + if (headerValue && headerValue.trim()) { + // Format: x-mcp-{alias}-{header_name} + const mcpHeaderName = `x-mcp-${safeAlias}-${headerName.toLowerCase()}`; + customHeaders[mcpHeaderName] = headerValue; + } + }); + } + } + return Object.keys(customHeaders).length > 0 ? customHeaders : undefined; }; @@ -54,15 +114,55 @@ const MCPToolsViewer = ({ error: mcpToolsError, refetch: refetchTools, } = useQuery({ - queryKey: ["mcpTools", serverId, passthroughHeaders], - queryFn: () => { + queryKey: ["mcpTools", serverId, passthroughHeaders, oauthToken], + queryFn: async () => { if (!accessToken) throw new Error("Access Token required"); - return listMCPTools(accessToken, serverId, buildCustomHeaders()); + const result = await listMCPTools(accessToken, serverId, buildCustomHeaders()); + // listMCPTools never throws — surface error responses as thrown errors + // here so useQuery's retry/onError can react (e.g. clear the cached + // OAuth token on 401). + if (result?.error) { + const status = (result as { status?: number }).status; + if (status === 401) { + removeToken(serverId, userID); + } + const enhancedError = new Error( + result.message || result.error || "Failed to fetch MCP tools", + ) as Error & { + status?: number; + statusText?: string; + details?: any; + }; + enhancedError.status = status; + enhancedError.statusText = (result as any).statusText; + enhancedError.details = (result as any).details; + throw enhancedError; + } + return result; }, - enabled: !!accessToken, + // For OAuth servers, block the query until a session token is available + enabled: !!accessToken && (!isOAuth || oauthToken !== null), staleTime: 30000, // Consider data fresh for 30 seconds + retry: (failureCount, error: any) => { + // Don't retry on 401 — token is invalid, user must re-authenticate + if (error?.status === 401 || error?.response?.status === 401) return false; + return failureCount < 2; + }, }); + // If the tools query fails with 401, the cached OAuth token is invalid — + // clear it so the auth gate is shown again and the user can re-authenticate. + useEffect(() => { + const err = mcpToolsError as + | (Error & { status?: number; response?: { status?: number } }) + | null; + const status = err?.status ?? err?.response?.status; + if (status === 401) { + removeToken(serverId, userID); + setOauthToken(null); + } + }, [mcpToolsError, serverId, userID]); + // Mutation for calling a tool const { mutate: executeTool, isPending: isCallingTool } = useMutation({ mutationFn: async (args: { tool: MCPTool; arguments: Record }) => { @@ -85,9 +185,14 @@ const MCPToolsViewer = ({ setToolResult(data.content); setToolError(null); }, - onError: (error: Error) => { + onError: (error: Error & { status?: number; response?: { status?: number } }) => { setToolError(error); setToolResult(null); + // On 401, clear the cached token so the auth gate is shown again + if (error?.status === 401 || (error as any)?.response?.status === 401) { + removeToken(serverId, userID); + setOauthToken(null); + } }, }); @@ -197,7 +302,31 @@ const MCPToolsViewer = ({ )} - {/* Search Bar */} + {/* OAuth Auth Gate — shown when token is absent for OAuth servers */} + {isOAuth && !oauthToken && ( +
+ +

Authentication required

+

+ Authenticate to view available tools +

+ + Authorize + + {oauthError && ( +

{oauthError}

+ )} +
+ )} + + {/* Search Bar — only shown when tools are loaded */} + {!isOAuth || oauthToken ? <> {toolsData.length > 0 && (
-

Error: {mcpToolsResponse.message}

+

+ Error: {mcpToolsResponse?.message || (mcpToolsError as Error)?.message} +

)} {/* No Tools State */} - {!isLoadingTools && !mcpToolsResponse?.error && (!toolsData || toolsData.length === 0) && ( + {!isLoadingTools && !mcpToolsResponse?.error && !mcpToolsError && (!toolsData || toolsData.length === 0) && (
@@ -315,6 +446,7 @@ const MCPToolsViewer = ({ )} )} + : null}
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 7cfe08d9ee5..9a8f2e8f514 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -163,6 +163,13 @@ export interface MCPToolsViewerProps { serverId: string; accessToken: string | null; auth_type?: string | null; + /** + * When set, indicates the server uses the OAuth2 M2M (client_credentials) + * flow — the backend handles token acquisition internally, so the UI must + * not gate tool listing behind an interactive PKCE authorization. Mirrors + * the heuristic used in `mcp_server_edit.tsx` (`token_url` set => M2M). + */ + tokenUrl?: string | null; userRole: string | null; userID: string | null; serverAlias?: string | null; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 756348f4937..57e7d51123e 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -7068,46 +7068,33 @@ export const testSearchToolConnection = async (accessToken: string, litellmParam }; export const listMCPTools = async ( - accessToken: string, + accessToken: string, serverId: string, - customHeaders?: Record + customHeaders?: Record, ) => { + // Construct base URL + let url = proxyBaseUrl + ? `${proxyBaseUrl}/mcp-rest/tools/list?server_id=${serverId}` + : `/mcp-rest/tools/list?server_id=${serverId}`; + + console.log("Fetching MCP tools from:", url); + + const headers: Record = { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + ...customHeaders, // Merge custom headers for passthrough auth + }; + + let response: Response; try { - // Construct base URL - let url = proxyBaseUrl - ? `${proxyBaseUrl}/mcp-rest/tools/list?server_id=${serverId}` - : `/mcp-rest/tools/list?server_id=${serverId}`; - - console.log("Fetching MCP tools from:", url); - - const headers: Record = { - [globalLitellmHeaderName]: `Bearer ${accessToken}`, - "Content-Type": "application/json", - ...customHeaders, // Merge custom headers for passthrough auth - }; - - const response = await fetch(url, { + response = await fetch(url, { method: "GET", headers, }); - - const data = await response.json(); - console.log("Fetched MCP tools response:", data); - - if (!response.ok) { - // If the server returned an error response, use it - if (data.error && data.message) { - throw new Error(data.message); - } - // Otherwise use a generic error - throw new Error("Failed to fetch MCP tools"); - } - - // Return the full response object which includes tools, error, message, and stack_trace - return data; } catch (error) { - console.error("Failed to fetch MCP tools:", error); - // Return an error response in the same format as the API + // Network-level failure (no HTTP response). Preserve legacy shape so the + // caller can render a generic error message without crashing. + console.error("Failed to fetch MCP tools (network error):", error); return { tools: [], error: "network_error", @@ -7115,6 +7102,44 @@ export const listMCPTools = async ( stack_trace: null, }; } + + let data: any = null; + try { + data = await response.json(); + } catch (parseError) { + console.error("Failed to parse MCP tools response:", parseError); + return { + tools: [], + error: "parse_error", + message: "Failed to parse MCP tools response", + status: response.status, + statusText: response.statusText, + stack_trace: null, + }; + } + console.log("Fetched MCP tools response:", data); + + if (!response.ok) { + // Preserve the legacy "never throws" contract so existing callers + // (e.g. MCPToolPermissions, MCPAppsPanel, MCPConnectPicker) can continue + // to inspect `result.error` / `result.message`. Attach `status` so + // callers that need to react to auth failures (e.g. the useQuery in + // mcp_tools.tsx) can still detect 401s from the returned object. + const errorMessage = + (data && (data.message || data.error)) || "Failed to fetch MCP tools"; + return { + tools: [], + error: (data && data.error) || `http_${response.status}`, + message: errorMessage, + status: response.status, + statusText: response.statusText, + details: data, + stack_trace: null, + }; + } + + // Return the full response object which includes tools, error, message, and stack_trace + return data; }; interface CallMCPToolOptions { diff --git a/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts b/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts new file mode 100644 index 00000000000..3aff8af6eef --- /dev/null +++ b/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts @@ -0,0 +1,39 @@ +/** + * Shared utilities for MCP OAuth2 PKCE flow hooks. + * + * These helpers are used by both useToolsOAuthFlow and useUserMcpOAuthFlow + * to avoid divergence in URL construction and storage cleanup logic. + */ + +import { getProxyBaseUrl, serverRootPath } from "@/components/networking"; + +/** + * Build the OAuth callback URL for the current UI deployment. + * + * In the browser, derive the `/ui` prefix from the current pathname so the + * callback works regardless of how the proxy is mounted. Outside the browser + * (SSR), fall back to the configured proxy base URL and server root path. + */ +export const buildCallbackUrl = (): string => { + if (typeof window !== "undefined") { + const path = window.location.pathname || ""; + const idx = path.indexOf("/ui"); + const prefix = idx >= 0 ? path.slice(0, idx + 3).replace(/\/+$/, "") : ""; + return `${window.location.origin}${prefix}/mcp/oauth/callback`; + } + const base = (getProxyBaseUrl() || "").replace(/\/+$/, ""); + const root = serverRootPath && serverRootPath !== "/" ? serverRootPath : ""; + return `${base}${root}/ui/mcp/oauth/callback`; +}; + +/** + * Remove the given keys from sessionStorage, ignoring errors (e.g. storage + * disabled by browser privacy settings). + */ +export const clearStorage = (...keys: string[]): void => { + keys.forEach((k) => { + try { + window.sessionStorage.removeItem(k); + } catch (_) {} + }); +}; diff --git a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx index 7edeade4cbd..11efaba53b3 100644 --- a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx @@ -224,12 +224,21 @@ export const useMcpOAuthFlow = ({ if (!storedPayload) { return; } - + + // Guard: the callback page writes to the admin result key for *all* OAuth + // flows (including the tools re-auth flow). Only proceed if this hook's + // own flow state exists, meaning startOAuthFlow() was actually called here. + // Without this guard, a tools re-auth redirect triggers a spurious + // "OAuth session state was lost" error from this hook. + const storedFlowState = getStorageItem(FLOW_STATE_KEY); + if (!storedFlowState) { + return; + } + // Mark as processing processingRef.current = true; payload = JSON.parse(storedPayload); - const storedFlowState = getStorageItem(FLOW_STATE_KEY); - flowState = storedFlowState ? JSON.parse(storedFlowState) : null; + flowState = JSON.parse(storedFlowState); } catch (err) { clearStoredFlow(); processingRef.current = false; diff --git a/ui/litellm-dashboard/src/hooks/useToolsOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useToolsOAuthFlow.tsx new file mode 100644 index 00000000000..66e59b80db4 --- /dev/null +++ b/ui/litellm-dashboard/src/hooks/useToolsOAuthFlow.tsx @@ -0,0 +1,232 @@ +"use client"; + +/** + * OAuth2 PKCE flow for the Tools screen re-authentication path. + * + * Unlike useUserMcpOAuthFlow (used in the chat panel), this hook: + * - stores the resulting token in sessionStorage via mcpTokenStore only + * - does NOT call storeMCPOAuthUserCredential (no backend DB write) + * - uses "litellm-tools-mcp-oauth-result" as its result key to avoid + * collisions with the admin and user flows + * + * The OAuth callback page (src/app/mcp/oauth/callback/page.tsx) writes + * to this key so this hook can pick up the result after the redirect. + */ + +import { useCallback, useEffect, useRef, useState } from "react"; +import { + buildMcpOAuthAuthorizeUrl, + exchangeMcpOAuthToken, + registerMcpOAuthClient, +} from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { extractErrorMessage } from "@/utils/errorUtils"; +import { generateCodeChallenge, generateCodeVerifier } from "@/utils/pkce"; +import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; +import { setToken } from "@/utils/mcpTokenStore"; +import { buildCallbackUrl, clearStorage } from "./mcpOAuthUtils"; + +export type ToolsOAuthStatus = "idle" | "authorizing" | "exchanging" | "success" | "error"; + +interface UseToolsOAuthFlowOptions { + accessToken: string; + serverId: string; + serverAlias?: string | null; + userId?: string | null; + scopes?: string[]; + clientId?: string | null; + onSuccess: (accessToken: string) => void; +} + +interface UseToolsOAuthFlowResult { + startOAuthFlow: () => Promise; + status: ToolsOAuthStatus; + error: string | null; +} + +const FLOW_STATE_KEY = "litellm-tools-mcp-oauth-flow-state"; +const RESULT_KEY = "litellm-tools-mcp-oauth-result"; +const RETURN_URL_KEY = "litellm-mcp-oauth-return-url"; + +type StoredFlowState = { + state: string; + codeVerifier: string; + serverId: string; + redirectUri: string; + clientId?: string; + clientSecret?: string; + scopes?: string[]; +}; + +export const useToolsOAuthFlow = ({ + accessToken, + serverId, + serverAlias, + userId, + scopes, + clientId: preClientId, + onSuccess, +}: UseToolsOAuthFlowOptions): UseToolsOAuthFlowResult => { + const [status, setStatus] = useState("idle"); + const [error, setError] = useState(null); + const processingRef = useRef(false); + const onSuccessRef = useRef(onSuccess); + onSuccessRef.current = onSuccess; + + const startOAuthFlow = useCallback(async () => { + if (typeof window === "undefined") return; + try { + setStatus("authorizing"); + setError(null); + + let clientId: string | undefined = preClientId ?? undefined; + let clientSecret: string | undefined; + + if (!clientId) { + try { + const reg = await registerMcpOAuthClient(accessToken, serverId, { + client_name: serverAlias || serverId, + grant_types: ["authorization_code", "refresh_token"], + response_types: ["code"], + token_endpoint_auth_method: "none", + }); + clientId = reg?.client_id; + clientSecret = reg?.client_secret; + } catch (_) { + // Registration is optional; proceed without client_id + } + } + + const verifier = generateCodeVerifier(); + const challenge = await generateCodeChallenge(verifier); + const state = crypto.randomUUID(); + const redirectUri = buildCallbackUrl(); + const scopeString = scopes?.filter((s) => s.trim()).join(" "); + + const authorizeUrl = buildMcpOAuthAuthorizeUrl({ + serverId, + clientId, + redirectUri, + state, + codeChallenge: challenge, + scope: scopeString, + }); + + const flowState: StoredFlowState = { + state, + codeVerifier: verifier, + serverId, + redirectUri, + clientId, + clientSecret, + scopes, + }; + + setSecureItem(FLOW_STATE_KEY, JSON.stringify(flowState)); + // Return to the current page (Tools tab) after the OAuth redirect + setSecureItem(RETURN_URL_KEY, window.location.href); + + window.location.href = authorizeUrl; + } catch (err) { + const msg = extractErrorMessage(err); + setError(msg); + setStatus("error"); + NotificationsManager.error(msg); + } + }, [accessToken, serverId, serverAlias, scopes, preClientId]); + + const resumeOAuthFlow = useCallback(async () => { + if (typeof window === "undefined" || processingRef.current) return; + + const storedResult = getSecureItem(RESULT_KEY); + if (!storedResult) return; + + // The callback page writes to this result key for every OAuth flow (including + // the admin server-creation flow). Guard: only proceed if *this* hook's flow + // state exists, meaning startOAuthFlow() was actually called from the Tools screen. + // Without this guard, a stale result written during server creation would trigger + // "OAuth session state was lost" when the user navigates to the Tools tab. + const rawFlowState = getSecureItem(FLOW_STATE_KEY); + if (!rawFlowState) return; + + let peeked: StoredFlowState | null = null; + try { + peeked = JSON.parse(rawFlowState) as StoredFlowState; + if (peeked.serverId && peeked.serverId !== serverId) return; + } catch (_) {} + + processingRef.current = true; + clearStorage(RESULT_KEY); + + let payload: Record | null = null; + let flowState: StoredFlowState | null = null; + + try { + payload = JSON.parse(storedResult); + flowState = peeked; + } catch (_) { + setError("Failed to resume OAuth flow. Please retry."); + setStatus("error"); + processingRef.current = false; + clearStorage(FLOW_STATE_KEY); + return; + } + + try { + if (!flowState?.state || !flowState.codeVerifier || !flowState.serverId) { + throw new Error("OAuth session state was lost. Please retry."); + } + if (!payload?.state || payload.state !== flowState.state) { + throw new Error("OAuth state mismatch. Please retry."); + } + if (payload.error) { + throw new Error((payload.error_description as string) || (payload.error as string)); + } + if (!payload.code) { + throw new Error("Authorization code missing in callback."); + } + + setStatus("exchanging"); + const token = await exchangeMcpOAuthToken({ + serverId: flowState.serverId, + code: payload.code as string, + clientId: flowState.clientId, + clientSecret: flowState.clientSecret, + codeVerifier: flowState.codeVerifier, + redirectUri: flowState.redirectUri, + accessToken, + }); + + // Store in sessionStorage only — no backend DB write + setToken( + flowState.serverId, + { + access_token: token.access_token, + expires_in: token.expires_in, + refresh_token: token.refresh_token, + token_type: token.token_type, + }, + userId, + ); + + setStatus("success"); + setError(null); + NotificationsManager.success("Connected successfully"); + onSuccessRef.current(token.access_token); + } catch (err) { + const msg = extractErrorMessage(err); + setError(msg); + setStatus("error"); + NotificationsManager.error(msg); + } finally { + clearStorage(FLOW_STATE_KEY); + setTimeout(() => { processingRef.current = false; }, 1000); + } + }, [accessToken, serverId, userId]); + + useEffect(() => { + resumeOAuthFlow(); + }, [resumeOAuthFlow]); + + return { startOAuthFlow, status, error }; +}; diff --git a/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx index cf0a81dcadf..1dc7a5ee54a 100644 --- a/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx @@ -16,15 +16,14 @@ import { useCallback, useEffect, useRef, useState } from "react"; import { buildMcpOAuthAuthorizeUrl, exchangeMcpOAuthToken, - getProxyBaseUrl, registerMcpOAuthClient, - serverRootPath, storeMCPOAuthUserCredential, } from "@/components/networking"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { extractErrorMessage } from "@/utils/errorUtils"; import { generateCodeChallenge, generateCodeVerifier } from "@/utils/pkce"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; +import { buildCallbackUrl, clearStorage } from "./mcpOAuthUtils"; export type UserMcpOAuthStatus = "idle" | "authorizing" | "exchanging" | "success" | "error"; @@ -69,26 +68,6 @@ const getStorage = (key: string): string | null => { return getSecureItem(key); }; -const clearStorage = (...keys: string[]) => { - keys.forEach((k) => { - try { - window.sessionStorage.removeItem(k); - } catch (_) {} - }); -}; - -const buildCallbackUrl = (): string => { - if (typeof window !== "undefined") { - const path = window.location.pathname || ""; - const idx = path.indexOf("/ui"); - const prefix = idx >= 0 ? path.slice(0, idx + 3).replace(/\/+$/, "") : ""; - return `${window.location.origin}${prefix}/mcp/oauth/callback`; - } - const base = (getProxyBaseUrl() || "").replace(/\/+$/, ""); - const root = serverRootPath && serverRootPath !== "/" ? serverRootPath : ""; - return `${base}${root}/ui/mcp/oauth/callback`; -}; - export const useUserMcpOAuthFlow = ({ accessToken, serverId, @@ -176,13 +155,17 @@ export const useUserMcpOAuthFlow = ({ // mount and would compete for the same RESULT_KEY. Peek at the stored // flow state first: only the hook instance whose serverId matches the one // that initiated the OAuth flow should consume the result. + // Guard: only proceed if this hook's flow state exists (startOAuthFlow was + // called from this hook). Without the guard, a tools re-auth redirect writes + // to the user result key too, and every OAuth2ConnectButton instance would try + // to resume a flow that was never started here. const rawFlowState = getStorage(FLOW_STATE_KEY); - if (rawFlowState) { - try { - const peeked = JSON.parse(rawFlowState) as StoredFlowState; - if (peeked.serverId && peeked.serverId !== serverId) return; - } catch (_) {} - } + if (!rawFlowState) return; + + try { + const peeked = JSON.parse(rawFlowState) as StoredFlowState; + if (peeked.serverId && peeked.serverId !== serverId) return; + } catch (_) {} processingRef.current = true; clearStorage(RESULT_KEY); diff --git a/ui/litellm-dashboard/src/utils/cookieUtils.test.ts b/ui/litellm-dashboard/src/utils/cookieUtils.test.ts index c7bd27a6a85..28e2fc771c2 100644 --- a/ui/litellm-dashboard/src/utils/cookieUtils.test.ts +++ b/ui/litellm-dashboard/src/utils/cookieUtils.test.ts @@ -1,5 +1,6 @@ import { describe, it, expect, beforeEach, vi } from "vitest"; import { clearTokenCookies, getCookie, storeLoginToken } from "./cookieUtils"; +import { getToken, setToken } from "./mcpTokenStore"; describe("cookieUtils", () => { beforeEach(() => { @@ -20,6 +21,15 @@ describe("cookieUtils", () => { expect(getCookie("token")).toBeNull(); }); + it("should clear MCP session tokens on logout", () => { + setToken("server-1", { access_token: "mcp-tok" }, "user-a"); + expect(getToken("server-1", "user-a")).not.toBeNull(); + + clearTokenCookies(); + + expect(getToken("server-1", "user-a")).toBeNull(); + }); + it("should clear token cookie from /ui path", () => { document.cookie = "token=test-token-value; path=/ui"; clearTokenCookies(); diff --git a/ui/litellm-dashboard/src/utils/cookieUtils.ts b/ui/litellm-dashboard/src/utils/cookieUtils.ts index b4493744ad4..da232e72e2a 100644 --- a/ui/litellm-dashboard/src/utils/cookieUtils.ts +++ b/ui/litellm-dashboard/src/utils/cookieUtils.ts @@ -2,6 +2,8 @@ * Utility functions for managing cookies */ +import { clearAllMcpTokens } from "./mcpTokenStore"; + /** * Returns the cookie path for the UI. * Derives the path from window.location.pathname so it works when @@ -67,6 +69,7 @@ export function clearTokenCookies() { // sessionStorage may be unavailable } + clearAllMcpTokens(); } /** diff --git a/ui/litellm-dashboard/src/utils/mcpHeaderUtils.test.ts b/ui/litellm-dashboard/src/utils/mcpHeaderUtils.test.ts new file mode 100644 index 00000000000..b730b9c097c --- /dev/null +++ b/ui/litellm-dashboard/src/utils/mcpHeaderUtils.test.ts @@ -0,0 +1,16 @@ +import { describe, expect, it } from "vitest"; +import { sanitizeMcpAliasForHeader } from "./mcpHeaderUtils"; + +describe("sanitizeMcpAliasForHeader", () => { + it("lowercases and replaces spaces with underscores", () => { + expect(sanitizeMcpAliasForHeader("My Server")).toBe("my_server"); + }); + + it("replaces invalid characters for header token segments", () => { + expect(sanitizeMcpAliasForHeader("GitHub-MCP!")).toBe("github_mcp"); + }); + + it("preserves underscores and digits", () => { + expect(sanitizeMcpAliasForHeader("github_mcp2")).toBe("github_mcp2"); + }); +}); diff --git a/ui/litellm-dashboard/src/utils/mcpHeaderUtils.ts b/ui/litellm-dashboard/src/utils/mcpHeaderUtils.ts new file mode 100644 index 00000000000..76c752a4b61 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/mcpHeaderUtils.ts @@ -0,0 +1,14 @@ +/** + * Sanitize an MCP server alias for use in HTTP header names (x-mcp-{alias}-...). + * RFC 7230 tchar allows token chars; aliases with spaces or hyphens break parsing + * because the backend splits on the first dash after the x-mcp- prefix. + * Keep in sync with litellm.proxy._experimental.mcp_server.utils.sanitize_mcp_alias_for_header. + */ +export function sanitizeMcpAliasForHeader(alias: string): string { + return alias + .toLowerCase() + .trim() + .replace(/[^a-z0-9_]/g, "_") + .replace(/_+/g, "_") + .replace(/^_|_$/g, ""); +} diff --git a/ui/litellm-dashboard/src/utils/mcpTokenStore.test.ts b/ui/litellm-dashboard/src/utils/mcpTokenStore.test.ts new file mode 100644 index 00000000000..1c61e9b1a26 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/mcpTokenStore.test.ts @@ -0,0 +1,46 @@ +import { afterEach, beforeEach, describe, expect, it } from "vitest"; +import { + clearAllMcpTokens, + getToken, + isTokenValid, + removeToken, + setToken, +} from "./mcpTokenStore"; + +describe("mcpTokenStore", () => { + beforeEach(() => { + sessionStorage.clear(); + }); + + afterEach(() => { + sessionStorage.clear(); + }); + + it("scopes tokens by user id", () => { + setToken("server-a", { access_token: "user1-token" }, "user-1"); + setToken("server-a", { access_token: "user2-token" }, "user-2"); + + expect(getToken("server-a", "user-1")?.access_token).toBe("user1-token"); + expect(getToken("server-a", "user-2")?.access_token).toBe("user2-token"); + expect(getToken("server-a", "user-3")).toBeNull(); + }); + + it("validates expiry per user scope", () => { + setToken("server-a", { access_token: "tok", expires_in: 3600 }, "user-1"); + expect(isTokenValid("server-a", "user-1")).toBe(true); + removeToken("server-a", "user-1"); + expect(isTokenValid("server-a", "user-1")).toBe(false); + }); + + it("clearAllMcpTokens removes every mcp-session-token entry", () => { + setToken("s1", { access_token: "a" }, "u1"); + setToken("s2", { access_token: "b" }, "u2"); + sessionStorage.setItem("unrelated", "keep"); + + clearAllMcpTokens(); + + expect(getToken("s1", "u1")).toBeNull(); + expect(getToken("s2", "u2")).toBeNull(); + expect(sessionStorage.getItem("unrelated")).toBe("keep"); + }); +}); diff --git a/ui/litellm-dashboard/src/utils/mcpTokenStore.ts b/ui/litellm-dashboard/src/utils/mcpTokenStore.ts new file mode 100644 index 00000000000..0279922cd07 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/mcpTokenStore.ts @@ -0,0 +1,93 @@ +/** + * Session-storage-backed OAuth token store for MCP servers. + * Tokens are keyed by LiteLLM user id + server_id and cleared when the browser + * session ends (tab/window close). Never written to localStorage. + */ + +const KEY_PREFIX = "mcp-session-token:"; + +interface StoredToken { + access_token: string; + expires_at: number; + refresh_token?: string; + token_type: string; +} + +interface TokenInput { + access_token: string; + expires_in?: number; + refresh_token?: string; + token_type?: string; +} + +const DEFAULT_TTL_MS = 3600 * 1000; // 1 hour + +function storageKey(serverId: string, userId?: string | null): string { + const userPart = userId?.trim() || "_anonymous"; + return `${KEY_PREFIX}${userPart}:${serverId}`; +} + +export function setToken( + serverId: string, + data: TokenInput, + userId?: string | null, +): void { + if (typeof window === "undefined") return; + const stored: StoredToken = { + access_token: data.access_token, + expires_at: Date.now() + (data.expires_in != null ? data.expires_in * 1000 : DEFAULT_TTL_MS), + token_type: data.token_type ?? "bearer", + ...(data.refresh_token ? { refresh_token: data.refresh_token } : {}), + }; + try { + window.sessionStorage.setItem(storageKey(serverId, userId), JSON.stringify(stored)); + } catch { + // Silently ignore storage errors (private browsing, quota exceeded, etc.) + } +} + +export function getToken( + serverId: string, + userId?: string | null, +): StoredToken | null { + if (typeof window === "undefined") return null; + try { + const raw = window.sessionStorage.getItem(storageKey(serverId, userId)); + if (!raw) return null; + return JSON.parse(raw) as StoredToken; + } catch { + return null; + } +} + +export function removeToken(serverId: string, userId?: string | null): void { + if (typeof window === "undefined") return; + try { + window.sessionStorage.removeItem(storageKey(serverId, userId)); + } catch { + // Silently ignore + } +} + +export function isTokenValid(serverId: string, userId?: string | null): boolean { + const token = getToken(serverId, userId); + if (!token) return false; + return token.expires_at > Date.now(); +} + +/** Remove all MCP session tokens (e.g. on logout or user switch). */ +export function clearAllMcpTokens(): void { + if (typeof window === "undefined") return; + try { + const keysToRemove: string[] = []; + for (let i = 0; i < window.sessionStorage.length; i++) { + const key = window.sessionStorage.key(i); + if (key?.startsWith(KEY_PREFIX)) { + keysToRemove.push(key); + } + } + keysToRemove.forEach((key) => window.sessionStorage.removeItem(key)); + } catch { + // Silently ignore + } +}