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: