Draft OpenID Connect support.
This commit is contained in:
@@ -0,0 +1,272 @@
|
||||
"""
|
||||
OIDC Provider endpoints.
|
||||
|
||||
Implements OpenID Connect 1.0 Authorization Code flow:
|
||||
- POST /token - Token endpoint (code exchange)
|
||||
- GET /userinfo - UserInfo endpoint (bearer token)
|
||||
|
||||
Authorization is handled by /auth/restricted/ which passes OIDC params to
|
||||
the /auth/ws/authenticate WebSocket.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import logging
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import Body, Depends, FastAPI, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.security import HTTPBearer
|
||||
|
||||
from paskia import db
|
||||
from paskia.util import oidjwt
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||
|
||||
|
||||
def _get_issuer(request: Request) -> str:
|
||||
"""Build issuer URL from request."""
|
||||
scheme = request.headers.get("x-forwarded-proto", request.url.scheme)
|
||||
host = request.headers.get("host", request.url.netloc)
|
||||
return f"{scheme}://{host}"
|
||||
|
||||
|
||||
def _verify_pkce(code_verifier: str, code_challenge: str, method: str) -> bool:
|
||||
"""Verify PKCE code_verifier against stored code_challenge."""
|
||||
if method == "plain":
|
||||
return code_verifier == code_challenge
|
||||
elif method == "S256":
|
||||
# SHA256 hash, base64url encode
|
||||
digest = hashlib.sha256(code_verifier.encode("ascii")).digest()
|
||||
computed = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
|
||||
return computed == code_challenge
|
||||
return False
|
||||
|
||||
|
||||
def _parse_client_credentials(
|
||||
request: Request,
|
||||
client_id: str | None,
|
||||
client_secret: str | None,
|
||||
) -> tuple[str, str]:
|
||||
"""Extract client credentials from request (Basic auth or body params)."""
|
||||
auth_header = request.headers.get("authorization", "")
|
||||
if auth_header.lower().startswith("basic "):
|
||||
try:
|
||||
decoded = base64.b64decode(auth_header[6:]).decode("utf-8")
|
||||
client_id, client_secret = decoded.split(":", 1)
|
||||
except Exception:
|
||||
raise HTTPException(401, "Invalid Authorization header")
|
||||
|
||||
if not client_id or not client_secret:
|
||||
raise HTTPException(401, "Missing client credentials")
|
||||
|
||||
return client_id, client_secret
|
||||
|
||||
|
||||
@app.post("/token")
|
||||
async def token(
|
||||
request: Request,
|
||||
grant_type: str = Body(..., embed=False),
|
||||
code: str | None = Body(None, embed=False),
|
||||
redirect_uri: str | None = Body(None, embed=False),
|
||||
client_id: str | None = Body(None, embed=False),
|
||||
client_secret: str | None = Body(None, embed=False),
|
||||
code_verifier: str | None = Body(None, embed=False),
|
||||
):
|
||||
"""OIDC Token endpoint.
|
||||
|
||||
Exchanges authorization code for tokens.
|
||||
Supports client_secret_post and client_secret_basic authentication.
|
||||
"""
|
||||
# Parse form data (OAuth uses application/x-www-form-urlencoded)
|
||||
content_type = request.headers.get("content-type", "")
|
||||
if "application/x-www-form-urlencoded" in content_type:
|
||||
form = await request.form()
|
||||
grant_type = form.get("grant_type", grant_type)
|
||||
code = form.get("code", code)
|
||||
redirect_uri = form.get("redirect_uri", redirect_uri)
|
||||
client_id = form.get("client_id", client_id)
|
||||
client_secret = form.get("client_secret", client_secret)
|
||||
code_verifier = form.get("code_verifier", code_verifier)
|
||||
|
||||
if grant_type != "authorization_code":
|
||||
return JSONResponse(
|
||||
{"error": "unsupported_grant_type"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
if not code:
|
||||
return JSONResponse(
|
||||
{"error": "invalid_request", "error_description": "Missing code"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
# Get client credentials
|
||||
client_id, client_secret = _parse_client_credentials(
|
||||
request, client_id, client_secret
|
||||
)
|
||||
|
||||
# Validate client
|
||||
try:
|
||||
client_uuid = UUID(client_id)
|
||||
except ValueError:
|
||||
return JSONResponse({"error": "invalid_client"}, status_code=401)
|
||||
|
||||
client = db.data().oid_clients.get(client_uuid)
|
||||
if not client or not client.verify_secret(client_secret):
|
||||
return JSONResponse({"error": "invalid_client"}, status_code=401)
|
||||
|
||||
# Consume auth code (atomic delete + return)
|
||||
auth_code = db.consume_oid_auth_code(code)
|
||||
if not auth_code:
|
||||
return JSONResponse(
|
||||
{"error": "invalid_grant", "error_description": "Code expired or invalid"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
# Verify client matches
|
||||
if auth_code.client_uuid != client.uuid:
|
||||
return JSONResponse({"error": "invalid_grant"}, status_code=400)
|
||||
|
||||
# Verify redirect_uri matches
|
||||
if redirect_uri and redirect_uri != auth_code.redirect_uri:
|
||||
return JSONResponse(
|
||||
{"error": "invalid_grant", "error_description": "redirect_uri mismatch"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
# Verify PKCE if code_challenge was provided
|
||||
if auth_code.code_challenge:
|
||||
if not code_verifier:
|
||||
return JSONResponse(
|
||||
{
|
||||
"error": "invalid_grant",
|
||||
"error_description": "Missing code_verifier",
|
||||
},
|
||||
status_code=400,
|
||||
)
|
||||
method = auth_code.code_challenge_method or "plain"
|
||||
if not _verify_pkce(code_verifier, auth_code.code_challenge, method):
|
||||
return JSONResponse(
|
||||
{
|
||||
"error": "invalid_grant",
|
||||
"error_description": "Invalid code_verifier",
|
||||
},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
# Get user
|
||||
user = db.data().users.get(auth_code.user_uuid)
|
||||
if not user:
|
||||
return JSONResponse(
|
||||
{"error": "invalid_grant", "error_description": "User not found"},
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
# Build issuer
|
||||
issuer = _get_issuer(request)
|
||||
|
||||
# Get user's permissions from role
|
||||
role = user.role
|
||||
org = role.org
|
||||
org_perm_uuids = {p.uuid for p in org.permissions}
|
||||
permissions = []
|
||||
for perm_uuid in role.permission_set:
|
||||
if perm_uuid not in org_perm_uuids:
|
||||
continue
|
||||
p = db.data().permissions.get(perm_uuid)
|
||||
if p:
|
||||
permissions.append(p.scope)
|
||||
|
||||
# Create ID token
|
||||
id_token = oidjwt.create_id_token(
|
||||
issuer=issuer,
|
||||
subject=user.uuid,
|
||||
audience=client_id,
|
||||
nonce=auth_code.nonce,
|
||||
name=user.display_name,
|
||||
preferred_username=user.preferred_username,
|
||||
email=user.email,
|
||||
permissions=permissions if permissions else None,
|
||||
)
|
||||
|
||||
# Create access token
|
||||
access_token = oidjwt.create_access_token(
|
||||
issuer=issuer,
|
||||
subject=user.uuid,
|
||||
audience=client_id,
|
||||
scope=auth_code.scope,
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
{
|
||||
"access_token": access_token,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
"id_token": id_token,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
bearer_auth = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
@app.get("/userinfo")
|
||||
async def userinfo(
|
||||
request: Request,
|
||||
credentials=Depends(bearer_auth),
|
||||
):
|
||||
"""OIDC UserInfo endpoint.
|
||||
|
||||
Returns claims about the authenticated user.
|
||||
Requires Bearer token from /token endpoint.
|
||||
"""
|
||||
if not credentials:
|
||||
raise HTTPException(401, "Bearer token required")
|
||||
|
||||
issuer = _get_issuer(request)
|
||||
payload = oidjwt.decode_access_token(credentials.credentials, issuer)
|
||||
if not payload:
|
||||
raise HTTPException(401, "Invalid or expired token")
|
||||
|
||||
# Get user
|
||||
try:
|
||||
user_uuid = UUID(payload["sub"])
|
||||
except (KeyError, ValueError):
|
||||
raise HTTPException(401, "Invalid token")
|
||||
|
||||
user = db.data().users.get(user_uuid)
|
||||
if not user:
|
||||
raise HTTPException(401, "User not found")
|
||||
|
||||
# Get user's permissions
|
||||
role = user.role
|
||||
org = role.org
|
||||
org_perm_uuids = {p.uuid for p in org.permissions}
|
||||
permissions = []
|
||||
for perm_uuid in role.permission_set:
|
||||
if perm_uuid not in org_perm_uuids:
|
||||
continue
|
||||
p = db.data().permissions.get(perm_uuid)
|
||||
if p:
|
||||
permissions.append(p.scope)
|
||||
|
||||
# Build userinfo response based on scope
|
||||
scope = payload.get("scope", "openid").split()
|
||||
response = {"sub": str(user.uuid)}
|
||||
|
||||
if "profile" in scope:
|
||||
response["name"] = user.display_name
|
||||
if user.preferred_username:
|
||||
response["preferred_username"] = user.preferred_username
|
||||
|
||||
if "email" in scope and user.email:
|
||||
response["email"] = user.email
|
||||
|
||||
# Always include permissions
|
||||
if permissions:
|
||||
response["permissions"] = permissions
|
||||
|
||||
return response
|
||||
Reference in New Issue
Block a user