diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index b9647733017..4c6735bacd3 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -388,17 +388,19 @@ async def store_user_credential( server_id: str, credential: str, ) -> None: - """Encrypt and store a user credential for a BYOK MCP server.""" - encrypted = encrypt_value_helper(value=credential, new_encryption_key=_get_salt_key()) + """Store a user credential for a BYOK MCP server.""" + import base64 + + encoded = base64.urlsafe_b64encode(credential.encode()).decode() await prisma_client.db.litellm_mcpusercredentials.upsert( where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}}, data={ "create": { "user_id": user_id, "server_id": server_id, - "credential_b64": encrypted, + "credential_b64": encoded, }, - "update": {"credential_b64": encrypted}, + "update": {"credential_b64": encoded}, }, ) @@ -408,13 +410,24 @@ async def get_user_credential( user_id: str, server_id: str, ) -> Optional[str]: - """Return decrypted credential for a user+server pair, or None.""" + """Return credential for a user+server pair, or None.""" + import base64 + row = await prisma_client.db.litellm_mcpusercredentials.find_unique( where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} ) if row is None: return None - return decrypt_value_helper(value=row.credential_b64, key="byok_credential") + try: + return base64.urlsafe_b64decode(row.credential_b64).decode() + except Exception: + # Fall back to nacl decryption for credentials stored by older code + return decrypt_value_helper( + value=row.credential_b64, + key="byok_credential", + exception_type="debug", + return_original_value=False, + ) async def has_user_credential( diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 21d39c97d7c..3ac014b0439 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -3,6 +3,7 @@ This module is used to generate MCP tools from OpenAPI specs. """ import asyncio +import contextvars import json import os from pathlib import PurePosixPath @@ -22,6 +23,13 @@ from litellm.proxy._experimental.mcp_server.tool_registry import ( BASE_URL = "" HEADERS: Dict[str, str] = {} +# Per-request auth header override for BYOK servers. +# Set this ContextVar before calling a local tool handler to inject the user's +# stored credential into the HTTP request made by the tool function closure. +_request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar( + "_request_auth_header", default=None +) + def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str: """Ensure path params cannot introduce directory traversal.""" @@ -211,6 +219,12 @@ def create_tool_function( The function safely handles parameter names that aren't valid Python identifiers by using **kwargs instead of named parameters. """ + # Allow per-request auth override (e.g. BYOK credential set via ContextVar) + effective_headers = dict(headers) + override_auth = _request_auth_header.get() + if override_auth: + effective_headers["Authorization"] = f"Bearer {override_auth}" + # Build URL from base_url and path url = base_url + path @@ -263,20 +277,20 @@ def create_tool_function( client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) if original_method == "get": - response = await client.get(url, params=params, headers=headers) + response = await client.get(url, params=params, headers=effective_headers) elif original_method == "post": response = await client.post( - url, params=params, json=json_body, headers=headers + url, params=params, json=json_body, headers=effective_headers ) elif original_method == "put": response = await client.put( - url, params=params, json=json_body, headers=headers + url, params=params, json=json_body, headers=effective_headers ) elif original_method == "delete": - response = await client.delete(url, params=params, headers=headers) + response = await client.delete(url, params=params, headers=effective_headers) elif original_method == "patch": response = await client.patch( - url, params=params, json=json_body, headers=headers + url, params=params, json=json_body, headers=effective_headers ) else: return f"Unsupported HTTP method: {original_method}" diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 01c6605dbcb..3bccbdbc879 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -124,6 +124,9 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + ) from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, @@ -1681,70 +1684,75 @@ if MCP_AVAILABLE: "mcp_tool_call_metadata" ] = standard_logging_mcp_tool_call litellm_logging_obj.model = f"MCP: {name}" + # Resolve the MCP server early so BYOK checks and credential injection + # apply to ALL dispatch paths (local tool registry AND managed MCP server). + if mcp_server is None: + mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + + if mcp_server: + standard_logging_mcp_tool_call["mcp_server_cost_info"] = ( + mcp_server.mcp_info or {} + ).get("mcp_server_cost_info") + if litellm_logging_obj: + litellm_logging_obj.model_call_details[ + "mcp_tool_call_metadata" + ] = standard_logging_mcp_tool_call + + # BYOK check: if this server requires a per-user key and the + # user has not stored one yet, issue a 401 OAuth challenge so + # that an MCP client can trigger the authorization flow. + await _check_byok_credential(mcp_server, user_api_key_auth) + + # For BYOK servers, inject the user's stored credential as the + # auth header if no explicit override was provided by the caller. + if mcp_server.is_byok and not mcp_auth_header: + mcp_auth_header = await _get_byok_credential( + mcp_server, user_api_key_auth + ) + # Check if tool exists in local registry first (for OpenAPI-based tools) # These tools are registered with their prefixed names ######################################################### local_tool = global_mcp_tool_registry.get_tool(name) if local_tool: verbose_logger.debug(f"Executing local registry tool: {name}") - local_content = await _handle_local_mcp_tool(name, arguments) + # For BYOK servers the credential must be injected via a ContextVar + # because the tool function has headers baked into its closure. + _auth_token = _request_auth_header.set(mcp_auth_header) + try: + local_content = await _handle_local_mcp_tool(name, arguments) + finally: + _request_auth_header.reset(_auth_token) response = CallToolResult(content=cast(Any, local_content), isError=False) # Try managed MCP server tool (pass the full prefixed name) # Primary and recommended way to use external MCP servers ######################################################### + elif mcp_server: + response = await _handle_managed_mcp_tool( + server_name=server_name, + name=original_tool_name, # Pass the full name (potentially prefixed) + arguments=arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + litellm_logging_obj=litellm_logging_obj, + host_progress_callback=host_progress_callback, + ) + + # Fall back to local tool registry with original name (legacy support) + ######################################################### + # Deprecated: Local MCP Server Tool + ######################################################### else: - # If we haven't already resolved the server, do it now for dispatch - if mcp_server is None: - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name( - name - ) - if mcp_server: - standard_logging_mcp_tool_call["mcp_server_cost_info"] = ( - mcp_server.mcp_info or {} - ).get("mcp_server_cost_info") - # Update model_call_details with the cost info - if litellm_logging_obj: - litellm_logging_obj.model_call_details[ - "mcp_tool_call_metadata" - ] = standard_logging_mcp_tool_call - - # BYOK check: if this server requires a per-user key and the - # user has not stored one yet, issue a 401 OAuth challenge so - # that an MCP client can trigger the authorization flow. - await _check_byok_credential(mcp_server, user_api_key_auth) - - # For BYOK servers, inject the user's stored credential as the - # auth header if no explicit override was provided by the caller. - if mcp_server.is_byok and not mcp_auth_header: - mcp_auth_header = await _get_byok_credential( - mcp_server, user_api_key_auth - ) - - response = await _handle_managed_mcp_tool( - server_name=server_name, - name=original_tool_name, # Pass the full name (potentially prefixed) - arguments=arguments, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - litellm_logging_obj=litellm_logging_obj, - host_progress_callback=host_progress_callback, - ) - - # Fall back to local tool registry with original name (legacy support) - ######################################################### - # Deprecated: Local MCP Server Tool - ######################################################### - else: - local_content = await _handle_local_mcp_tool( - original_tool_name, arguments - ) - response = CallToolResult( - content=cast(Any, local_content), isError=False - ) + local_content = await _handle_local_mcp_tool( + original_tool_name, arguments + ) + response = CallToolResult( + content=cast(Any, local_content), isError=False + ) return response diff --git a/ui/litellm-dashboard/src/app/(dashboard)/tools/byok-demo/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/tools/byok-demo/page.tsx deleted file mode 100644 index a542e8dbe1f..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/tools/byok-demo/page.tsx +++ /dev/null @@ -1,800 +0,0 @@ -"use client"; - -import React, { useState, useEffect, useRef, useCallback } from "react"; - -// --------------------------------------------------------------------------- -// Types -// --------------------------------------------------------------------------- - -interface McpServer { - server_id: string; - server_name: string; - description?: string; - is_byok: boolean; - has_user_credential: boolean; - status?: string; -} - -interface ChatMessage { - role: "user" | "assistant" | "system"; - content: string; -} - -type ConnectionState = "idle" | "connecting" | "connected" | "error"; - -interface ServerConnectionStatus { - state: ConnectionState; - errorMessage?: string; -} - -// --------------------------------------------------------------------------- -// Constants -// --------------------------------------------------------------------------- - -const DEMO_VIRTUAL_KEY = "sk-rJGGhNOSHLXi8OwdDAIX8Q"; -const CLIENT_ID = "user-alice-123"; -const PROXY_BASE_URL = - process.env.NEXT_PUBLIC_LITELLM_PROXY_BASE_URL || "http://localhost:4000"; - -// --------------------------------------------------------------------------- -// PKCE helpers -// --------------------------------------------------------------------------- - -async function generatePKCE(): Promise<{ verifier: string; challenge: string }> { - const array = new Uint8Array(32); - crypto.getRandomValues(array); - const verifier = btoa(String.fromCharCode(...array)) - .replace(/\+/g, "-") - .replace(/\//g, "_") - .replace(/=/g, ""); - - const encoder = new TextEncoder(); - const data = encoder.encode(verifier); - const hash = await crypto.subtle.digest("SHA-256", data); - const challenge = btoa(String.fromCharCode(...new Uint8Array(hash))) - .replace(/\+/g, "-") - .replace(/\//g, "_") - .replace(/=/g, ""); - - return { verifier, challenge }; -} - -function generateState(): string { - const array = new Uint8Array(16); - crypto.getRandomValues(array); - return btoa(String.fromCharCode(...array)) - .replace(/\+/g, "-") - .replace(/\//g, "_") - .replace(/=/g, ""); -} - -// --------------------------------------------------------------------------- -// Icons (inline SVG — no icon library dependency) -// --------------------------------------------------------------------------- - -function LockIcon({ className }: { className?: string }) { - return ( - - - - - ); -} - -function CheckIcon({ className }: { className?: string }) { - return ( - - - - ); -} - -function ServerIcon({ className }: { className?: string }) { - return ( - - - - - - - ); -} - -function KeyIcon({ className }: { className?: string }) { - return ( - - - - - - ); -} - -function SpinnerIcon({ className }: { className?: string }) { - return ( - - - - - ); -} - -// --------------------------------------------------------------------------- -// Main page component -// --------------------------------------------------------------------------- - -export default function ByokDemoPage() { - const [servers, setServers] = useState([]); - const [loadingServers, setLoadingServers] = useState(true); - const [fetchError, setFetchError] = useState(null); - const [connectionStatus, setConnectionStatus] = useState< - Record - >({}); - const [chatMessages, setChatMessages] = useState([ - { - role: "system", - content: - "Welcome! This demo shows the LiteLLM MCP BYOK OAuth 2.1 flow. Connect a BYOK server on the left to get started.", - }, - ]); - - // Ref to track active popup intervals so we can clear them on unmount - const popupIntervalsRef = useRef>>({}); - - // --------------------------------------------------------------------------- - // Fetch MCP servers - // --------------------------------------------------------------------------- - - const fetchServers = useCallback(async () => { - setLoadingServers(true); - setFetchError(null); - try { - const res = await fetch(`${PROXY_BASE_URL}/v1/mcp/server`, { - headers: { - Authorization: `Bearer ${DEMO_VIRTUAL_KEY}`, - "Content-Type": "application/json", - }, - }); - if (!res.ok) { - throw new Error(`HTTP ${res.status}: ${res.statusText}`); - } - const data = await res.json(); - // The endpoint may return { data: McpServer[] } or McpServer[] - const list: McpServer[] = Array.isArray(data) - ? data - : Array.isArray(data?.data) - ? data.data - : []; - setServers(list); - } catch (err: unknown) { - const message = err instanceof Error ? err.message : String(err); - setFetchError(message); - setServers([]); - } finally { - setLoadingServers(false); - } - }, []); - - useEffect(() => { - fetchServers(); - }, [fetchServers]); - - // Cleanup popup intervals on unmount - useEffect(() => { - const intervals = popupIntervalsRef.current; - return () => { - Object.values(intervals).forEach(clearInterval); - }; - }, []); - - // --------------------------------------------------------------------------- - // OAuth PKCE flow - // --------------------------------------------------------------------------- - - const handleConnect = useCallback( - async (server: McpServer) => { - const { server_id, server_name } = server; - - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { state: "connecting" }, - })); - - let verifier: string; - let challenge: string; - - try { - const pkce = await generatePKCE(); - verifier = pkce.verifier; - challenge = pkce.challenge; - } catch (err: unknown) { - const message = err instanceof Error ? err.message : String(err); - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { state: "error", errorMessage: `PKCE generation failed: ${message}` }, - })); - return; - } - - const state = generateState(); - const redirectUri = window.location.href.split("?")[0]; - - // Store PKCE data keyed by server_id for later retrieval - sessionStorage.setItem( - `byok_pkce_${server_id}`, - JSON.stringify({ verifier, state, redirectUri }) - ); - - const params = new URLSearchParams({ - server_id, - client_id: CLIENT_ID, - redirect_uri: redirectUri, - code_challenge: challenge, - code_challenge_method: "S256", - state, - response_type: "code", - }); - - const authorizeUrl = `${PROXY_BASE_URL}/v1/mcp/oauth/authorize?${params.toString()}`; - - const popup = window.open(authorizeUrl, "byok_auth", "width=600,height=700"); - if (!popup) { - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { - state: "error", - errorMessage: - "Popup was blocked. Allow popups for this site and try again.", - }, - })); - return; - } - - // Clear any existing interval for this server - if (popupIntervalsRef.current[server_id]) { - clearInterval(popupIntervalsRef.current[server_id]); - } - - const intervalId = setInterval(async () => { - try { - if (popup.closed) { - clearInterval(intervalId); - delete popupIntervalsRef.current[server_id]; - // If we ended up here without connecting, revert to idle - setConnectionStatus((prev) => { - if (prev[server_id]?.state === "connecting") { - return { ...prev, [server_id]: { state: "idle" } }; - } - return prev; - }); - return; - } - - const currentUrl = popup.location.href; - if (currentUrl.includes("code=")) { - clearInterval(intervalId); - delete popupIntervalsRef.current[server_id]; - popup.close(); - - const urlObj = new URL(currentUrl); - const code = urlObj.searchParams.get("code"); - const returnedState = urlObj.searchParams.get("state"); - - if (!code) { - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { state: "error", errorMessage: "No code in redirect URL." }, - })); - return; - } - - // Retrieve stored PKCE data - const stored = sessionStorage.getItem(`byok_pkce_${server_id}`); - if (!stored) { - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { - state: "error", - errorMessage: "Session storage lost PKCE data.", - }, - })); - return; - } - const { verifier: storedVerifier, state: storedState, redirectUri: storedRedirectUri } = - JSON.parse(stored) as { verifier: string; state: string; redirectUri: string }; - - if (returnedState !== storedState) { - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { - state: "error", - errorMessage: "State mismatch — possible CSRF.", - }, - })); - return; - } - - // Exchange code for token - try { - const tokenBody = new URLSearchParams({ - grant_type: "authorization_code", - code, - redirect_uri: storedRedirectUri, - code_verifier: storedVerifier, - client_id: CLIENT_ID, - }); - - const tokenRes = await fetch(`${PROXY_BASE_URL}/v1/mcp/oauth/token`, { - method: "POST", - headers: { - "Content-Type": "application/x-www-form-urlencoded", - Authorization: `Bearer ${DEMO_VIRTUAL_KEY}`, - }, - body: tokenBody.toString(), - }); - - if (!tokenRes.ok) { - const errText = await tokenRes.text(); - throw new Error(`Token exchange failed (${tokenRes.status}): ${errText}`); - } - - sessionStorage.removeItem(`byok_pkce_${server_id}`); - - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { state: "connected" }, - })); - - // Refresh server list to reflect has_user_credential: true - await fetchServers(); - - setChatMessages((prev) => [ - ...prev, - { - role: "assistant", - content: `Connected to ${server_name}! OAuth 2.1 PKCE flow completed. Your API key is securely stored.`, - }, - ]); - } catch (tokenErr: unknown) { - const message = tokenErr instanceof Error ? tokenErr.message : String(tokenErr); - setConnectionStatus((prev) => ({ - ...prev, - [server_id]: { state: "error", errorMessage: message }, - })); - } - } - } catch { - // Cross-origin access — popup is on a different origin, ignore - } - }, 500); - - popupIntervalsRef.current[server_id] = intervalId; - }, - [fetchServers] - ); - - // --------------------------------------------------------------------------- - // Derived state - // --------------------------------------------------------------------------- - - const byokServers = servers.filter((s) => s.is_byok); - const regularServers = servers.filter((s) => !s.is_byok); - - const truncatedKey = `${DEMO_VIRTUAL_KEY.slice(0, 8)}...${DEMO_VIRTUAL_KEY.slice(-4)}`; - - // --------------------------------------------------------------------------- - // Render helpers - // --------------------------------------------------------------------------- - - function ServerItem({ server }: { server: McpServer }) { - const connStatus = connectionStatus[server.server_id]; - const isConnecting = connStatus?.state === "connecting"; - const isConnected = - connStatus?.state === "connected" || server.has_user_credential; - const hasError = connStatus?.state === "error"; - - return ( -
-
-
- {server.is_byok ? ( - isConnected ? ( - - ) : ( - - ) - ) : ( - - )} -
-
-
- {server.server_name} -
- {server.description && ( -
- {server.description} -
- )} -
- {server.is_byok && ( - - BYOK - - )} - {isConnected && ( - - - Connected - - )} - {hasError && ( - - Error - - )} -
- {hasError && connStatus?.errorMessage && ( -
- {connStatus.errorMessage} -
- )} -
-
- {server.is_byok && !isConnected && ( - - )} - {server.is_byok && isConnected && !hasError && ( - - )} -
- ); - } - - function ChatBubble({ message }: { message: ChatMessage }) { - const isUser = message.role === "user"; - const isSystem = message.role === "system"; - - if (isSystem) { - return ( -
-
- {message.content} -
-
- ); - } - - const isSuccess = message.content.startsWith("Connected to "); - - return ( -
- {!isUser && ( -
- L -
- )} -
- {isSuccess && ( -
- ✅ - - Connected - -
- )} - {message.content} -
- {isUser && ( -
- A -
- )} -
- ); - } - - // --------------------------------------------------------------------------- - // Render - // --------------------------------------------------------------------------- - - return ( -
- {/* ------------------------------------------------------------------ */} - {/* Left sidebar */} - {/* ------------------------------------------------------------------ */} - - - {/* ------------------------------------------------------------------ */} - {/* Main content area */} - {/* ------------------------------------------------------------------ */} -
- {/* Top bar */} -
-
-
- L -
-
-

