mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(discoverable_endpoints.py): use encryption + encoding to securely handle state + redirect uri without storing in db
would needlessly flood the db
This commit is contained in:
parent
0d36e6cbe9
commit
1c939c70e5
1 changed files with 60 additions and 14 deletions
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
from typing import Optional
|
||||
from urllib.parse import urlencode, urlparse, urlunparse
|
||||
|
||||
|
|
@ -5,13 +6,54 @@ import httpx
|
|||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
|
||||
|
||||
STATE_MAP = {}
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
router = APIRouter(
|
||||
tags=["mcp"],
|
||||
)
|
||||
|
||||
|
||||
def encode_state_with_base_url(base_url: str, original_state: str) -> str:
|
||||
"""
|
||||
Encode the base_url and original state using encryption.
|
||||
|
||||
Args:
|
||||
base_url: The base URL to encode
|
||||
original_state: The original state parameter
|
||||
|
||||
Returns:
|
||||
An encrypted string that encodes both values
|
||||
"""
|
||||
state_data = {"base_url": base_url, "original_state": original_state}
|
||||
state_json = json.dumps(state_data, sort_keys=True)
|
||||
encrypted_state = encrypt_value_helper(state_json)
|
||||
return encrypted_state
|
||||
|
||||
|
||||
def decode_state_hash(encrypted_state: str) -> tuple[str, str]:
|
||||
"""
|
||||
Decode an encrypted state to retrieve the base_url and original state.
|
||||
|
||||
Args:
|
||||
encrypted_state: The encrypted string to decode
|
||||
|
||||
Returns:
|
||||
A tuple of (base_url, original_state)
|
||||
|
||||
Raises:
|
||||
Exception: If decryption fails or data is malformed
|
||||
"""
|
||||
decrypted_json = decrypt_value_helper(encrypted_state, "oauth_state")
|
||||
if decrypted_json is None:
|
||||
raise ValueError("Failed to decrypt state parameter")
|
||||
|
||||
state_data = json.loads(decrypted_json)
|
||||
return state_data["base_url"], state_data["original_state"]
|
||||
|
||||
|
||||
@router.get("/{mcp_server_name}/authorize")
|
||||
@router.get("/authorize")
|
||||
async def authorize(
|
||||
|
|
@ -44,12 +86,15 @@ async def authorize(
|
|||
parsed = urlparse(redirect_uri)
|
||||
base_url = urlunparse(parsed._replace(query=""))
|
||||
request_base_url = str(request.base_url).rstrip("/")
|
||||
STATE_MAP[state] = base_url
|
||||
|
||||
# Encode the base_url and original state in a unique hash
|
||||
encoded_state = encode_state_with_base_url(base_url, state)
|
||||
|
||||
params = {
|
||||
"client_id": mcp_server.client_id,
|
||||
"redirect_uri": f"{request_base_url}/callback",
|
||||
"scope": " ".join(mcp_server.scopes),
|
||||
"state": state,
|
||||
"state": encoded_state,
|
||||
}
|
||||
return RedirectResponse(f"{mcp_server.authorization_url}?{urlencode(params)}")
|
||||
|
||||
|
|
@ -114,21 +159,22 @@ async def token_endpoint(
|
|||
|
||||
@router.get("/callback")
|
||||
async def callback(code: str, state: str):
|
||||
# Exchange code for token with GitHub
|
||||
params = {"code": code, "state": state}
|
||||
try:
|
||||
# Decode the state hash to get base_url and original state
|
||||
base_url, original_state = decode_state_hash(state)
|
||||
|
||||
# Forward token to Claude ephemeral endpoint
|
||||
redirect_uri = STATE_MAP.pop(state, None)
|
||||
|
||||
if redirect_uri:
|
||||
complete_returned_url = f"{redirect_uri}?{urlencode(params)}"
|
||||
# Exchange code for token with GitHub
|
||||
params = {"code": code, "state": original_state}
|
||||
|
||||
# Forward token to Claude ephemeral endpoint
|
||||
complete_returned_url = f"{base_url}?{urlencode(params)}"
|
||||
return RedirectResponse(url=complete_returned_url, status_code=302)
|
||||
|
||||
# fallback if redirect_uri not found
|
||||
return HTMLResponse(
|
||||
"<html><body>Authentication incomplete. You can close this window.</body></html>"
|
||||
)
|
||||
except Exception:
|
||||
# fallback if state hash not found
|
||||
return HTMLResponse(
|
||||
"<html><body>Authentication incomplete. You can close this window.</body></html>"
|
||||
)
|
||||
|
||||
|
||||
# ------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue