Files
paskia/paskia/fastapi/oid.py
T

575 lines
18 KiB
Python

"""
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/oidc which passes OIDC params to
the /auth/ws/authenticate WebSocket.
"""
import base64
import hashlib
import logging
from datetime import UTC, datetime
from uuid import UUID
import base64url
from fastapi import Body, Depends, FastAPI, HTTPException, Request
from fastapi.responses import JSONResponse
from fastapi.security import HTTPBearer
from paskia import authcode, db
from paskia.config import SESSION_LIFETIME
from paskia.db.structs import Session
from paskia.util import oidjwt
from paskia.util.crypto import hash_secret
_logger = logging.getLogger(__name__)
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
def _oidc_session_by_token(
token: str, client_uuid: UUID | None = None
) -> Session | None:
"""Look up an OIDC session by token (refresh token value)."""
key = hash_secret("oidc", token)
s = db.data().sessions.get(key)
if not s or s.client_uuid is None:
return None
if client_uuid is not None and s.client_uuid != client_uuid:
return None
return s
def _oidc_session_by_sid(sid: bytes, client_uuid: UUID | None = None) -> Session | None:
"""Look up an OIDC session by sid (for backchannel logout)."""
for s in db.data().sessions.values():
if s.client_uuid is None:
continue
if client_uuid is not None and s.client_uuid != client_uuid:
continue
if hash_secret("oidc", s.key) == sid:
return s
return 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) -> bool:
"""Verify PKCE code_verifier against stored code_challenge (S256 only)."""
digest = hashlib.sha256(code_verifier.encode("ascii")).digest()
computed = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
return computed == code_challenge
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),
refresh_token: str | None = Body(None, embed=False),
):
"""OIDC Token endpoint.
Supports:
- grant_type=authorization_code: Exchange code for tokens
- grant_type=refresh_token: Refresh access token using sid
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)
refresh_token = form.get("refresh_token", refresh_token)
# Get client credentials (required for all grant types)
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)
if grant_type == "authorization_code":
return await _handle_authorization_code(
request, client, client_id, code, redirect_uri, code_verifier
)
elif grant_type == "refresh_token":
return await _handle_refresh_token(request, client, client_id, refresh_token)
else:
return JSONResponse(
{"error": "unsupported_grant_type"},
status_code=400,
)
async def _handle_authorization_code(
request: Request,
client,
client_id: str,
code: str | None,
redirect_uri: str | None,
code_verifier: str | None,
):
"""Handle grant_type=authorization_code."""
if not code:
return JSONResponse(
{"error": "invalid_request", "error_description": "Missing code"},
status_code=400,
)
# Consume auth code (atomic delete + return)
auth_code = authcode.codes.pop(code, None)
if not auth_code:
return JSONResponse(
{"error": "invalid_grant", "error_description": "Code expired or invalid"},
status_code=400,
)
# Look up the OIDC session by token
session = _oidc_session_by_token(auth_code.session_key, client.uuid)
if not session:
return JSONResponse(
{
"error": "invalid_grant",
"error_description": "Session not found or not OIDC",
},
status_code=400,
)
# Verify redirect_uri matches (OIDC only)
if auth_code.oidc and redirect_uri and redirect_uri != auth_code.oidc.redirect_uri:
return JSONResponse(
{"error": "invalid_grant", "error_description": "redirect_uri mismatch"},
status_code=400,
)
# Verify PKCE (OIDC only)
if auth_code.oidc:
if not code_verifier:
return JSONResponse(
{
"error": "invalid_grant",
"error_description": "Missing code_verifier",
},
status_code=400,
)
method = auth_code.oidc.code_challenge_method
if method != "S256":
return JSONResponse(
{
"error": "invalid_grant",
"error_description": "Only S256 code_challenge_method is supported",
},
status_code=400,
)
if not _verify_pkce(code_verifier, auth_code.oidc.code_challenge):
return JSONResponse(
{
"error": "invalid_grant",
"error_description": "Invalid code_verifier",
},
status_code=400,
)
# Get user from session
user = db.data().users.get(session.user_uuid)
if not user:
return JSONResponse(
{"error": "invalid_grant", "error_description": "User not found"},
status_code=400,
)
# Derive sid from session key
sid = base64url.enc(hash_secret("oidc", session.key))
return _build_token_response(
request,
user,
client_id,
auth_code.session_key,
sid,
auth_code.oidc.nonce if auth_code.oidc else None,
auth_code.oidc.scope if auth_code.oidc else None,
credential_uuid=session.credential_uuid,
)
async def _handle_refresh_token(
request: Request,
client,
client_id: str,
refresh_token_value: str | None,
):
"""Handle grant_type=refresh_token.
The refresh_token is the session secret. On refresh:
- Validates session exists and belongs to client
- Extends session expiry (24h sliding window)
- Records current IP and user_agent
- Issues new access_token and id_token
"""
if not refresh_token_value:
return JSONResponse(
{"error": "invalid_request", "error_description": "Missing refresh_token"},
status_code=400,
)
# Look up session by refresh token
session = _oidc_session_by_token(refresh_token_value, client.uuid)
if not session:
return JSONResponse(
{
"error": "invalid_grant",
"error_description": "Invalid or expired refresh_token",
},
status_code=400,
)
# Check session not expired
now = datetime.now(UTC)
if session.expiry < now:
return JSONResponse(
{"error": "invalid_grant", "error_description": "Refresh token expired"},
status_code=400,
)
# Get user
user = db.data().users.get(session.user_uuid)
if not user:
return JSONResponse(
{"error": "invalid_grant", "error_description": "User not found"},
status_code=400,
)
# Refresh the session - extend expiry and record IP/user_agent
ip = request.headers.get("x-forwarded-for", "").split(",")[0].strip()
if not ip:
ip = request.client.host if request.client else ""
user_agent = request.headers.get("user-agent", "")
db.update_session(
session.key,
ip=ip,
user_agent=user_agent,
expiry=now + SESSION_LIFETIME,
)
_logger.info("OIDC session refreshed: %s", session.key)
# Base64url encode session's derived sid for JWT claim
sid_str = base64url.enc(hash_secret("oidc", session.key))
return _build_token_response(
request,
user,
client_id,
refresh_token_value,
sid_str,
nonce=None,
scope="openid",
credential_uuid=session.credential_uuid,
)
def _build_token_response(
request: Request,
user,
client_id: str,
secret: str,
sid: str,
nonce: str | None,
scope: str,
credential_uuid: UUID | None = None,
):
"""Build the token response with access_token, id_token, and refresh_token."""
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)
# Get credential's last_used as auth_time
auth_time = None
if credential_uuid:
try:
credential = db.data().credentials[credential_uuid]
if credential.last_used:
auth_time = credential.last_used
except KeyError:
pass
# Create ID token
id_token = oidjwt.create_id_token(
issuer=issuer,
subject=user.uuid,
audience=client_id,
nonce=nonce,
sid=sid,
name=user.display_name,
preferred_username=user.preferred_username,
email=user.email,
permissions=permissions if permissions else None,
auth_time=auth_time,
)
# Create access token
access_token = oidjwt.create_access_token(
issuer=issuer,
subject=user.uuid,
audience=client_id,
scope=scope,
)
return JSONResponse(
{
"access_token": access_token,
"token_type": "Bearer",
"expires_in": 3600,
"refresh_token": secret,
"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")
# Verify audience is a valid client
aud = payload.get("aud")
if not aud:
raise HTTPException(401, "Invalid token (missing aud claim)")
try:
client_uuid = UUID(aud)
except ValueError:
raise HTTPException(401, "Invalid token (invalid aud format)")
if not db.data().oid_clients.get(client_uuid):
raise HTTPException(401, "Invalid token (unknown client)")
# 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
@app.post("/backchannel-logout")
async def backchannel_logout(
request: Request,
logout_token: str | None = Body(None, embed=False),
):
"""OIDC Back-Channel Logout endpoint.
Receives a logout_token JWT from the RP and invalidates the session.
The logout_token must contain either 'sid' (session ID) or 'sub' (user ID).
"""
# Parse form data
content_type = request.headers.get("content-type", "")
if "application/x-www-form-urlencoded" in content_type:
form = await request.form()
logout_token = form.get("logout_token", logout_token)
if not logout_token:
return JSONResponse(
{"error": "invalid_request", "error_description": "Missing logout_token"},
status_code=400,
)
# Decode and verify the logout token
issuer = _get_issuer(request)
payload = oidjwt.decode_access_token(logout_token, issuer)
if not payload:
return JSONResponse(
{"error": "invalid_request", "error_description": "Invalid logout_token"},
status_code=400,
)
# Validate required claims
sid = payload.get("sid")
sub = payload.get("sub")
# Verify audience is a valid client (if present)
aud = payload.get("aud")
client_uuid = None
if aud:
try:
client_uuid = UUID(aud)
if not db.data().oid_clients.get(client_uuid):
return JSONResponse(
{
"error": "invalid_request",
"error_description": "Unknown client in logout_token",
},
status_code=400,
)
except ValueError:
return JSONResponse(
{
"error": "invalid_request",
"error_description": "Invalid client format in logout_token",
},
status_code=400,
)
if not sid and not sub:
return JSONResponse(
{
"error": "invalid_request",
"error_description": "logout_token must contain sid or sub",
},
status_code=400,
)
# Delete session(s)
deleted = 0
if sid:
# Decode sid from base64url to bytes
try:
sid_bytes = base64url.dec(sid)
except Exception:
return JSONResponse(
{"error": "invalid_request", "error_description": "Invalid sid format"},
status_code=400,
)
# Delete specific session by sid
session = _oidc_session_by_sid(sid_bytes, client_uuid)
if session:
db.delete_session(session.key)
deleted = 1
_logger.info("Back-channel logout: deleted session %s", sid)
elif sub:
# Delete all OIDC sessions for this user/client
try:
user_uuid = UUID(sub)
except ValueError:
return JSONResponse(
{"error": "invalid_request", "error_description": "Invalid sub claim"},
status_code=400,
)
# Find and delete matching sessions
sessions_to_delete = [
s
for s in db.data().sessions.values()
if s.user_uuid == user_uuid
and s.client_uuid is not None
and (client_uuid is None or s.client_uuid == client_uuid)
]
for session in sessions_to_delete:
db.delete_session(session.key)
deleted += 1
if deleted:
_logger.info(
"Back-channel logout: deleted %d sessions for user %s", deleted, sub
)
# Return 200 OK even if no sessions were found (per spec)
return JSONResponse({"deleted": deleted})