fix: address greptile review feedback (greploop iteration 1)

- XSS: escape all user-supplied values in _build_authorize_html() with html.escape()
- Open redirect: validate redirect_uri scheme and URL-encode code/state in redirect
- N+1 query: batch BYOK credential lookup into single find_many() call
- Critical path DB: add 60s TTL in-memory cache to _check_byok_credential()
- Encrypt BYOK credentials at rest using encrypt_value_helper/decrypt_value_helper
This commit is contained in:
Ishaan Jaffer 2026-03-04 20:05:15 -08:00
parent ef1eb973b7
commit fff3a55228
4 changed files with 74 additions and 17 deletions

View file

@ -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'<div class="access-item"><span class="check">&#10003;</span>{item}</div>'
f'<div class="access-item"><span class="check">&#10003;</span>{e(item)}</div>'
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'<a class="help-link" href="{help_url}" target="_blank">Where do I find my API key? &#8599;</a>'
help_link_html = f'<a class="help-link" href="{e(help_url)}" target="_blank">Where do I find my API key? &#8599;</a>'
return f"""<!DOCTYPE html>
<html lang="en">
@ -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)

View file

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

View file

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

View file

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