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
This commit is contained in:
Ishaan Jaffer 2026-03-04 19:53:43 -08:00
parent 6e59fe839d
commit 37b87789e5
19 changed files with 2379 additions and 2 deletions

View file

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

View file

@ -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 = """<!DOCTYPE html>
<html>
<head><title>Connect to {server_name} — LiteLLM</title>
<style>
body {{ font-family: system-ui; background: #0f172a; display: flex; justify-content: center; align-items: center; height: 100vh; margin: 0; }}
.card {{ background: #1e293b; border-radius: 12px; padding: 32px; width: 400px; color: white; }}
h2 {{ margin: 0 0 8px; font-size: 20px; }}
p {{ color: #94a3b8; margin: 0 0 24px; font-size: 14px; }}
label {{ font-size: 13px; color: #cbd5e1; display: block; margin-bottom: 6px; }}
input[type=password] {{ width: 100%; padding: 10px; border-radius: 8px; border: 1px solid #334155; background: #0f172a; color: white; font-size: 14px; box-sizing: border-box; }}
button {{ width: 100%; margin-top: 20px; padding: 12px; background: #3b82f6; border: none; border-radius: 8px; color: white; font-size: 15px; cursor: pointer; }}
button:hover {{ background: #2563eb; }}
.note {{ font-size: 12px; color: #64748b; margin-top: 16px; text-align: center; }}
</style></head>
<body>
<div class="card">
<h2>Connect to {server_name}</h2>
<p>Enter your {server_name} API key to authorize this connection.</p>
<form method="POST">
<input type="hidden" name="client_id" value="{client_id}">
<input type="hidden" name="redirect_uri" value="{redirect_uri}">
<input type="hidden" name="code_challenge" value="{code_challenge}">
<input type="hidden" name="code_challenge_method" value="{code_challenge_method}">
<input type="hidden" name="state" value="{state}">
<input type="hidden" name="server_id" value="{server_id}">
<label>{server_name} API Key</label>
<input type="password" name="api_key" placeholder="Enter your API key" required autofocus>
<button type="submit">Connect &amp; Authorize</button>
</form>
<p class="note">Your key is encrypted at rest and never shared with third parties.</p>
</div>
</body>
</html>"""
# ---------------------------------------------------------------------------
# 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,
}
)

View file

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

View file

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

View file

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

View file

@ -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 ########

View file

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

View file

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

View file

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

View file

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

View file

@ -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 &amp; 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)

View file

@ -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: <DatabaseOutlined style={{ fontSize: 18 }} />,
roles: all_admin_roles,
},
{
key: "29",
page: "byok-demo",
label: "BYOK Demo",
icon: <KeyOutlined style={{ fontSize: 18 }} />,
},
],
},
{

View file

@ -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 (
<svg
className={className}
xmlns="http://www.w3.org/2000/svg"
viewBox="0 0 24 24"
fill="none"
stroke="currentColor"
strokeWidth={2}
strokeLinecap="round"
strokeLinejoin="round"
>
<rect x="3" y="11" width="18" height="11" rx="2" ry="2" />
<path d="M7 11V7a5 5 0 0 1 10 0v4" />
</svg>
);
}
function CheckIcon({ className }: { className?: string }) {
return (
<svg
className={className}
xmlns="http://www.w3.org/2000/svg"
viewBox="0 0 24 24"
fill="none"
stroke="currentColor"
strokeWidth={2.5}
strokeLinecap="round"
strokeLinejoin="round"
>
<polyline points="20 6 9 17 4 12" />
</svg>
);
}
function ServerIcon({ className }: { className?: string }) {
return (
<svg
className={className}
xmlns="http://www.w3.org/2000/svg"
viewBox="0 0 24 24"
fill="none"
stroke="currentColor"
strokeWidth={2}
strokeLinecap="round"
strokeLinejoin="round"
>
<rect x="2" y="2" width="20" height="8" rx="2" ry="2" />
<rect x="2" y="14" width="20" height="8" rx="2" ry="2" />
<line x1="6" y1="6" x2="6.01" y2="6" />
<line x1="6" y1="18" x2="6.01" y2="18" />
</svg>
);
}
function KeyIcon({ className }: { className?: string }) {
return (
<svg
className={className}
xmlns="http://www.w3.org/2000/svg"
viewBox="0 0 24 24"
fill="none"
stroke="currentColor"
strokeWidth={2}
strokeLinecap="round"
strokeLinejoin="round"
>
<circle cx="7.5" cy="15.5" r="5.5" />
<path d="M21 2L11 12" />
<path d="M15 6l1 1" />
</svg>
);
}
function SpinnerIcon({ className }: { className?: string }) {
return (
<svg
className={className}
xmlns="http://www.w3.org/2000/svg"
fill="none"
viewBox="0 0 24 24"
>
<circle
className="opacity-25"
cx="12"
cy="12"
r="10"
stroke="currentColor"
strokeWidth="4"
/>
<path
className="opacity-75"
fill="currentColor"
d="M4 12a8 8 0 018-8V0C5.373 0 0 5.373 0 12h4zm2 5.291A7.962 7.962 0 014 12H0c0 3.042 1.135 5.824 3 7.938l3-2.647z"
/>
</svg>
);
}
// ---------------------------------------------------------------------------
// Main page component
// ---------------------------------------------------------------------------
export default function ByokDemoPage() {
const [servers, setServers] = useState<McpServer[]>([]);
const [loadingServers, setLoadingServers] = useState(true);
const [fetchError, setFetchError] = useState<string | null>(null);
const [connectionStatus, setConnectionStatus] = useState<
Record<string, ServerConnectionStatus>
>({});
const [chatMessages, setChatMessages] = useState<ChatMessage[]>([
{
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<Record<string, ReturnType<typeof setInterval>>>({});
// ---------------------------------------------------------------------------
// 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 (
<div
className={`rounded-lg p-3 mb-2 border transition-colors ${
isConnected
? "border-emerald-500/40 bg-emerald-900/20"
: hasError
? "border-red-500/40 bg-red-900/10"
: "border-slate-700 bg-slate-800/50"
}`}
>
<div className="flex items-start gap-2">
<div className="mt-0.5 flex-shrink-0">
{server.is_byok ? (
isConnected ? (
<CheckIcon className="w-4 h-4 text-emerald-400" />
) : (
<LockIcon className="w-4 h-4 text-amber-400" />
)
) : (
<ServerIcon className="w-4 h-4 text-slate-400" />
)}
</div>
<div className="flex-1 min-w-0">
<div className="text-sm font-medium text-slate-200 truncate">
{server.server_name}
</div>
{server.description && (
<div className="text-xs text-slate-500 mt-0.5 truncate">
{server.description}
</div>
)}
<div className="flex items-center gap-2 mt-1.5">
{server.is_byok && (
<span className="inline-flex items-center px-1.5 py-0.5 rounded text-[10px] font-medium bg-amber-900/40 text-amber-300 border border-amber-700/50">
BYOK
</span>
)}
{isConnected && (
<span className="inline-flex items-center gap-1 px-1.5 py-0.5 rounded text-[10px] font-medium bg-emerald-900/40 text-emerald-300 border border-emerald-700/50">
<CheckIcon className="w-2.5 h-2.5" />
Connected
</span>
)}
{hasError && (
<span className="inline-flex items-center px-1.5 py-0.5 rounded text-[10px] font-medium bg-red-900/40 text-red-300 border border-red-700/50">
Error
</span>
)}
</div>
{hasError && connStatus?.errorMessage && (
<div className="text-[11px] text-red-400 mt-1.5 leading-snug">
{connStatus.errorMessage}
</div>
)}
</div>
</div>
{server.is_byok && !isConnected && (
<button
onClick={() => handleConnect(server)}
disabled={isConnecting}
className={`mt-2.5 w-full flex items-center justify-center gap-1.5 px-3 py-1.5 rounded-md text-xs font-medium transition-colors ${
isConnecting
? "bg-slate-700 text-slate-400 cursor-not-allowed"
: "bg-indigo-600 hover:bg-indigo-500 text-white"
}`}
>
{isConnecting ? (
<>
<SpinnerIcon className="w-3 h-3 animate-spin" />
Connecting…
</>
) : (
<>
<KeyIcon className="w-3 h-3" />
Connect
</>
)}
</button>
)}
{server.is_byok && isConnected && !hasError && (
<button
onClick={() => {
setConnectionStatus((prev) => ({
...prev,
[server.server_id]: { state: "idle" },
}));
handleConnect(server);
}}
className="mt-2.5 w-full flex items-center justify-center gap-1.5 px-3 py-1.5 rounded-md text-xs font-medium text-slate-400 hover:text-slate-200 border border-slate-700 hover:border-slate-500 transition-colors"
>
Reconnect
</button>
)}
</div>
);
}
function ChatBubble({ message }: { message: ChatMessage }) {
const isUser = message.role === "user";
const isSystem = message.role === "system";
if (isSystem) {
return (
<div className="flex justify-center mb-4">
<div className="bg-slate-800 border border-slate-700 rounded-lg px-4 py-3 max-w-lg text-sm text-slate-400 text-center">
{message.content}
</div>
</div>
);
}
const isSuccess = message.content.startsWith("Connected to ");
return (
<div className={`flex mb-4 ${isUser ? "justify-end" : "justify-start"}`}>
{!isUser && (
<div className="w-7 h-7 rounded-full bg-indigo-600 flex items-center justify-center flex-shrink-0 mr-2 mt-0.5">
<span className="text-xs font-bold text-white">L</span>
</div>
)}
<div
className={`rounded-2xl px-4 py-2.5 max-w-md text-sm leading-relaxed ${
isUser
? "bg-indigo-600 text-white rounded-br-sm"
: isSuccess
? "bg-emerald-900/40 border border-emerald-700/50 text-emerald-200 rounded-bl-sm"
: "bg-slate-800 border border-slate-700 text-slate-200 rounded-bl-sm"
}`}
>
{isSuccess && (
<div className="flex items-center gap-1.5 mb-1.5">
<span className="text-emerald-400 text-base">✅</span>
<span className="text-emerald-300 font-medium text-xs uppercase tracking-wide">
Connected
</span>
</div>
)}
{message.content}
</div>
{isUser && (
<div className="w-7 h-7 rounded-full bg-slate-600 flex items-center justify-center flex-shrink-0 ml-2 mt-0.5">
<span className="text-xs font-bold text-white">A</span>
</div>
)}
</div>
);
}
// ---------------------------------------------------------------------------
// Render
// ---------------------------------------------------------------------------
return (
<div className="flex h-screen bg-[#0f172a] text-slate-100 overflow-hidden">
{/* ------------------------------------------------------------------ */}
{/* Left sidebar */}
{/* ------------------------------------------------------------------ */}
<aside className="w-72 flex-shrink-0 flex flex-col border-r border-slate-800 bg-[#0d1526]">
{/* Header */}
<div className="px-4 pt-5 pb-4 border-b border-slate-800">
<h2 className="text-xs font-semibold uppercase tracking-widest text-slate-500 mb-1">
MCP Tools
</h2>
<p className="text-[11px] text-slate-600">
via {PROXY_BASE_URL.replace(/https?:\/\//, "")}
</p>
</div>
{/* Server list */}
<div className="flex-1 overflow-y-auto px-3 py-3">
{loadingServers ? (
<div className="flex flex-col items-center justify-center gap-2 py-10 text-slate-600">
<SpinnerIcon className="w-6 h-6 animate-spin" />
<span className="text-xs">Loading servers…</span>
</div>
) : fetchError ? (
<div className="rounded-lg border border-red-800/50 bg-red-900/10 p-3">
<div className="text-xs font-medium text-red-400 mb-1">
Could not fetch servers
</div>
<div className="text-[11px] text-red-500 leading-snug">{fetchError}</div>
<button
onClick={fetchServers}
className="mt-2 text-[11px] text-indigo-400 hover:text-indigo-300 underline"
>
Retry
</button>
</div>
) : (
<>
{byokServers.length > 0 && (
<div className="mb-4">
<div className="text-[10px] font-semibold uppercase tracking-widest text-amber-500/80 mb-2 px-1">
Requires your key
</div>
{byokServers.map((s) => (
<ServerItem key={s.server_id} server={s} />
))}
</div>
)}
{regularServers.length > 0 && (
<div className="mb-4">
<div className="text-[10px] font-semibold uppercase tracking-widest text-slate-500 mb-2 px-1">
Available
</div>
{regularServers.map((s) => (
<ServerItem key={s.server_id} server={s} />
))}
</div>
)}
{servers.length === 0 && (
<div className="text-center py-10 text-slate-600 text-xs">
No MCP servers found.
</div>
)}
</>
)}
</div>
{/* Footer: virtual key info */}
<div className="px-4 py-3 border-t border-slate-800">
<div className="flex items-center gap-2">
<KeyIcon className="w-3.5 h-3.5 text-slate-500 flex-shrink-0" />
<span className="text-[11px] text-slate-500 font-mono truncate">
{truncatedKey}
</span>
</div>
<div className="text-[10px] text-slate-700 mt-0.5">Demo virtual key</div>
</div>
</aside>
{/* ------------------------------------------------------------------ */}
{/* Main content area */}
{/* ------------------------------------------------------------------ */}
<div className="flex-1 flex flex-col overflow-hidden">
{/* Top bar */}
<header className="flex items-center justify-between px-6 py-3.5 border-b border-slate-800 bg-[#0d1526] flex-shrink-0">
<div className="flex items-center gap-3">
<div className="w-7 h-7 rounded-lg bg-indigo-600 flex items-center justify-center">
<span className="text-sm font-bold text-white">L</span>
</div>
<div>
<h1 className="text-sm font-semibold text-slate-100 leading-tight">
LiteLLM MCP Demo
</h1>
<p className="text-[11px] text-slate-500 leading-tight">
External chat UI
</p>
</div>
</div>
<div className="flex items-center gap-2">
<span className="inline-flex items-center px-2.5 py-1 rounded-full text-[11px] font-medium bg-indigo-900/50 text-indigo-300 border border-indigo-700/50">
BYOK OAuth Flow
</span>
<span className="inline-flex items-center px-2.5 py-1 rounded-full text-[11px] font-mono bg-slate-800 text-slate-400 border border-slate-700">
{truncatedKey}
</span>
</div>
</header>
{/* Chat messages */}
<div className="flex-1 overflow-y-auto px-6 py-6">
{/* Explainer card */}
<div className="mb-6 rounded-xl border border-slate-700 bg-slate-800/50 p-5 max-w-2xl mx-auto">
<h3 className="text-sm font-semibold text-slate-200 mb-2">
How this demo works
</h3>
<ol className="text-xs text-slate-400 space-y-1.5 list-none">
{[
"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) => (
<li key={i} className="flex gap-2">
<span className="flex-shrink-0 w-4 h-4 rounded-full bg-indigo-900/60 border border-indigo-700/50 text-indigo-400 text-[9px] font-bold flex items-center justify-center mt-0.5">
{i + 1}
</span>
<span>{step}</span>
</li>
))}
</ol>
</div>
{/* Chat messages */}
<div className="max-w-2xl mx-auto">
{chatMessages.map((msg, idx) => (
<ChatBubble key={idx} message={msg} />
))}
</div>
</div>
{/* Chat input (UI only — no LLM call in this demo) */}
<div className="px-6 py-4 border-t border-slate-800 flex-shrink-0">
<div className="max-w-2xl mx-auto">
<div className="flex gap-2">
<input
type="text"
placeholder="Type a message… (demo — not connected to an LLM)"
className="flex-1 bg-slate-800 border border-slate-700 rounded-xl px-4 py-2.5 text-sm text-slate-300 placeholder-slate-600 focus:outline-none focus:border-indigo-500 focus:ring-1 focus:ring-indigo-500/50 transition-colors"
onKeyDown={(e) => {
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 = "";
}
}}
/>
<button
className="px-4 py-2.5 bg-indigo-600 hover:bg-indigo-500 text-white text-sm font-medium rounded-xl transition-colors flex-shrink-0"
onClick={(e) => {
const input = (e.currentTarget.previousSibling as HTMLInputElement);
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 = "";
}}
>
Send
</button>
</div>
<p className="text-[11px] text-slate-700 mt-2 text-center">
Demo only — chat responses are simulated. MCP tool calls require a connected server.
</p>
</div>
</div>
</div>
</div>
);
}

View file

@ -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<ByokCredentialModalProps> = ({
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 (
<Modal
open={open}
onCancel={handleClose}
footer={null}
width={480}
closeIcon={null}
className="byok-modal"
>
<div className="relative p-2">
{/* Step dots + close */}
<div className="flex items-center justify-between mb-6">
{step === 2 ? (
<button
onClick={() => setStep(1)}
className="flex items-center gap-1 text-gray-500 hover:text-gray-800 text-sm"
>
<ArrowLeftOutlined /> Back
</button>
) : (
<div />
)}
<div className="flex items-center gap-1.5">
<div className={`w-2 h-2 rounded-full ${step === 1 ? "bg-blue-500" : "bg-gray-300"}`} />
<div className={`w-2 h-2 rounded-full ${step === 2 ? "bg-blue-500" : "bg-gray-300"}`} />
</div>
<button onClick={handleClose} className="text-gray-400 hover:text-gray-600">
<CloseOutlined />
</button>
</div>
{step === 1 ? (
<div className="text-center">
{/* Logos */}
<div className="flex items-center justify-center gap-3 mb-6">
<div className="w-14 h-14 rounded-xl bg-gradient-to-br from-teal-400 to-cyan-600 flex items-center justify-center text-white font-bold text-xl shadow">
L
</div>
<ArrowRightOutlined className="text-gray-400 text-lg" />
<div className="w-14 h-14 rounded-xl bg-gradient-to-br from-blue-600 to-indigo-800 flex items-center justify-center text-white font-bold text-xl shadow">
{firstLetter}
</div>
</div>
<h2 className="text-2xl font-bold text-gray-900 mb-2">Connect {serverDisplayName}</h2>
<p className="text-gray-500 mb-6">
LiteLLM needs access to {serverDisplayName} to complete your request.
</p>
{/* How it works */}
<div className="bg-gray-50 rounded-xl p-4 text-left mb-4">
<div className="flex items-start gap-3">
<div className="mt-0.5">
<svg width="20" height="20" viewBox="0 0 24 24" fill="none" className="text-gray-500">
<rect x="2" y="4" width="20" height="16" rx="2" stroke="currentColor" strokeWidth="2" />
<path d="M8 4v16M16 4v16" stroke="currentColor" strokeWidth="2" />
</svg>
</div>
<div>
<p className="font-semibold text-gray-800 mb-1">How it works</p>
<p className="text-gray-500 text-sm">
LiteLLM acts as a secure bridge. Your requests are routed through our MCP client directly to{" "}
{serverDisplayName}&apos;s API.
</p>
</div>
</div>
</div>
{/* Requested access */}
{server.byok_description && server.byok_description.length > 0 && (
<div className="bg-gray-50 rounded-xl p-4 text-left mb-6">
<p className="text-xs font-semibold text-gray-500 uppercase tracking-widest mb-3 flex items-center gap-2">
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" className="text-green-500">
<path
d="M12 2L12 22M2 12L22 12"
stroke="currentColor"
strokeWidth="2"
strokeLinecap="round"
/>
<circle cx="12" cy="12" r="9" stroke="currentColor" strokeWidth="2" />
</svg>
Requested Access
</p>
<ul className="space-y-2">
{server.byok_description.map((item, i) => (
<li key={i} className="flex items-center gap-2 text-sm text-gray-700">
<CheckOutlined className="text-green-500 flex-shrink-0" />
{item}
</li>
))}
</ul>
</div>
)}
<button
onClick={() => setStep(2)}
className="w-full bg-gray-900 hover:bg-gray-700 text-white font-medium py-3 px-6 rounded-xl flex items-center justify-center gap-2 transition-colors"
>
Continue to Authentication <ArrowRightOutlined />
</button>
<button
onClick={handleClose}
className="mt-3 w-full text-gray-400 hover:text-gray-600 text-sm py-2"
>
Cancel
</button>
</div>
) : (
<div>
{/* Key icon */}
<div className="w-12 h-12 rounded-full bg-blue-50 flex items-center justify-center mb-4">
<KeyOutlined className="text-blue-400 text-xl" />
</div>
<h2 className="text-2xl font-bold text-gray-900 mb-2">Provide API Key</h2>
<p className="text-gray-500 mb-6">
Enter your {serverDisplayName} API key to authorize this connection.
</p>
<div className="mb-4">
<label className="block text-sm font-semibold text-gray-800 mb-2">
{serverDisplayName} API Key
</label>
<Input.Password
placeholder="Enter your API key"
value={apiKey}
onChange={(e) => setApiKey(e.target.value)}
size="large"
className="rounded-lg"
/>
{server.byok_api_key_help_url && (
<a
href={server.byok_api_key_help_url}
target="_blank"
rel="noopener noreferrer"
className="text-blue-500 hover:text-blue-700 text-sm mt-2 flex items-center gap-1"
>
Where do I find my API key? <LinkOutlined />
</a>
)}
</div>
{/* Save toggle */}
<div className="bg-gray-50 rounded-xl p-4 flex items-center justify-between mb-4">
<div className="flex items-center gap-3">
<svg width="20" height="20" viewBox="0 0 24 24" fill="none" className="text-gray-500">
<path
d="M12 2C8.13 2 5 5.13 5 9c0 5.25 7 13 7 13s7-7.75 7-13c0-3.87-3.13-7-7-7zm0 9.5c-1.38 0-2.5-1.12-2.5-2.5s1.12-2.5 2.5-2.5 2.5 1.12 2.5 2.5-1.12 2.5-2.5 2.5z"
fill="currentColor"
/>
</svg>
<span className="text-sm font-medium text-gray-800">Save key for future use</span>
</div>
<Switch checked={saveKey} onChange={setSaveKey} />
</div>
{/* Security note */}
<div className="bg-blue-50 rounded-xl p-4 flex items-start gap-3 mb-6">
<LockOutlined className="text-blue-400 mt-0.5 flex-shrink-0" />
<p className="text-sm text-blue-700">
Your key is encrypted at rest and transmitted securely. It is never shared with third parties.
</p>
</div>
<button
onClick={handleAuthorize}
disabled={loading}
className="w-full bg-blue-500 hover:bg-blue-600 disabled:opacity-60 text-white font-medium py-3 px-6 rounded-xl flex items-center justify-center gap-2 transition-colors"
>
<LockOutlined /> Connect &amp; Authorize
</button>
</div>
)}
</div>
</Modal>
);
};
export default ByokCredentialModal;

View file

@ -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<CreateMCPServerProps> = ({
</Form.Item>
)}
{/* BYOK toggle - only for OpenAPI */}
{transportType === TRANSPORT.OPENAPI && (
<>
<Form.Item
label={
<span className="text-sm font-medium text-gray-700 flex items-center gap-2">
BYOK (Bring Your Own Key)
<Tooltip title="When enabled, each user provides their own API key for this service. Keys are stored per-user and never shared.">
<InfoCircleOutlined className="text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name="is_byok"
valuePropName="checked"
>
<Switch />
</Form.Item>
<Form.Item noStyle shouldUpdate={(prev, cur) => 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" && (
<div className="mb-4 p-3 bg-blue-50 rounded-lg text-sm text-blue-700 flex items-start gap-2">
<InfoCircleOutlined className="mt-0.5 flex-shrink-0" />
<span>
User keys will be sent as:{" "}
<code className="font-mono bg-blue-100 px-1 rounded">
{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}"}
</code>
{!getFieldValue("auth_type") && "Set Authentication Type below to specify the format."}
</span>
</div>
)}
{!getFieldValue("auth_type") && (
<div className="mb-4 p-3 bg-yellow-50 rounded-lg text-sm text-yellow-700 flex items-start gap-2">
<InfoCircleOutlined className="mt-0.5 flex-shrink-0" />
<span>Set the <strong>Authentication Type</strong> below to specify how user keys are sent (e.g., Bearer Token, API Key header).</span>
</div>
)}
<Form.Item
label={
<span className="text-sm font-medium text-gray-700">
Access Description
<Tooltip title="List of permissions shown to users in the connection modal (e.g. 'Create and manage Jira issues')">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name="byok_description"
>
<Select
mode="tags"
placeholder="Add access description items (press Enter after each)"
className="w-full"
tokenSeparators={[","]}
/>
</Form.Item>
<Form.Item
label={
<span className="text-sm font-medium text-gray-700">
API Key Help URL
<Tooltip title="Optional link shown to users to help them find their API key">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name="byok_api_key_help_url"
>
<Input placeholder="https://docs.example.com/api-keys" />
</Form.Item>
</>
) : null
}
</Form.Item>
</>
)}
{/* Authentication - show for HTTP, SSE, and OpenAPI */}
{transportType !== "stdio" && transportType !== "" && (
<Form.Item

View file

@ -4,6 +4,7 @@ import { Icon } from "@tremor/react";
import { PencilAltIcon, TrashIcon } from "@heroicons/react/outline";
import { getMaskedAndFullUrl } from "./utils";
import { Tooltip } from "antd";
import { CheckOutlined } from "@ant-design/icons";
export const mcpServerColumns = (
userRole: string,
@ -11,6 +12,7 @@ export const mcpServerColumns = (
onEdit: (serverId: string) => void,
onDelete: (serverId: string) => void,
isLoadingHealth?: boolean,
onByokConnect?: (server: MCPServer) => void,
): ColumnDef<MCPServer>[] => [
{
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 <span className="text-gray-300 text-xs">—</span>;
}
if (server.has_user_credential) {
return (
<div className="flex items-center gap-2">
<span className="text-green-600 text-xs font-medium flex items-center gap-1">
<CheckOutlined /> Connected
</span>
{onByokConnect && (
<button
className="text-xs text-gray-400 hover:text-blue-500 underline"
onClick={() => onByokConnect(server)}
>
Reconnect
</button>
)}
</div>
);
}
return onByokConnect ? (
<button
className="text-xs bg-blue-500 hover:bg-blue-600 text-white px-3 py-1 rounded-lg font-medium"
onClick={() => onByokConnect(server)}
>
Connect
</button>
) : null;
},
},
{
id: "actions",
header: "Actions",

View file

@ -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<MCPServerProps> = ({ accessToken, userRole, userID })
const [isDiscoveryVisible, setDiscoveryVisible] = useState(false);
const [prefillData, setPrefillData] = useState<DiscoverableMCPServer | null>(null);
const [isDeletingServer, setIsDeletingServer] = useState(false);
const [byokModalServer, setByokModalServer] = useState<MCPServer | null>(null);
const isInternalUser = userRole === "Internal User";
useEffect(() => {
@ -170,6 +172,7 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
},
handleDelete,
isLoadingHealth,
(server: MCPServer) => setByokModalServer(server),
),
[userRole, isLoadingHealth],
);
@ -427,6 +430,19 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
</TabPanel>
</TabPanels>
</TabGroup>
{byokModalServer && (
<ByokCredentialModal
server={byokModalServer}
open={!!byokModalServer}
onClose={() => setByokModalServer(null)}
onSuccess={(_serverId) => {
refetch();
setByokModalServer(null);
}}
accessToken={accessToken || ""}
/>
)}
</div>
);
};

View file

@ -178,6 +178,12 @@ export interface MCPServer {
command?: string | null;
args?: string[] | null;
env?: Record<string, string> | 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 {

View file

@ -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<ChatUIProps> = ({
fixedModel,
}) => {
const [mcpServers, setMCPServers] = useState<MCPServer[]>([]);
const [byokModalServer, setByokModalServer] = useState<MCPServer | null>(null);
const [selectedMCPServers, setSelectedMCPServers] = useState<string[]>(() => {
const saved = sessionStorage.getItem("selectedMCPServers");
try {
@ -1746,6 +1748,49 @@ const ChatUI: React.FC<ChatUIProps> = ({
})}
</div>
)}
{/* 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;
}) && (
<div className="mt-3 space-y-2">
{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 (
<div key={serverId} className="border border-blue-100 rounded p-2 bg-blue-50 flex items-center justify-between">
<Text className="text-xs text-blue-700">
{serverName} requires your API key
</Text>
{server.has_user_credential ? (
<div className="flex items-center gap-2">
<span className="text-green-600 text-xs font-medium flex items-center gap-1">
<KeyOutlined /> Connected
</span>
<button
className="text-xs text-gray-400 hover:text-blue-500 underline"
onClick={() => setByokModalServer(server)}
>
Reconnect
</button>
</div>
) : (
<button
className="text-xs bg-blue-500 hover:bg-blue-600 text-white px-3 py-1 rounded-lg font-medium"
onClick={() => setByokModalServer(server)}
>
Connect
</button>
)}
</div>
);
})}
</div>
)}
</div>
<div>
@ -2498,6 +2543,20 @@ const ChatUI: React.FC<ChatUIProps> = ({
{generatedCode}
</SyntaxHighlighter>
</Modal>
{byokModalServer && (
<ByokCredentialModal
server={byokModalServer}
open={!!byokModalServer}
onClose={() => setByokModalServer(null)}
onSuccess={(_serverId) => {
// Refresh MCP servers to pick up updated has_user_credential
loadMCPServers();
setByokModalServer(null);
}}
accessToken={accessToken || ""}
/>
)}
</div>
);
};