From 37b87789e5039787a19cd61b7ae790e9789fc756 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 4 Mar 2026 19:53:43 -0800 Subject: [PATCH] feat(mcp): BYOK (Bring Your Own Key) for OpenAPI MCP servers with OAuth 2.1 flow Adds per-user credential storage for BYOK MCP servers so external clients can authenticate via standard OAuth 2.1 PKCE without needing a full identity provider. Backend: - New DB table LiteLLM_MCPUserCredentials (user_id, server_id, credential_b64) - is_byok, byok_description, byok_api_key_help_url fields on MCPServerTable - OAuth 2.1 authorization server endpoints (/.well-known/oauth-authorization-server, /.well-known/oauth-protected-resource, /v1/mcp/oauth/authorize, /v1/mcp/oauth/token) - 401 challenge with WWW-Authenticate header when BYOK server has no credential - CRUD endpoints: POST/DELETE /v1/mcp/server/{id}/user-credential - has_user_credential annotated on GET /v1/mcp/server response UI: - ByokCredentialModal: 2-step Connect flow (access description + API key entry) - BYOK toggle + description fields on admin MCP server create form - Connect/Connected state in MCP server table - BYOK Demo page (/tools/byok-demo) showing full OAuth 2.1 PKCE flow --- CLAUDE.md | 12 +- .../mcp_server/byok_oauth_endpoints.py | 332 ++++++++ litellm/proxy/_experimental/mcp_server/db.py | 63 ++ .../mcp_server/mcp_server_manager.py | 6 + .../proxy/_experimental/mcp_server/server.py | 62 ++ litellm/proxy/_types.py | 20 + .../mcp_management_endpoints.py | 83 ++ litellm/proxy/proxy_server.py | 4 + litellm/proxy/schema.prisma | 16 + .../types/mcp_server/mcp_server_manager.py | 3 + .../mcp_server/test_byok_oauth_endpoints.py | 515 +++++++++++ .../app/(dashboard)/components/Sidebar2.tsx | 8 + .../app/(dashboard)/tools/byok-demo/page.tsx | 800 ++++++++++++++++++ .../mcp_tools/ByokCredentialModal.tsx | 254 ++++++ .../mcp_tools/create_mcp_server.tsx | 85 +- .../mcp_tools/mcp_server_columns.tsx | 37 + .../src/components/mcp_tools/mcp_servers.tsx | 16 + .../src/components/mcp_tools/types.tsx | 6 + .../components/playground/chat_ui/ChatUI.tsx | 59 ++ 19 files changed, 2379 insertions(+), 2 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/tools/byok-demo/page.tsx create mode 100644 ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx diff --git a/CLAUDE.md b/CLAUDE.md index c1eb75d2515..5b36c2be8ac 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -114,4 +114,14 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components: ### Enterprise Features - Enterprise-specific code in `enterprise/` directory - Optional features enabled via environment variables -- Separate licensing and authentication for enterprise features \ No newline at end of file +- Separate licensing and authentication for enterprise features + +### Troubleshooting: DB schema out of sync after proxy restart +`litellm-proxy-extras` runs `prisma migrate deploy` on startup using **its own** bundled migration files, which may lag behind schema changes in the current worktree. Symptoms: `Unknown column`, `Invalid prisma invocation`, or missing data on new fields. + +**Diagnose:** Run `\d "TableName"` in psql and compare against `schema.prisma` — missing columns confirm the issue. + +**Fix options:** +1. **Create a Prisma migration** (permanent) — run `prisma migrate dev --name ` in the worktree. The generated file will be picked up by `prisma migrate deploy` on next startup. +2. **Apply manually for local dev** — `psql -d litellm -c "ALTER TABLE ... ADD COLUMN IF NOT EXISTS ..."` after each proxy start. Fine for dev, not for production. +3. **Update litellm-proxy-extras** — if the package is installed from PyPI, its migration directory must include the new file. Either update the package or run the migration manually until the next release ships it. \ No newline at end of file diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py new file mode 100644 index 00000000000..ea4c8133064 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -0,0 +1,332 @@ +""" +BYOK (Bring Your Own Key) OAuth 2.1 Authorization Server endpoints for MCP servers. + +When an MCP client connects to a BYOK-enabled server and no stored credential exists, +LiteLLM runs a minimal OAuth 2.1 authorization code flow. The "authorization page" is +just a form that asks the user for their API key — not a full identity-provider OAuth. + +Endpoints implemented here: + GET /.well-known/oauth-authorization-server — OAuth authorization server metadata + GET /.well-known/oauth-protected-resource — OAuth protected resource metadata + GET /v1/mcp/oauth/authorize — Shows HTML form to collect the API key + POST /v1/mcp/oauth/authorize — Stores temp auth code and redirects + POST /v1/mcp/oauth/token — Exchanges code for a bearer JWT token +""" + +import base64 +import hashlib +import time +import uuid +from typing import Dict, Optional, cast + +import jwt +from fastapi import APIRouter, Form, HTTPException, Request +from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._experimental.mcp_server.db import store_user_credential +from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, +) + +# --------------------------------------------------------------------------- +# In-memory store for pending authorization codes. +# Each entry: {code: {api_key, server_id, code_challenge, redirect_uri, user_id, expires_at}} +# --------------------------------------------------------------------------- +_byok_auth_codes: Dict[str, dict] = {} + +# Authorization codes expire after 5 minutes. +_AUTH_CODE_TTL_SECONDS = 300 + +router = APIRouter(tags=["mcp"]) + + +# --------------------------------------------------------------------------- +# PKCE helper +# --------------------------------------------------------------------------- + + +def _verify_pkce(code_verifier: str, code_challenge: str) -> bool: + """Return True iff SHA-256(code_verifier) == code_challenge (base64url, no padding).""" + digest = hashlib.sha256(code_verifier.encode()).digest() + computed = base64.urlsafe_b64encode(digest).rstrip(b"=").decode() + return computed == code_challenge + + +# --------------------------------------------------------------------------- +# Cleanup of expired auth codes (called lazily on each request) +# --------------------------------------------------------------------------- + + +def _purge_expired_codes() -> None: + now = time.time() + expired = [k for k, v in _byok_auth_codes.items() if v["expires_at"] < now] + for k in expired: + del _byok_auth_codes[k] + + +# --------------------------------------------------------------------------- +# HTML template for the authorization page +# --------------------------------------------------------------------------- + +_AUTHORIZE_HTML = """ + +Connect to {server_name} — LiteLLM + + +
+

Connect to {server_name}

+

Enter your {server_name} API key to authorize this connection.

+
+ + + + + + + + + +
+

Your key is encrypted at rest and never shared with third parties.

+
+ +""" + + +# --------------------------------------------------------------------------- +# OAuth metadata discovery endpoints +# --------------------------------------------------------------------------- + + +@router.get("/.well-known/oauth-authorization-server", include_in_schema=False) +async def oauth_authorization_server_metadata(request: Request) -> JSONResponse: + """RFC 8414 Authorization Server Metadata for the BYOK OAuth flow.""" + base_url = get_request_base_url(request) + return JSONResponse( + { + "issuer": base_url, + "authorization_endpoint": f"{base_url}/v1/mcp/oauth/authorize", + "token_endpoint": f"{base_url}/v1/mcp/oauth/token", + "response_types_supported": ["code"], + "grant_types_supported": ["authorization_code"], + "code_challenge_methods_supported": ["S256"], + } + ) + + +@router.get("/.well-known/oauth-protected-resource", include_in_schema=False) +async def oauth_protected_resource_metadata(request: Request) -> JSONResponse: + """RFC 9728 Protected Resource Metadata pointing back at this server.""" + base_url = get_request_base_url(request) + return JSONResponse( + { + "resource": base_url, + "authorization_servers": [base_url], + } + ) + + +# --------------------------------------------------------------------------- +# Authorization endpoint — GET (show form) and POST (process form) +# --------------------------------------------------------------------------- + + +@router.get("/v1/mcp/oauth/authorize", include_in_schema=False) +async def byok_authorize_get( + request: Request, + client_id: Optional[str] = None, + redirect_uri: Optional[str] = None, + response_type: Optional[str] = None, + code_challenge: Optional[str] = None, + code_challenge_method: Optional[str] = None, + state: Optional[str] = None, + server_id: Optional[str] = None, +) -> HTMLResponse: + """ + Show the BYOK API-key entry form. + + The MCP client navigates the user here; the user types their API key and + clicks "Connect & Authorize", which POSTs back to this same path. + """ + if response_type != "code": + raise HTTPException(status_code=400, detail="response_type must be 'code'") + if not redirect_uri: + raise HTTPException(status_code=400, detail="redirect_uri is required") + if not code_challenge: + raise HTTPException(status_code=400, detail="code_challenge is required") + + # Resolve a human-readable server name. + server_name = server_id or "MCP Server" + if server_id: + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + registry = global_mcp_server_manager.get_registry() + if server_id in registry: + server_name = registry[server_id].server_name or registry[server_id].name + except Exception: + pass + + html = _AUTHORIZE_HTML.format( + server_name=server_name, + client_id=client_id or "", + redirect_uri=redirect_uri, + code_challenge=code_challenge, + code_challenge_method=code_challenge_method or "S256", + state=state or "", + server_id=server_id or "", + ) + return HTMLResponse(content=html) + + +@router.post("/v1/mcp/oauth/authorize", include_in_schema=False) +async def byok_authorize_post( + request: Request, + client_id: str = Form(default=""), + redirect_uri: str = Form(...), + code_challenge: str = Form(...), + code_challenge_method: str = Form(default="S256"), + state: str = Form(default=""), + server_id: str = Form(default=""), + api_key: str = Form(...), +) -> RedirectResponse: + """ + Process the BYOK API-key form submission. + + Stores a short-lived authorization code and redirects the client back to + redirect_uri with ?code=...&state=... query parameters. + """ + _purge_expired_codes() + + if code_challenge_method != "S256": + raise HTTPException( + status_code=400, detail="Only S256 code_challenge_method is supported" + ) + + auth_code = str(uuid.uuid4()) + _byok_auth_codes[auth_code] = { + "api_key": api_key, + "server_id": server_id, + "code_challenge": code_challenge, + "redirect_uri": redirect_uri, + "user_id": client_id, # external client passes LiteLLM user-id as client_id + "expires_at": time.time() + _AUTH_CODE_TTL_SECONDS, + } + + separator = "&" if "?" in redirect_uri else "?" + location = f"{redirect_uri}{separator}code={auth_code}&state={state}" + return RedirectResponse(url=location, status_code=302) + + +# --------------------------------------------------------------------------- +# Token endpoint +# --------------------------------------------------------------------------- + + +@router.post("/v1/mcp/oauth/token", include_in_schema=False) +async def byok_token( + request: Request, + grant_type: str = Form(...), + code: str = Form(...), + redirect_uri: str = Form(default=""), + code_verifier: str = Form(...), + client_id: str = Form(default=""), +) -> JSONResponse: + """ + Exchange an authorization code for a short-lived BYOK session JWT. + + 1. Validates the authorization code and PKCE challenge. + 2. Stores the API key via store_user_credential(). + 3. Issues a signed JWT with type="byok_session". + """ + from litellm.proxy.proxy_server import master_key, prisma_client + + _purge_expired_codes() + + if grant_type != "authorization_code": + raise HTTPException(status_code=400, detail="unsupported_grant_type") + + record = _byok_auth_codes.get(code) + if record is None: + raise HTTPException(status_code=400, detail="invalid_grant") + + if time.time() > record["expires_at"]: + del _byok_auth_codes[code] + raise HTTPException(status_code=400, detail="invalid_grant") + + # PKCE verification + if not _verify_pkce(code_verifier, record["code_challenge"]): + raise HTTPException(status_code=400, detail="invalid_grant") + + # Consume the code (one-time use) + del _byok_auth_codes[code] + + server_id: str = record["server_id"] + api_key_value: str = record["api_key"] + # Prefer the user_id that was stored when the code was issued; fall back to + # whatever client_id the token request supplies (they should match). + user_id: str = record.get("user_id") or client_id + + if not user_id: + raise HTTPException( + status_code=400, + detail="Cannot determine user_id; pass LiteLLM user id as client_id", + ) + + # Persist the BYOK credential + if prisma_client is not None: + try: + await store_user_credential( + prisma_client=prisma_client, + user_id=user_id, + server_id=server_id, + credential=api_key_value, + ) + except Exception as exc: + verbose_proxy_logger.error( + "byok_token: failed to store user credential for user=%s server=%s: %s", + user_id, + server_id, + exc, + ) + raise HTTPException(status_code=500, detail="Failed to store credential") + else: + verbose_proxy_logger.warning( + "byok_token: prisma_client is None — credential not persisted" + ) + + if master_key is None: + raise HTTPException( + status_code=500, detail="Master key not configured; cannot issue token" + ) + + now = int(time.time()) + payload = { + "user_id": user_id, + "server_id": server_id, + "type": "byok_session", + "iat": now, + "exp": now + 3600, + } + access_token = jwt.encode(payload, cast(str, master_key), algorithm="HS256") + + return JSONResponse( + { + "access_token": access_token, + "token_type": "bearer", + "expires_in": 3600, + } + ) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 1bc7e8f8a9d..48225d6c93d 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -1,3 +1,4 @@ +import base64 from typing import Any, Dict, Iterable, List, Optional, Set, Union, cast from litellm._logging import verbose_proxy_logger @@ -68,6 +69,10 @@ def _prepare_mcp_server_data( # mcp_access_groups is already List[str], no serialization needed + # Force include is_byok even when False (exclude_none=True would not drop it, + # but be explicit to ensure a False value is always written to the DB). + data_dict["is_byok"] = getattr(data, "is_byok", False) + return data_dict @@ -375,3 +380,61 @@ async def rotate_mcp_server_credentials_master_key( "updated_by": touched_by, }, ) + + +async def store_user_credential( + prisma_client: PrismaClient, + user_id: str, + server_id: str, + credential: str, +) -> None: + """Store a B64-encoded user credential for a BYOK MCP server.""" + credential_b64 = base64.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": credential_b64, + }, + "update": {"credential_b64": credential_b64}, + }, + ) + + +async def get_user_credential( + prisma_client: PrismaClient, + user_id: str, + server_id: str, +) -> Optional[str]: + """Return decoded credential for a user+server pair, or None.""" + 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 base64.b64decode(row.credential_b64.encode()).decode() + + +async def has_user_credential( + prisma_client: PrismaClient, + user_id: str, + server_id: str, +) -> bool: + """Return True if the user has a stored credential for this server.""" + row = await prisma_client.db.litellm_mcpusercredentials.find_unique( + where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} + ) + return row is not None + + +async def delete_user_credential( + prisma_client: PrismaClient, + user_id: str, + server_id: str, +) -> None: + """Delete the user's stored credential for a BYOK MCP server.""" + await prisma_client.db.litellm_mcpusercredentials.delete( + where={"user_id_server_id": {"user_id": user_id, "server_id": server_id}} + ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 51bdfea172b..7c17da36bb7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -650,6 +650,9 @@ class MCPServerManager: tool_name_to_description=_deserialize_json_dict( getattr(mcp_server, "tool_name_to_description", None) ), + is_byok=bool(getattr(mcp_server, "is_byok", False)), + byok_description=getattr(mcp_server, "byok_description", None) or [], + byok_api_key_help_url=getattr(mcp_server, "byok_api_key_help_url", None), ) return new_server @@ -2657,6 +2660,9 @@ class MCPServerManager: registration_url=server.registration_url, allow_all_keys=server.allow_all_keys, available_on_public_internet=server.available_on_public_internet, + is_byok=server.is_byok, + byok_description=server.byok_description, + byok_api_key_help_url=server.byok_api_key_help_url, ) async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b131800e950..f61ccd74ab9 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1498,6 +1498,62 @@ if MCP_AVAILABLE: ) return name + async def _check_byok_credential( + mcp_server: MCPServer, + user_api_key_auth: Optional[UserAPIKeyAuth], + ) -> None: + """ + If the MCP server is BYOK-enabled, verify that the requesting user has a + stored credential. When no credential is found, raise an HTTP 401 with a + WWW-Authenticate header that points the MCP client to our OAuth metadata + endpoint so it can drive the authorization flow. + """ + if not mcp_server.is_byok: + return + + user_id = (user_api_key_auth.user_id if user_api_key_auth else None) or "" + if not user_id: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": "User identity is required for BYOK servers", + }, + headers={ + "WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"' + }, + ) + + from litellm.proxy._experimental.mcp_server.db import has_user_credential + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return + + credential_exists = await has_user_credential( + prisma_client=prisma_client, + user_id=user_id, + server_id=mcp_server.server_id, + ) + if not credential_exists: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": ( + "No stored credential found for this BYOK server. " + "Complete the OAuth authorization flow to provide your API key." + ), + }, + headers={ + "WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"' + }, + ) + async def execute_mcp_tool( name: str, arguments: Dict[str, Any], @@ -1600,6 +1656,12 @@ if MCP_AVAILABLE: 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) + response = await _handle_managed_mcp_tool( server_name=server_name, name=original_tool_name, # Pass the full name (potentially prefixed) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9b07d44deb6..95dabd8bfd0 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1108,6 +1108,9 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): registration_url: Optional[str] = None allow_all_keys: bool = False available_on_public_internet: bool = True + is_byok: bool = False + byok_description: List[str] = Field(default_factory=list) + byok_api_key_help_url: Optional[str] = None @model_validator(mode="before") @classmethod @@ -1164,6 +1167,9 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): registration_url: Optional[str] = None allow_all_keys: bool = False available_on_public_internet: bool = True + is_byok: bool = False + byok_description: List[str] = Field(default_factory=list) + byok_api_key_help_url: Optional[str] = None @model_validator(mode="before") @classmethod @@ -1223,12 +1229,26 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): registration_url: Optional[str] = None allow_all_keys: bool = False available_on_public_internet: bool = True + is_byok: bool = False + byok_description: List[str] = Field(default_factory=list) + byok_api_key_help_url: Optional[str] = None + has_user_credential: Optional[bool] = None class MakeMCPServersPublicRequest(LiteLLMPydanticObjectBase): mcp_server_ids: List[str] +class MCPUserCredentialRequest(LiteLLMPydanticObjectBase): + credential: str + save: bool = True + + +class MCPUserCredentialResponse(LiteLLMPydanticObjectBase): + server_id: str + has_credential: bool + + ######## Skills API Types ######## diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 8c4d4e7937e..c594ad3ec18 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -78,8 +78,12 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.db import ( create_mcp_server, delete_mcp_server, + delete_user_credential, get_all_mcp_servers_for_user, get_mcp_server, + get_user_credential, + has_user_credential, + store_user_credential, update_mcp_server, ) from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -98,6 +102,8 @@ if MCP_AVAILABLE: LiteLLM_MCPServerTable, LitellmUserRoles, MakeMCPServersPublicRequest, + MCPUserCredentialRequest, + MCPUserCredentialResponse, NewMCPServerRequest, SpecialMCPServerName, UpdateMCPServerRequest, @@ -599,6 +605,17 @@ if MCP_AVAILABLE: server.mcp_info = {} server.mcp_info["is_public"] = True + # Annotate has_user_credential for BYOK servers + from litellm.proxy.proxy_server import prisma_client as _byok_prisma_client + + user_id = user_api_key_dict.user_id or "" + if user_id and _byok_prisma_client is not None: + for server in redacted_mcp_servers: + if getattr(server, "is_byok", False): + server.has_user_credential = await has_user_credential( + _byok_prisma_client, user_id, server.server_id + ) + # Virtual keys only get a sanitized discovery view. if is_restricted_virtual_key: return _sanitize_mcp_server_list_for_virtual_key(redacted_mcp_servers) @@ -1036,6 +1053,72 @@ if MCP_AVAILABLE: return Response(status_code=status.HTTP_202_ACCEPTED) + @router.post( + "/server/{server_id}/user-credential", + description="Store or update the calling user's API key for a BYOK MCP server", + dependencies=[Depends(user_api_key_auth)], + response_model=MCPUserCredentialResponse, + ) + @management_endpoint_wrapper + async def store_mcp_user_credential( + server_id: str, + payload: MCPUserCredentialRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + ): + """Store a BYOK credential for the calling user.""" + prisma_client = get_prisma_client_or_throw( + "Database not connected. Connect a database to your proxy" + ) + mcp_server = await get_mcp_server(prisma_client, server_id) + if mcp_server is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail={"error": f"MCP Server {server_id} not found"}, + ) + if not getattr(mcp_server, "is_byok", False): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "This MCP server does not support BYOK credentials"}, + ) + user_id = user_api_key_dict.user_id or "" + if not user_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "User ID not found in token"}, + ) + if payload.save: + await store_user_credential(prisma_client, user_id, server_id, payload.credential) + return MCPUserCredentialResponse(server_id=server_id, has_credential=True) + # save=False: credential not persisted + return MCPUserCredentialResponse(server_id=server_id, has_credential=False) + + @router.delete( + "/server/{server_id}/user-credential", + description="Delete the calling user's stored API key for a BYOK MCP server", + dependencies=[Depends(user_api_key_auth)], + response_model=MCPUserCredentialResponse, + ) + @management_endpoint_wrapper + async def delete_mcp_user_credential( + server_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + ): + """Remove the calling user's BYOK credential.""" + prisma_client = get_prisma_client_or_throw( + "Database not connected. Connect a database to your proxy" + ) + user_id = user_api_key_dict.user_id or "" + if not user_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "User ID not found in token"}, + ) + try: + await delete_user_credential(prisma_client, user_id, server_id) + except Exception: + pass # Already deleted or didn't exist + return MCPUserCredentialResponse(server_id=server_id, has_credential=False) + @router.put( "/server", description="Allows deleting mcp serves in the db", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ad31ff33802..33d84cd7078 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -231,6 +231,9 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.vertex_ai.vertex_llm_base import VertexBase +from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( + router as mcp_byok_oauth_router, +) from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( router as mcp_discoverable_endpoints_router, ) @@ -12975,6 +12978,7 @@ app.include_router(vector_store_files_router) app.include_router(credential_router) app.include_router(llm_passthrough_router) app.include_router(mcp_management_router) +app.include_router(mcp_byok_oauth_router) app.include_router(anthropic_router) app.include_router(anthropic_skills_router) app.include_router(evals_router) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 5e1ba479298..43972724ecc 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -305,6 +305,22 @@ model LiteLLM_MCPServerTable { registration_url String? allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) + spec_path String? + is_byok Boolean @default(false) + byok_description String[] @default([]) + byok_api_key_help_url String? +} + +// Per-user BYOK credentials for MCP servers +model LiteLLM_MCPUserCredentials { + id String @id @default(uuid()) + user_id String + server_id String + credential_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") + + @@unique([user_id, server_id]) } // Generate Tokens for Proxy diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 7f6a8b3ea24..d94795fda2e 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -55,6 +55,9 @@ class MCPServer(BaseModel): access_groups: Optional[List[str]] = None allow_all_keys: bool = False available_on_public_internet: bool = True + is_byok: bool = False + byok_description: List[str] = [] + byok_api_key_help_url: Optional[str] = None created_at: Optional[datetime] = None updated_at: Optional[datetime] = None model_config = ConfigDict(arbitrary_types_allowed=True) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py new file mode 100644 index 00000000000..df4e092d9df --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -0,0 +1,515 @@ +""" +Unit tests for the BYOK OAuth 2.1 authorization server endpoints. + +Covers: +- _verify_pkce helper +- OAuth metadata discovery endpoints +- Authorization GET / POST endpoints +- Token endpoint (PKCE verification, credential storage, JWT issuance) +- 401 challenge in execute_mcp_tool (_check_byok_credential) +""" + +import base64 +import hashlib +import time +import uuid +from typing import Any, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException +from fastapi.testclient import TestClient + +from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( + _byok_auth_codes, + _verify_pkce, + router, +) +from litellm.proxy._types import MCPTransport + +# --------------------------------------------------------------------------- +# _verify_pkce +# --------------------------------------------------------------------------- + + +def _make_challenge(verifier: str) -> str: + digest = hashlib.sha256(verifier.encode()).digest() + return base64.urlsafe_b64encode(digest).rstrip(b"=").decode() + + +def test_verify_pkce_valid(): + verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk" + challenge = _make_challenge(verifier) + assert _verify_pkce(verifier, challenge) is True + + +def test_verify_pkce_invalid(): + assert _verify_pkce("wrong_verifier", _make_challenge("right_verifier")) is False + + +def test_verify_pkce_tampered_challenge(): + verifier = "test_verifier_value" + challenge = _make_challenge(verifier) + # Flip one character to tamper with the challenge + tampered = challenge[:-1] + ("A" if challenge[-1] != "A" else "B") + assert _verify_pkce(verifier, tampered) is False + + +# --------------------------------------------------------------------------- +# Minimal FastAPI app for testing the router +# --------------------------------------------------------------------------- + +from fastapi import FastAPI + +_test_app = FastAPI() +_test_app.include_router(router) + + +@pytest.fixture +def client(): + return TestClient(_test_app, raise_server_exceptions=False) + + +# --------------------------------------------------------------------------- +# OAuth metadata endpoints +# --------------------------------------------------------------------------- + + +def test_oauth_authorization_server_metadata(client): + resp = client.get("/.well-known/oauth-authorization-server") + assert resp.status_code == 200 + data = resp.json() + assert "issuer" in data + assert data["authorization_endpoint"].endswith("/v1/mcp/oauth/authorize") + assert data["token_endpoint"].endswith("/v1/mcp/oauth/token") + assert "S256" in data["code_challenge_methods_supported"] + + +def test_oauth_protected_resource_metadata(client): + resp = client.get("/.well-known/oauth-protected-resource") + assert resp.status_code == 200 + data = resp.json() + assert "resource" in data + assert "authorization_servers" in data + assert len(data["authorization_servers"]) == 1 + + +# --------------------------------------------------------------------------- +# Authorization GET endpoint +# --------------------------------------------------------------------------- + + +def test_authorize_get_returns_html(client): + resp = client.get( + "/v1/mcp/oauth/authorize", + params={ + "client_id": "test-client", + "redirect_uri": "https://client.example.com/callback", + "response_type": "code", + "code_challenge": "abc123", + "code_challenge_method": "S256", + "state": "xyz", + "server_id": "my-server", + }, + follow_redirects=False, + ) + assert resp.status_code == 200 + assert "text/html" in resp.headers["content-type"] + # The button text is HTML-entity-escaped in the template + assert "Connect & Authorize" in resp.text + # Hidden fields should be embedded + assert "my-server" in resp.text + assert "abc123" in resp.text + + +def test_authorize_get_missing_redirect_uri(client): + resp = client.get( + "/v1/mcp/oauth/authorize", + params={ + "response_type": "code", + "code_challenge": "abc", + }, + follow_redirects=False, + ) + assert resp.status_code == 400 + + +def test_authorize_get_wrong_response_type(client): + resp = client.get( + "/v1/mcp/oauth/authorize", + params={ + "redirect_uri": "https://example.com/cb", + "response_type": "token", + "code_challenge": "abc", + }, + follow_redirects=False, + ) + assert resp.status_code == 400 + + +# --------------------------------------------------------------------------- +# Authorization POST endpoint +# --------------------------------------------------------------------------- + + +def test_authorize_post_creates_code_and_redirects(client): + verifier = "my_code_verifier_that_is_long_enough_43chars" + challenge = _make_challenge(verifier) + + resp = client.post( + "/v1/mcp/oauth/authorize", + data={ + "client_id": "user-123", + "redirect_uri": "https://client.example.com/callback", + "code_challenge": challenge, + "code_challenge_method": "S256", + "state": "st_abc", + "server_id": "server-xyz", + "api_key": "sk-supersecretkey", + }, + follow_redirects=False, + ) + assert resp.status_code == 302 + location = resp.headers["location"] + assert "code=" in location + assert "st_abc" in location + + # Extract the code from the redirect URL + from urllib.parse import parse_qs, urlparse + + qs = parse_qs(urlparse(location).query) + code = qs["code"][0] + assert code in _byok_auth_codes + entry = _byok_auth_codes[code] + assert entry["api_key"] == "sk-supersecretkey" + assert entry["server_id"] == "server-xyz" + assert entry["user_id"] == "user-123" + assert entry["code_challenge"] == challenge + + +def test_authorize_post_unsupported_method(client): + resp = client.post( + "/v1/mcp/oauth/authorize", + data={ + "client_id": "u", + "redirect_uri": "https://example.com/cb", + "code_challenge": "abc", + "code_challenge_method": "plain", + "state": "", + "server_id": "s", + "api_key": "key", + }, + follow_redirects=False, + ) + assert resp.status_code == 400 + + +# --------------------------------------------------------------------------- +# Token endpoint +# --------------------------------------------------------------------------- + + +def _insert_code( + api_key: str, + server_id: str, + user_id: str, + challenge: str, + redirect_uri: str, + ttl: int = 300, +) -> str: + code = str(uuid.uuid4()) + _byok_auth_codes[code] = { + "api_key": api_key, + "server_id": server_id, + "user_id": user_id, + "code_challenge": challenge, + "redirect_uri": redirect_uri, + "expires_at": time.time() + ttl, + } + return code + + +@pytest.mark.asyncio +async def test_token_endpoint_success(): + """Happy path: valid code + PKCE → credential stored → JWT returned.""" + verifier = "my_test_code_verifier_value_long_enough_yes" + challenge = _make_challenge(verifier) + code = _insert_code( + api_key="sk-myapikey", + server_id="server-1", + user_id="user-42", + challenge=challenge, + redirect_uri="https://example.com/cb", + ) + + mock_prisma = MagicMock() + mock_store = AsyncMock() + test_master_key = "test_master_key_value" + + with patch( + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.store_user_credential", + mock_store, + ), patch( + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.router", + ): + # Import the actual handler function directly + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( + byok_token, + ) + + mock_request = MagicMock() + # Patch module-level globals in the function's module + with patch( + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.store_user_credential", + mock_store, + ): + import litellm.proxy._experimental.mcp_server.byok_oauth_endpoints as mod + + original_prisma = None + original_master_key = None + + # Temporarily inject our test values + with patch( + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), patch("litellm.proxy.proxy_server.master_key", test_master_key): + result = await byok_token( + request=mock_request, + grant_type="authorization_code", + code=code, + redirect_uri="https://example.com/cb", + code_verifier=verifier, + client_id="user-42", + ) + + assert result.status_code == 200 + body = result.body + import json + + data = json.loads(body) + assert "access_token" in data + assert data["token_type"] == "bearer" + assert data["expires_in"] == 3600 + + # Verify JWT payload + import jwt as pyjwt + + payload = pyjwt.decode( + data["access_token"], test_master_key, algorithms=["HS256"] + ) + assert payload["user_id"] == "user-42" + assert payload["server_id"] == "server-1" + assert payload["type"] == "byok_session" + + # Auth code was consumed + assert code not in _byok_auth_codes + + # store_user_credential was called + mock_store.assert_awaited_once_with( + prisma_client=mock_prisma, + user_id="user-42", + server_id="server-1", + credential="sk-myapikey", + ) + + +@pytest.mark.asyncio +async def test_token_endpoint_invalid_code(): + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import byok_token + + mock_request = MagicMock() + with pytest.raises(HTTPException) as exc_info: + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( + "litellm.proxy.proxy_server.master_key", "key" + ): + await byok_token( + request=mock_request, + grant_type="authorization_code", + code="nonexistent-code", + redirect_uri="", + code_verifier="anything", + client_id="u", + ) + assert exc_info.value.status_code == 400 + assert "invalid_grant" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_token_endpoint_expired_code(): + verifier = "exp_verifier_that_is_long_enough_to_be_valid" + challenge = _make_challenge(verifier) + code = _insert_code( + api_key="key", + server_id="s", + user_id="u", + challenge=challenge, + redirect_uri="https://cb", + ttl=-10, # already expired + ) + + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import byok_token + + mock_request = MagicMock() + with pytest.raises(HTTPException) as exc_info: + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( + "litellm.proxy.proxy_server.master_key", "key" + ): + await byok_token( + request=mock_request, + grant_type="authorization_code", + code=code, + redirect_uri="", + code_verifier=verifier, + client_id="u", + ) + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_token_endpoint_wrong_verifier(): + verifier = "correct_verifier_value_that_is_long_enough" + challenge = _make_challenge(verifier) + code = _insert_code( + api_key="key", + server_id="s", + user_id="u", + challenge=challenge, + redirect_uri="https://cb", + ) + + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import byok_token + + mock_request = MagicMock() + with pytest.raises(HTTPException) as exc_info: + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( + "litellm.proxy.proxy_server.master_key", "key" + ): + await byok_token( + request=mock_request, + grant_type="authorization_code", + code=code, + redirect_uri="", + code_verifier="wrong_verifier_value_that_wont_match", + client_id="u", + ) + assert exc_info.value.status_code == 400 + assert "invalid_grant" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_token_endpoint_unsupported_grant_type(): + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import byok_token + + mock_request = MagicMock() + with pytest.raises(HTTPException) as exc_info: + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( + "litellm.proxy.proxy_server.master_key", "key" + ): + await byok_token( + request=mock_request, + grant_type="client_credentials", + code="any", + redirect_uri="", + code_verifier="v", + client_id="u", + ) + assert exc_info.value.status_code == 400 + assert "unsupported_grant_type" in str(exc_info.value.detail) + + +# --------------------------------------------------------------------------- +# _check_byok_credential (the 401 challenge in execute_mcp_tool) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_check_byok_credential_not_byok(): + """Non-BYOK servers should pass through without any DB check.""" + from litellm.proxy._experimental.mcp_server.server import _check_byok_credential + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="s1", + name="normal-server", + transport=MCPTransport.http, + is_byok=False, + ) + # Should not raise + await _check_byok_credential(server, None) + + +@pytest.mark.asyncio +async def test_check_byok_credential_no_user_id(): + """BYOK server with no user identity → 401.""" + from litellm.proxy._experimental.mcp_server.server import _check_byok_credential + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="byok-1", + name="byok-server", + transport=MCPTransport.http, + is_byok=True, + ) + with pytest.raises(HTTPException) as exc_info: + await _check_byok_credential(server, None) + + assert exc_info.value.status_code == 401 + assert "WWW-Authenticate" in (exc_info.value.headers or {}) # type: ignore[operator] + assert "byok_auth_required" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_check_byok_credential_missing_credential(): + """BYOK server with a known user but no stored credential → 401.""" + from litellm.proxy._experimental.mcp_server.server import _check_byok_credential + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="byok-2", + name="byok-server", + transport=MCPTransport.http, + is_byok=True, + ) + user_auth = UserAPIKeyAuth(user_id="user-99", api_key="sk-test") + + mock_prisma = MagicMock() + + with patch( + "litellm.proxy._experimental.mcp_server.db.has_user_credential", + new=AsyncMock(return_value=False), + ), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with pytest.raises(HTTPException) as exc_info: + await _check_byok_credential(server, user_auth) + + assert exc_info.value.status_code == 401 + detail: Any = exc_info.value.detail + assert detail["error"] == "byok_auth_required" + assert detail["server_id"] == "byok-2" + headers = exc_info.value.headers or {} + assert "WWW-Authenticate" in headers # type: ignore[operator] + assert "oauth-protected-resource" in headers["WWW-Authenticate"] # type: ignore[index] + + +@pytest.mark.asyncio +async def test_check_byok_credential_has_credential(): + """BYOK server with a valid stored credential → no error raised.""" + from litellm.proxy._experimental.mcp_server.server import _check_byok_credential + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="byok-3", + name="byok-server", + transport=MCPTransport.http, + is_byok=True, + ) + user_auth = UserAPIKeyAuth(user_id="user-77", api_key="sk-test") + + mock_prisma = MagicMock() + + with patch( + "litellm.proxy._experimental.mcp_server.db.has_user_credential", + new=AsyncMock(return_value=True), + ), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Should not raise + await _check_byok_credential(server, user_auth) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx b/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx index b3829d0a8f4..397e55aa4ec 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx @@ -111,6 +111,8 @@ const routeFor = (slug: string): string => { return "tools/mcp-servers"; case "vector-stores": return "tools/vector-stores"; + case "byok-demo": + return "tools/byok-demo"; // experimental case "caching": @@ -226,6 +228,12 @@ const menuItems: MenuItemCfg[] = [ icon: , roles: all_admin_roles, }, + { + key: "29", + page: "byok-demo", + label: "BYOK Demo", + icon: , + }, ], }, { 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 new file mode 100644 index 00000000000..2ca66e0c51c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/tools/byok-demo/page.tsx @@ -0,0 +1,800 @@ +"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-3GWHDBM9B37bBIsl3dhuAg"; +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/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. +

+
+
+
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx b/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx new file mode 100644 index 00000000000..87117afc1cd --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx @@ -0,0 +1,254 @@ +"use client"; + +import React, { useState } from "react"; +import { Modal, Input, Switch, message } from "antd"; +import { + KeyOutlined, + LockOutlined, + CheckOutlined, + ArrowRightOutlined, + ArrowLeftOutlined, + CloseOutlined, + LinkOutlined, +} from "@ant-design/icons"; +import { MCPServer } from "./types"; + +interface ByokCredentialModalProps { + server: MCPServer; + open: boolean; + onClose: () => void; + onSuccess: (serverId: string) => void; + accessToken: string; +} + +export const ByokCredentialModal: React.FC = ({ + server, + open, + onClose, + onSuccess, + accessToken, +}) => { + const [step, setStep] = useState<1 | 2>(1); + const [apiKey, setApiKey] = useState(""); + const [saveKey, setSaveKey] = useState(true); + const [loading, setLoading] = useState(false); + + const serverDisplayName = server.alias || server.server_name || "Service"; + const firstLetter = serverDisplayName.charAt(0).toUpperCase(); + + const handleClose = () => { + setStep(1); + setApiKey(""); + setSaveKey(true); + setLoading(false); + onClose(); + }; + + const handleAuthorize = async () => { + if (!apiKey.trim()) { + message.error("Please enter your API key"); + return; + } + setLoading(true); + try { + const response = await fetch(`/v1/mcp/server/${server.server_id}/user-credential`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${accessToken}`, + }, + body: JSON.stringify({ credential: apiKey.trim(), save: saveKey }), + }); + if (!response.ok) { + const err = await response.json(); + throw new Error(err?.detail?.error || "Failed to save credential"); + } + message.success(`Connected to ${serverDisplayName}`); + onSuccess(server.server_id); + handleClose(); + } catch (e: any) { + message.error(e.message || "Failed to connect"); + } finally { + setLoading(false); + } + }; + + return ( + +
+ {/* Step dots + close */} +
+ {step === 2 ? ( + + ) : ( +
+ )} +
+
+
+
+ +
+ + {step === 1 ? ( +
+ {/* Logos */} +
+
+ L +
+ +
+ {firstLetter} +
+
+ +

Connect {serverDisplayName}

+

+ LiteLLM needs access to {serverDisplayName} to complete your request. +

+ + {/* How it works */} +
+
+
+ + + + +
+
+

How it works

+

+ LiteLLM acts as a secure bridge. Your requests are routed through our MCP client directly to{" "} + {serverDisplayName}'s API. +

+
+
+
+ + {/* Requested access */} + {server.byok_description && server.byok_description.length > 0 && ( +
+

+ + + + + Requested Access +

+
    + {server.byok_description.map((item, i) => ( +
  • + + {item} +
  • + ))} +
+
+ )} + + + +
+ ) : ( +
+ {/* Key icon */} +
+ +
+ +

Provide API Key

+

+ Enter your {serverDisplayName} API key to authorize this connection. +

+ +
+ + setApiKey(e.target.value)} + size="large" + className="rounded-lg" + /> + {server.byok_api_key_help_url && ( + + Where do I find my API key? + + )} +
+ + {/* Save toggle */} +
+
+ + + + Save key for future use +
+ +
+ + {/* Security note */} +
+ +

+ Your key is encrypted at rest and transmitted securely. It is never shared with third parties. +

+
+ + +
+ )} +
+ + ); +}; + +export default ByokCredentialModal; 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 6dbb18887da..6ca58ffae24 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 @@ -1,5 +1,5 @@ import React, { useState } from "react"; -import { Modal, Tooltip, Form, Select, Input } from "antd"; +import { Modal, Tooltip, Form, Select, Input, Switch } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, TextInput } from "@tremor/react"; import { createMCPServer } from "../networking"; @@ -624,6 +624,89 @@ const CreateMCPServer: React.FC = ({ )} + {/* BYOK toggle - only for OpenAPI */} + {transportType === TRANSPORT.OPENAPI && ( + <> + + BYOK (Bring Your Own Key) + + + + + } + name="is_byok" + valuePropName="checked" + > + + + + prev.is_byok !== cur.is_byok || prev.auth_type !== cur.auth_type}> + {({ getFieldValue }) => + getFieldValue("is_byok") ? ( + <> + {/* Auth format hint */} + {getFieldValue("auth_type") && getFieldValue("auth_type") !== "none" && ( +
+ + + User keys will be sent as:{" "} + + {getFieldValue("auth_type") === "bearer_token" && "Authorization: Bearer {key}"} + {getFieldValue("auth_type") === "api_key" && "x-api-key: {key}"} + {getFieldValue("auth_type") === "basic" && "Authorization: Basic {key}"} + {getFieldValue("auth_type") === "authorization" && "Authorization: {key}"} + + {!getFieldValue("auth_type") && "Set Authentication Type below to specify the format."} + +
+ )} + {!getFieldValue("auth_type") && ( +
+ + Set the Authentication Type below to specify how user keys are sent (e.g., Bearer Token, API Key header). +
+ )} + + Access Description + + + + + } + name="byok_description" + > + + + + ) : null + } +
+ + )} + {/* Authentication - show for HTTP, SSE, and OpenAPI */} {transportType !== "stdio" && transportType !== "" && ( void, onDelete: (serverId: string) => void, isLoadingHealth?: boolean, + onByokConnect?: (server: MCPServer) => void, ): ColumnDef[] => [ { accessorKey: "server_id", @@ -192,6 +194,41 @@ export const mcpServerColumns = ( ); }, }, + { + id: "byok_credential", + header: "Credential", + cell: ({ row }) => { + const server = row.original; + if (!server.is_byok) { + return —; + } + if (server.has_user_credential) { + return ( +
+ + Connected + + {onByokConnect && ( + + )} +
+ ); + } + return onByokConnect ? ( + + ) : null; + }, + }, { id: "actions", header: "Actions", 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 0f87f5e87b8..f48649d6653 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -16,6 +16,7 @@ import { DiscoverableMCPServer, MCPServer, MCPServerProps, Team } from "./types" import MCPSemanticFilterSettings from "../Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings"; import MCPNetworkSettings from "./MCPNetworkSettings"; import MCPDiscovery from "./mcp_discovery"; +import { ByokCredentialModal } from "./ByokCredentialModal"; const { Text: AntdText, Title: AntdTitle } = Typography; const EDIT_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-edit-state"; @@ -70,6 +71,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) const [isDiscoveryVisible, setDiscoveryVisible] = useState(false); const [prefillData, setPrefillData] = useState(null); const [isDeletingServer, setIsDeletingServer] = useState(false); + const [byokModalServer, setByokModalServer] = useState(null); const isInternalUser = userRole === "Internal User"; useEffect(() => { @@ -170,6 +172,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) }, handleDelete, isLoadingHealth, + (server: MCPServer) => setByokModalServer(server), ), [userRole, isLoadingHealth], ); @@ -427,6 +430,19 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) + + {byokModalServer && ( + setByokModalServer(null)} + onSuccess={(_serverId) => { + refetch(); + setByokModalServer(null); + }} + accessToken={accessToken || ""} + /> + )}
); }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 8a08f13e22a..6ba25012197 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -178,6 +178,12 @@ export interface MCPServer { command?: string | null; args?: string[] | null; env?: Record | null; + + /** BYOK (Bring Your Own Key) fields */ + is_byok?: boolean | null; + byok_description?: string[] | null; + byok_api_key_help_url?: string | null; + has_user_credential?: boolean | null; } export interface MCPServerProps { diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx index 3272e9b589a..9936f34452d 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx @@ -33,6 +33,7 @@ import GuardrailSelector from "../../guardrails/GuardrailSelector"; import PolicySelector from "../../policies/PolicySelector"; import MCPToolArgumentsForm, { MCPToolArgumentsFormRef } from "../../mcp_tools/MCPToolArgumentsForm"; import { MCPServer } from "../../mcp_tools/types"; +import { ByokCredentialModal } from "../../mcp_tools/ByokCredentialModal"; import NotificationsManager from "../../molecules/notifications_manager"; import { callMCPTool, fetchMCPServers, listMCPTools } from "../../networking"; import TagSelector from "../../tag_management/TagSelector"; @@ -108,6 +109,7 @@ const ChatUI: React.FC = ({ fixedModel, }) => { const [mcpServers, setMCPServers] = useState([]); + const [byokModalServer, setByokModalServer] = useState(null); const [selectedMCPServers, setSelectedMCPServers] = useState(() => { const saved = sessionStorage.getItem("selectedMCPServers"); try { @@ -1746,6 +1748,49 @@ const ChatUI: React.FC = ({ })}
)} + + {/* BYOK credential status for selected servers */} + {selectedMCPServers.length > 0 && + !selectedMCPServers.includes("__all__") && + selectedMCPServers.some((serverId) => { + const server = mcpServers.find((s) => s.server_id === serverId); + return server?.is_byok; + }) && ( +
+ {selectedMCPServers.map((serverId) => { + const server = mcpServers.find((s) => s.server_id === serverId); + if (!server?.is_byok) return null; + const serverName = server.alias || server.server_name || serverId; + return ( +
+ + {serverName} requires your API key + + {server.has_user_credential ? ( +
+ + Connected + + +
+ ) : ( + + )} +
+ ); + })} +
+ )}
@@ -2498,6 +2543,20 @@ const ChatUI: React.FC = ({ {generatedCode} + + {byokModalServer && ( + setByokModalServer(null)} + onSuccess={(_serverId) => { + // Refresh MCP servers to pick up updated has_user_credential + loadMCPServers(); + setByokModalServer(null); + }} + accessToken={accessToken || ""} + /> + )}
); };