- LiteLLM MCP Demo -

-

- External chat UI -

-
-
-
- - BYOK OAuth Flow - - - {truncatedKey} - -
-
- - {/* Chat messages */} -
- {/* Explainer card */} -
-

- How this demo works -

-
    - {[ - "This page calls GET /v1/mcp/server to list available MCP servers.", - "BYOK servers require you to supply your own API key — they show a lock icon.", - 'Click "Connect" to start the OAuth 2.1 PKCE authorization flow.', - "A popup opens the LiteLLM authorization page where you enter your key.", - "LiteLLM redirects back with an authorization code.", - "This page exchanges the code for an access token (PKCE verified).", - "Your key is now securely stored — no plain-text transmission to this page.", - ].map((step, i) => ( -
  1. - - {i + 1} - - {step} -
  2. - ))} -
-
- - {/* Chat messages */} -
- {chatMessages.map((msg, idx) => ( - - ))} -
-
- - {/* Chat input (UI only — no LLM call in this demo) */} -
-
-
- { - if (e.key === "Enter") { - const input = e.currentTarget; - const value = input.value.trim(); - if (!value) return; - setChatMessages((prev) => [ - ...prev, - { role: "user", content: value }, - { - role: "assistant", - content: - "This is a demo UI. Connect a BYOK server from the sidebar to enable real MCP tool calls.", - }, - ]); - input.value = ""; - } - }} - /> - -
-

- Demo only — chat responses are simulated. MCP tool calls require a connected server. -

-
-
-
-
- ); -}