mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
ef1eb973b7
commit
fff3a55228
4 changed files with 74 additions and 17 deletions
|
|
@ -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">✓</span>{item}</div>'
|
||||
f'<div class="access-item"><span class="check">✓</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? ↗</a>'
|
||||
help_link_html = f'<a class="help-link" href="{e(help_url)}" target="_blank">Where do I find my API key? ↗</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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue