Draft OpenID Connect support.
This commit is contained in:
+113
-2
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
from datetime import UTC, datetime
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import UUID
|
||||
|
||||
import msgspec
|
||||
@@ -197,7 +197,7 @@ class Role(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
class User(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
"""User data structure.
|
||||
|
||||
Mutable fields: display_name, role_uuid, last_seen, visits, theme
|
||||
Mutable fields: display_name, role_uuid, last_seen, visits, theme, email, preferred_username
|
||||
Immutable fields: created_at (set at creation, never modified)
|
||||
uuid is derived from created_at using uuid7.
|
||||
"""
|
||||
@@ -208,6 +208,8 @@ class User(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
last_seen: datetime | None = None
|
||||
visits: int = 0
|
||||
theme: str = "" # "" or "auto" = OS default, "light", "dark"
|
||||
email: str | None = None # OIDC email claim
|
||||
preferred_username: str | None = None # OIDC preferred_username claim
|
||||
|
||||
def __post_init__(self):
|
||||
if not hasattr(self, "uuid"):
|
||||
@@ -506,6 +508,107 @@ class ResetToken(msgspec.Struct, dict=True):
|
||||
return token, passphrase
|
||||
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# OIDC Provider structures
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
|
||||
class OIDClient(msgspec.Struct, dict=True):
|
||||
"""OIDC client (relying party) registration.
|
||||
|
||||
client_id is the dict key (UUID).
|
||||
"""
|
||||
|
||||
client_secret_hash: bytes
|
||||
name: str
|
||||
redirect_uris: list[str]
|
||||
created_at: datetime
|
||||
|
||||
def __post_init__(self):
|
||||
if not hasattr(self, "uuid"):
|
||||
self.uuid: UUID = _UUID_UNSET
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
name: str,
|
||||
redirect_uris: list[str],
|
||||
client_secret: str,
|
||||
created_at: datetime | None = None,
|
||||
) -> tuple[OIDClient, str]:
|
||||
"""Create a new OIDClient with hashed secret.
|
||||
|
||||
Returns (client, client_secret) tuple.
|
||||
"""
|
||||
now = created_at or datetime.now(UTC)
|
||||
secret_hash = hashlib.sha256(client_secret.encode()).digest()
|
||||
client = cls(
|
||||
client_secret_hash=secret_hash,
|
||||
name=name,
|
||||
redirect_uris=redirect_uris,
|
||||
created_at=now,
|
||||
)
|
||||
client.uuid = uuid7.create(now)
|
||||
return client, client_secret
|
||||
|
||||
def verify_secret(self, client_secret: str) -> bool:
|
||||
"""Verify a client secret against stored hash."""
|
||||
return secrets.compare_digest(
|
||||
self.client_secret_hash,
|
||||
hashlib.sha256(client_secret.encode()).digest(),
|
||||
)
|
||||
|
||||
|
||||
class OIDAuthCode(msgspec.Struct, dict=True):
|
||||
"""OIDC authorization code (short-lived, single-use).
|
||||
|
||||
code is the dict key (random string).
|
||||
"""
|
||||
|
||||
client_uuid: UUID = msgspec.field(name="client")
|
||||
user_uuid: UUID = msgspec.field(name="user")
|
||||
redirect_uri: str
|
||||
scope: str
|
||||
nonce: str | None = None
|
||||
code_challenge: str | None = None
|
||||
code_challenge_method: str | None = None
|
||||
created_at: datetime = msgspec.field(default_factory=lambda: datetime.now(UTC))
|
||||
expires_at: datetime = msgspec.field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
def __post_init__(self):
|
||||
if not hasattr(self, "code"):
|
||||
self.code: str = ""
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
client: UUID | OIDClient,
|
||||
user: UUID,
|
||||
redirect_uri: str,
|
||||
scope: str,
|
||||
nonce: str | None = None,
|
||||
code_challenge: str | None = None,
|
||||
code_challenge_method: str | None = None,
|
||||
lifetime_seconds: int = 600,
|
||||
) -> OIDAuthCode:
|
||||
"""Create a new auth code with 10-minute default lifetime."""
|
||||
now = datetime.now(UTC)
|
||||
client_uuid = client if isinstance(client, UUID) else client.uuid
|
||||
auth_code = cls(
|
||||
client_uuid=client_uuid,
|
||||
user_uuid=user,
|
||||
redirect_uri=redirect_uri,
|
||||
scope=scope,
|
||||
nonce=nonce,
|
||||
code_challenge=code_challenge,
|
||||
code_challenge_method=code_challenge_method,
|
||||
created_at=now,
|
||||
expires_at=now + timedelta(seconds=lifetime_seconds),
|
||||
)
|
||||
auth_code.code = secrets.token_urlsafe(32)
|
||||
return auth_code
|
||||
|
||||
|
||||
class SessionContext(msgspec.Struct):
|
||||
session: Session
|
||||
user: User
|
||||
@@ -541,6 +644,9 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
credentials: dict[UUID, Credential] = {}
|
||||
sessions: dict[str, Session] = {}
|
||||
reset_tokens: dict[bytes, ResetToken] = {}
|
||||
# OIDC provider data
|
||||
oid_clients: dict[UUID, OIDClient] = {}
|
||||
oid_auth_codes: dict[str, OIDAuthCode] = {}
|
||||
|
||||
def __post_init__(self):
|
||||
# Store reference for persistence (not serialized)
|
||||
@@ -560,6 +666,11 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
session.key = key
|
||||
for key, token in self.reset_tokens.items():
|
||||
token.key = key
|
||||
# OIDC
|
||||
for uuid, client in self.oid_clients.items():
|
||||
client.uuid = uuid
|
||||
for code, auth_code in self.oid_auth_codes.items():
|
||||
auth_code.code = code
|
||||
|
||||
def transaction(self, action, ctx=None, *, user=None):
|
||||
"""Wrap writes in transaction. Delegates to JsonlStore."""
|
||||
|
||||
Reference in New Issue
Block a user