diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py
index d4d327f0691..60cbc7cb560 100644
--- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py
+++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py
@@ -15,9 +15,11 @@ Endpoints implemented here:
import base64
import hashlib
+import html as _html_module
import time
import uuid
from typing import Dict, Optional, cast
+from urllib.parse import urlencode, urlparse
import jwt
from fastapi import APIRouter, Form, HTTPException, Request
@@ -79,9 +81,20 @@ def _build_authorize_html(
) -> str:
"""Build the 2-step BYOK OAuth authorization page HTML."""
+ # Escape all user-supplied / externally-derived values before interpolation
+ e = _html_module.escape
+ server_name = e(server_name)
+ server_initial = e(server_initial)
+ client_id = e(client_id)
+ redirect_uri = e(redirect_uri)
+ code_challenge = e(code_challenge)
+ code_challenge_method = e(code_challenge_method)
+ state = e(state)
+ server_id = e(server_id)
+
# Build access checklist rows
access_rows = "".join(
- f'
✓{item}
'
+ f'✓{e(item)}
'
for item in access_items
)
access_section = ""
@@ -98,7 +111,7 @@ def _build_authorize_html(
# Help link for step 2
help_link_html = ""
if help_url:
- help_link_html = f'Where do I find my API key? ↗'
+ help_link_html = f'Where do I find my API key? ↗'
return f"""
@@ -619,6 +632,11 @@ async def byok_authorize_post(
"""
_purge_expired_codes()
+ # Validate redirect_uri scheme to prevent open redirect
+ parsed_uri = urlparse(redirect_uri)
+ if parsed_uri.scheme not in ("http", "https"):
+ raise HTTPException(status_code=400, detail="Invalid redirect_uri scheme")
+
if code_challenge_method != "S256":
raise HTTPException(
status_code=400, detail="Only S256 code_challenge_method is supported"
@@ -634,8 +652,9 @@ async def byok_authorize_post(
"expires_at": time.time() + _AUTH_CODE_TTL_SECONDS,
}
+ params = urlencode({"code": auth_code, "state": state})
separator = "&" if "?" in redirect_uri else "?"
- location = f"{redirect_uri}{separator}code={auth_code}&state={state}"
+ location = f"{redirect_uri}{separator}{params}"
return RedirectResponse(url=location, status_code=302)
diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py
index 48225d6c93d..b9647733017 100644
--- a/litellm/proxy/_experimental/mcp_server/db.py
+++ b/litellm/proxy/_experimental/mcp_server/db.py
@@ -1,4 +1,3 @@
-import base64
from typing import Any, Dict, Iterable, List, Optional, Set, Union, cast
from litellm._logging import verbose_proxy_logger
@@ -14,6 +13,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
_get_salt_key,
+ decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.utils import PrismaClient
@@ -388,17 +388,17 @@ async def store_user_credential(
server_id: str,
credential: str,
) -> None:
- """Store a B64-encoded user credential for a BYOK MCP server."""
- credential_b64 = base64.b64encode(credential.encode()).decode()
+ """Encrypt and store a user credential for a BYOK MCP server."""
+ encrypted = encrypt_value_helper(value=credential, new_encryption_key=_get_salt_key())
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,
+ "credential_b64": encrypted,
},
- "update": {"credential_b64": credential_b64},
+ "update": {"credential_b64": encrypted},
},
)
@@ -408,13 +408,13 @@ async def get_user_credential(
user_id: str,
server_id: str,
) -> Optional[str]:
- """Return decoded credential for a user+server pair, or None."""
+ """Return decrypted 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()
+ return decrypt_value_helper(value=row.credential_b64, key="byok_credential")
async def has_user_credential(
diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py
index f61ccd74ab9..5679734b16f 100644
--- a/litellm/proxy/_experimental/mcp_server/server.py
+++ b/litellm/proxy/_experimental/mcp_server/server.py
@@ -5,6 +5,7 @@ LiteLLM MCP Server Routes
import asyncio
import contextlib
+import time
import traceback
import uuid
from datetime import datetime
@@ -54,6 +55,11 @@ from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
from litellm.utils import Rules, client, function_setup
+# Short-lived in-memory cache for BYOK credential existence checks.
+# Keyed by (user_id, server_id); value is (credential_exists, monotonic_timestamp).
+_byok_cred_cache: Dict[Tuple[str, str], Tuple[bool, float]] = {}
+_BYOK_CRED_CACHE_TTL = 60 # seconds
+
# Check if MCP is available
# "mcp" requires python 3.10 or higher, but several litellm users use python 3.8
# We're making this conditional import to avoid breaking users who use python 3.8.
@@ -1526,6 +1532,30 @@ if MCP_AVAILABLE:
},
)
+ # Check short-lived in-memory cache before hitting the DB on every tool call
+ cache_key = (user_id, mcp_server.server_id)
+ cached = _byok_cred_cache.get(cache_key)
+ if cached is not None:
+ credential_exists, ts = cached
+ if time.monotonic() - ts < _BYOK_CRED_CACHE_TTL:
+ 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"'
+ },
+ )
+ return
+
from litellm.proxy._experimental.mcp_server.db import has_user_credential
from litellm.proxy.proxy_server import prisma_client
@@ -1537,6 +1567,7 @@ if MCP_AVAILABLE:
user_id=user_id,
server_id=mcp_server.server_id,
)
+ _byok_cred_cache[cache_key] = (credential_exists, time.monotonic())
if not credential_exists:
raise HTTPException(
status_code=401,
diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
index c594ad3ec18..b08f6943358 100644
--- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py
@@ -82,7 +82,6 @@ if MCP_AVAILABLE:
get_all_mcp_servers_for_user,
get_mcp_server,
get_user_credential,
- has_user_credential,
store_user_credential,
update_mcp_server,
)
@@ -605,16 +604,24 @@ if MCP_AVAILABLE:
server.mcp_info = {}
server.mcp_info["is_public"] = True
- # Annotate has_user_credential for BYOK servers
+ # Annotate has_user_credential for BYOK servers (single batched query)
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
- )
+ byok_server_ids = [
+ s.server_id
+ for s in redacted_mcp_servers
+ if getattr(s, "is_byok", False)
+ ]
+ if byok_server_ids:
+ cred_rows = await _byok_prisma_client.db.litellm_mcpusercredentials.find_many(
+ where={"user_id": user_id, "server_id": {"in": byok_server_ids}}
+ )
+ cred_set = {r.server_id for r in cred_rows}
+ for server in redacted_mcp_servers:
+ if getattr(server, "is_byok", False):
+ server.has_user_credential = server.server_id in cred_set
# Virtual keys only get a sanitized discovery view.
if is_restricted_virtual_key: