mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
6e59fe839d
commit
37b87789e5
19 changed files with 2379 additions and 2 deletions
12
CLAUDE.md
12
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
|
||||
- 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.
|
||||
332
litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py
Normal file
332
litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py
Normal 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 & 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,
|
||||
}
|
||||
)
|
||||
|
|
@ -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}}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ########
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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 }} />,
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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}'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 & Authorize
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</Modal>
|
||||
);
|
||||
};
|
||||
|
||||
export default ByokCredentialModal;
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue