mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(mcp): annotate has_user_credential for OAuth2 servers + accept session cookie on authorize endpoint
This commit is contained in:
parent
3ef16098f2
commit
0e66c9b4b2
1 changed files with 58 additions and 6 deletions
|
|
@ -905,20 +905,26 @@ if MCP_AVAILABLE:
|
|||
|
||||
user_id = user_api_key_dict.user_id or ""
|
||||
if user_id and _byok_prisma_client is not None:
|
||||
byok_server_ids = [
|
||||
credentialed_server_ids = [
|
||||
s.server_id
|
||||
for s in redacted_mcp_servers
|
||||
if getattr(s, "is_byok", False)
|
||||
or getattr(s, "auth_type", None) == "oauth2"
|
||||
]
|
||||
if byok_server_ids:
|
||||
if credentialed_server_ids:
|
||||
cred_rows = (
|
||||
await _byok_prisma_client.db.litellm_mcpusercredentials.find_many(
|
||||
where={"user_id": user_id, "server_id": {"in": byok_server_ids}}
|
||||
where={
|
||||
"user_id": user_id,
|
||||
"server_id": {"in": credentialed_server_ids},
|
||||
}
|
||||
)
|
||||
)
|
||||
cred_set = {r.server_id for r in cred_rows}
|
||||
for server in redacted_mcp_servers:
|
||||
if getattr(server, "is_byok", False):
|
||||
if getattr(server, "is_byok", False) or getattr(
|
||||
server, "auth_type", None
|
||||
) == "oauth2":
|
||||
server.has_user_credential = server.server_id in cred_set
|
||||
|
||||
# Virtual keys only get a sanitized discovery view.
|
||||
|
|
@ -1494,15 +1500,61 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return server
|
||||
|
||||
async def _authorize_user_auth(request: Request) -> UserAPIKeyAuth:
|
||||
"""Accept Bearer header OR session cookie for browser-redirect authorize endpoint."""
|
||||
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import (
|
||||
_user_id_from_session_cookie,
|
||||
)
|
||||
from litellm.proxy.proxy_server import master_key as _master_key
|
||||
|
||||
auth_header = request.headers.get("Authorization") or request.headers.get(
|
||||
"x-litellm-api-key"
|
||||
)
|
||||
if auth_header:
|
||||
return await user_api_key_auth(request=request)
|
||||
|
||||
# Browser redirect: verify session cookie and build a UserAPIKeyAuth from it
|
||||
cookie_token = request.cookies.get("token")
|
||||
if not cookie_token or not _master_key:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail="login_required"
|
||||
)
|
||||
try:
|
||||
import jwt as pyjwt
|
||||
|
||||
payload = pyjwt.decode(
|
||||
cookie_token, _master_key, algorithms=["HS256"]
|
||||
)
|
||||
except Exception:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail="login_required"
|
||||
)
|
||||
|
||||
user_id = payload.get("user_id")
|
||||
user_role_str = payload.get("user_role", "")
|
||||
api_key = payload.get("key", "cookie_session")
|
||||
|
||||
role_map = {
|
||||
"proxy_admin": LitellmUserRoles.PROXY_ADMIN,
|
||||
"proxy_admin_viewer": LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
"internal_user": LitellmUserRoles.INTERNAL_USER,
|
||||
"internal_user_viewer": LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
}
|
||||
user_role = role_map.get(user_role_str, LitellmUserRoles.INTERNAL_USER)
|
||||
return UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
user_id=user_id,
|
||||
user_role=user_role,
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/server/oauth/{server_id}/authorize",
|
||||
include_in_schema=False,
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def mcp_authorize(
|
||||
request: Request,
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(_authorize_user_auth),
|
||||
client_id: Optional[str] = None,
|
||||
redirect_uri: str = Query(...),
|
||||
state: str = "",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue