Cleaner ed25519 key generation and storage in DB. Hide the value in DB logging and cleanup to remove non-color-support.

This commit is contained in:
Leo Vasanko
2026-02-17 23:39:08 +00:00
parent 3e364e0b6f
commit 8d463e0118
13 changed files with 135 additions and 118 deletions
+2 -2
View File
@@ -68,9 +68,9 @@ from paskia.db.operations import (
)
from paskia.db.structs import (
DB,
Client,
Config,
Credential,
OIDClient,
Org,
Permission,
ResetToken,
@@ -91,7 +91,7 @@ __all__ = [
"Config",
"Credential",
"DB",
"OIDClient",
"Client",
"Org",
"Permission",
"ResetToken",
+4
View File
@@ -8,6 +8,7 @@ import uuid7
import paskia.db.operations as _ops
from paskia.db.structs import Config, Org, Permission, ResetToken, Role, User
from paskia.util.crypto import secret_key
def bootstrap(
@@ -120,4 +121,7 @@ def bootstrap(
if config is not None:
_ops._db.config = config
# Generate OIDC signing key
_ops._db.oidc.key = secret_key()
return reset_passphrase
+15 -48
View File
@@ -49,11 +49,6 @@ def _is_uuid(value: str) -> bool:
return bool(_UUID_PATTERN.match(value))
def _uuid_suffix(uuid_str: str) -> str:
"""Get the last section of a UUID (after the last hyphen)."""
return uuid_str.rsplit("-", 1)[-1]
class UuidResolver:
"""Resolve UUIDs to display names or short suffixes.
@@ -148,14 +143,8 @@ class UuidResolver:
return None
def _use_color() -> bool:
"""Check if we should use color output."""
return sys.stderr.isatty()
def _format_value(
value: Any,
use_color: bool,
max_len: int = 60,
resolver: UuidResolver | None = None,
) -> str:
@@ -195,16 +184,14 @@ def _format_value(
if all_true:
parts.append(key_display)
else:
val_display = _format_value(v, use_color, max_len=30, resolver=resolver)
val_display = _format_value(v, max_len=30, resolver=resolver)
parts.append(f"{key_display}: {val_display}")
return "{" + ", ".join(parts) + "}"
if isinstance(value, list):
if not value:
return "[]"
parts = [
_format_value(v, use_color, max_len=30, resolver=resolver) for v in value
]
parts = [_format_value(v, max_len=30, resolver=resolver) for v in value]
return "[" + ", ".join(parts) + "]"
# Fallback for other types
@@ -214,9 +201,7 @@ def _format_value(
return text
def _format_path(
path: list[str], use_color: bool, resolver: UuidResolver | None = None
) -> str:
def _format_path(path: list[str], resolver: UuidResolver | None = None) -> str:
"""Format a path as dot notation with prefix in dark grey, final in default.
If resolver is provided, UUIDs in the path are replaced with display names.
@@ -228,8 +213,6 @@ def _format_path(
if resolver:
path = [resolver.resolve(p) if _is_uuid(p) else p for p in path]
if not use_color:
return ".".join(path)
if len(path) == 1:
return f"{_PATH_FINAL}{path[0]}{_RESET}"
prefix = ".".join(path[:-1])
@@ -324,7 +307,6 @@ def _format_change_lines(
change_type: str,
path: list[str],
value: Any,
use_color: bool,
resolver: UuidResolver | None = None,
) -> list[str]:
"""Format a single change as one or more lines.
@@ -332,6 +314,12 @@ def _format_change_lines(
If resolver is provided, UUIDs are replaced with display names.
"""
# Helper to format a value, checking for censored paths
def fmt_value(v: Any, child_path: list[str]) -> str:
if child_path[-2:] == ["oidc", "key"]:
return f"{_DIM}<hidden>{_RESET}"
return _format_value(v, resolver=resolver)
# Helper to format path with UUID replacement
def fmt_path(p: list[str]) -> list[str]:
if resolver:
@@ -341,8 +329,6 @@ def _format_change_lines(
formatted_path = fmt_path(path)
if change_type == "delete":
if not use_color:
return [f" {'.'.join(formatted_path)}"]
if len(formatted_path) == 1:
return [f" {_DELETE}{formatted_path[0]}{_RESET}"]
prefix = ".".join(formatted_path[:-1])
@@ -355,9 +341,7 @@ def _format_change_lines(
if isinstance(value, dict) and value:
lines = []
# First line: path with green final element and grey =
if not use_color:
lines.append(f" {'.'.join(formatted_path)} =")
elif len(formatted_path) == 1:
if len(formatted_path) == 1:
lines.append(f" {_ADD}{formatted_path[0]}{_RESET} {_DIM}={_RESET}")
else:
prefix = ".".join(formatted_path[:-1])
@@ -370,21 +354,16 @@ def _format_change_lines(
formatted_items = []
for k, v in value.items():
k_display = resolver.resolve(k) if resolver and _is_uuid(k) else k
v_str = _format_value(v, use_color, resolver=resolver)
v_str = fmt_value(v, path + [k])
formatted_items.append((k_display, v_str))
max_key_len = max(len(k) for k, _ in formatted_items)
field_width = max(max_key_len, 12) # minimum 12 chars
for k_display, v_str in formatted_items:
padding = " " * (field_width - len(k_display))
if use_color:
lines.append(f" {k_display}{_DIM}:{_RESET}{padding} {v_str}")
else:
lines.append(f" {k_display}:{padding} {v_str}")
return lines
else:
value_str = _format_value(value, use_color, resolver=resolver)
if not use_color:
return [f" {'.'.join(formatted_path)} = {value_str}"]
value_str = fmt_value(value, path)
if len(formatted_path) == 1:
return [
f" {_ADD}{formatted_path[0]}{_RESET} {_DIM}={_RESET} {value_str}"
@@ -396,11 +375,9 @@ def _format_change_lines(
]
# update: Existing item being updated - normal path colors
value_str = _format_value(value, use_color, resolver=resolver)
path_str = _format_path(path, use_color, resolver=resolver)
if use_color:
value_str = fmt_value(value, path)
path_str = _format_path(path, resolver=resolver)
return [f" {path_str} {_DIM}={_RESET} {value_str}"]
return [f" {path_str} = {value_str}"]
def format_diff(
@@ -417,7 +394,6 @@ def format_diff(
Returns a list of formatted lines (without newlines).
UUIDs are replaced with display names (using previous state for lookups).
"""
use_color = _use_color()
changes: list[tuple[str, list[str], Any]] = []
_collect_changes(diff, [], changes, previous)
@@ -430,27 +406,18 @@ def format_diff(
# Format each change
lines = []
for change_type, path, value in changes:
lines.extend(
_format_change_lines(change_type, path, value, use_color, resolver)
)
lines.extend(_format_change_lines(change_type, path, value, resolver))
return lines
def format_action_header(action: str, user_display: str | None = None) -> str:
"""Format the action header line."""
use_color = _use_color()
if use_color:
action_str = f"{_ACTION}{action}{_RESET}"
if user_display:
user_str = f"{_USER}{user_display}{_RESET}"
return f"{action_str} by {user_str}"
return action_str
else:
if user_display:
return f"{action} by {user_display}"
return action
def log_change(
+14 -2
View File
@@ -8,7 +8,7 @@ Each migration should be idempotent and only run when needed.
import base64
from collections.abc import Awaitable, Callable
from paskia.util.crypto import hash_secret
from paskia.util.crypto import hash_secret, secret_key
def migrate_v1(d: dict, **kwargs) -> None:
@@ -31,7 +31,19 @@ def migrate_v3(d: dict, **kwargs) -> None:
def migrate_v4(d: dict, **kwargs) -> None:
"""OpenID Connect support and hardened session keys."""
d["oid_clients"] = {}
# Migrate existing oid_clients and oidc_key if present (from old format)
existing_clients = d.pop("oid_clients", {})
existing_key = d.pop("oidc_key", None)
# Create OIDC structure
d["oidc"] = {"clients": existing_clients, "key": existing_key}
# Generate OIDC signing key if not present
if d["oidc"]["key"] is None:
key_bytes = secret_key()
d["oidc"]["key"] = base64.standard_b64encode(key_bytes).decode()
elif isinstance(d["oidc"]["key"], bytes):
# Existing key is bytes, encode to base64 string
d["oidc"]["key"] = base64.standard_b64encode(d["oidc"]["key"]).decode()
# Migrate sessions
d["sessions"] = {
base64.standard_b64encode(hash_secret("cookie", k)).decode(): v
for k, v in d["sessions"].items()
+14 -14
View File
@@ -19,9 +19,9 @@ from paskia.db.jsonl import (
)
from paskia.db.structs import (
DB,
Client,
Config,
Credential,
OIDClient,
Org,
Permission,
ResetToken,
@@ -697,12 +697,12 @@ def create_credential_session(
# -------------------------------------------------------------------------
def create_oid_client(client: OIDClient, *, ctx: SessionContext | None = None) -> None:
def create_oid_client(client: Client, *, ctx: SessionContext | None = None) -> None:
"""Create a new OIDC client."""
if client.uuid in _db.oid_clients:
if client.uuid in _db.oidc.clients:
raise ValueError(f"OIDC client {client.uuid} already exists")
with _db.transaction("admin:create_oid_client", ctx):
_db.oid_clients[client.uuid] = client
_db.oidc.clients[client.uuid] = client
def update_oid_client(
@@ -715,10 +715,10 @@ def update_oid_client(
ctx: SessionContext | None = None,
) -> None:
"""Update an OIDC client's name, redirect URIs, and/or secret."""
if client_uuid not in _db.oid_clients:
if client_uuid not in _db.oidc.clients:
raise ValueError(f"OIDC client {client_uuid} not found")
client = _db.oid_clients[client_uuid]
client = _db.oidc.clients[client_uuid]
changes = {}
if name is not None and name != client.name:
@@ -744,7 +744,7 @@ def update_oid_client(
with _db.transaction("admin:update_oid_client", ctx):
# Create updated client with new values
updated_client = OIDClient(
updated_client = Client(
client_secret_hash=secret_hash
if secret_hash is not None
else client.client_secret_hash,
@@ -755,7 +755,7 @@ def update_oid_client(
backchannel_logout_uri=new_logout_uri,
)
updated_client.uuid = client.uuid
_db.oid_clients[client_uuid] = updated_client
_db.oidc.clients[client_uuid] = updated_client
def reset_oid_client_secret(
@@ -765,23 +765,23 @@ def reset_oid_client_secret(
ctx: SessionContext | None = None,
) -> None:
"""Reset an OIDC client's secret."""
if client_uuid not in _db.oid_clients:
if client_uuid not in _db.oidc.clients:
raise ValueError(f"OIDC client {client_uuid} not found")
client = _db.oid_clients[client_uuid]
client = _db.oidc.clients[client_uuid]
with _db.transaction("admin:reset_oid_client_secret", ctx):
updated = OIDClient(
updated = Client(
client_secret_hash=new_secret_hash,
name=client.name,
redirect_uris=client.redirect_uris,
backchannel_logout_uri=client.backchannel_logout_uri,
)
updated.uuid = client.uuid
_db.oid_clients[client_uuid] = updated
_db.oidc.clients[client_uuid] = updated
def delete_oid_client(client_uuid: UUID, *, ctx: SessionContext | None = None) -> None:
"""Delete an OIDC client."""
if client_uuid not in _db.oid_clients:
if client_uuid not in _db.oidc.clients:
raise ValueError(f"OIDC client {client_uuid} not found")
with _db.transaction("admin:delete_oid_client", ctx):
del _db.oid_clients[client_uuid]
del _db.oidc.clients[client_uuid]
+9 -4
View File
@@ -540,7 +540,7 @@ class ResetToken(msgspec.Struct, dict=True):
# -------------------------------------------------------------------------
class OIDClient(msgspec.Struct, dict=True, omit_defaults=True):
class Client(msgspec.Struct, dict=True, omit_defaults=True):
"""OIDC client (relying party) registration.
client_id is the dict key (UUID).
@@ -563,7 +563,7 @@ class OIDClient(msgspec.Struct, dict=True, omit_defaults=True):
client_secret: str,
created_at: datetime | None = None,
backchannel_logout_uri: str | None = None,
) -> tuple[OIDClient, str]:
) -> tuple[Client, str]:
"""Create a new OIDClient with hashed secret.
Returns (client, client_secret) tuple.
@@ -596,6 +596,11 @@ class SessionContext(msgspec.Struct):
permissions: list[Permission] = []
class OIDC(msgspec.Struct, dict=True):
clients: dict[UUID, Client] = {}
key: bytes | None = None
class Config(msgspec.Struct, frozen=True, dict=True, omit_defaults=True):
"""Stored configuration for the instance."""
@@ -623,7 +628,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
sessions: dict[bytes, Session] = {}
reset_tokens: dict[bytes, ResetToken] = {}
# OIDC provider data
oid_clients: dict[UUID, OIDClient] = {}
oidc: OIDC = msgspec.field(default_factory=lambda: OIDC())
def __post_init__(self):
# Store reference for persistence (not serialized)
@@ -644,7 +649,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
for key, token in self.reset_tokens.items():
token.key = key
# OIDC
for uuid, client in self.oid_clients.items():
for uuid, client in self.oidc.clients.items():
client.uuid = uuid
def transaction(self, action, ctx=None, *, user=None):
+4 -4
View File
@@ -13,7 +13,7 @@ from paskia.db import Permission as PermDC
from paskia.db import Role as RoleDC
from paskia.db import User as UserDC
from paskia.db.operations import _UNSET
from paskia.db.structs import OIDClient
from paskia.db.structs import Client
from paskia.fastapi import authz
from paskia.fastapi.response import MsgspecResponse
from paskia.fastapi.session import AUTH_COOKIE
@@ -743,7 +743,7 @@ def _validate_permission_domain(domain: str | None) -> None:
# Allow OIDC client UUIDs (used for groups claim)
try:
client_uuid = UUID(domain)
if client_uuid in db.data().oid_clients:
if client_uuid in db.data().oidc.clients:
return
except ValueError:
pass
@@ -957,7 +957,7 @@ async def admin_list_oidc_clients(request: Request, auth=AUTH_COOKIE):
mode="forbidden",
)
clients = sorted(db.data().oid_clients.values(), key=lambda c: c.uuid)
clients = sorted(db.data().oidc.clients.values(), key=lambda c: c.uuid)
sessions = db.data().sessions
# Count active sessions per client
client_session_counts = {}
@@ -1036,7 +1036,7 @@ async def admin_create_oidc_client(
if backchannel_logout_uri and not backchannel_logout_uri.startswith("http"):
raise ValueError("backchannel_logout_uri must be an HTTP(S) URL")
client = OIDClient(
client = Client(
client_secret_hash=secret_hash,
name=name,
redirect_uris=redirect_uris,
+3 -3
View File
@@ -149,7 +149,7 @@ async def token(
except ValueError:
return JSONResponse({"error": "invalid_client"}, status_code=401)
client = db.data().oid_clients.get(client_uuid)
client = db.data().oidc.clients.get(client_uuid)
if not client or not client.verify_secret(client_secret):
return JSONResponse({"error": "invalid_client"}, status_code=401)
@@ -425,7 +425,7 @@ async def userinfo(
except ValueError:
raise HTTPException(401, "Invalid token (invalid aud format)")
if not db.data().oid_clients.get(client_uuid):
if not db.data().oidc.clients.get(client_uuid):
raise HTTPException(401, "Invalid token (unknown client)")
# Get user
@@ -510,7 +510,7 @@ async def backchannel_logout(
if aud:
try:
client_uuid = UUID(aud)
if not db.data().oid_clients.get(client_uuid):
if not db.data().oidc.clients.get(client_uuid):
return JSONResponse(
{
"error": "invalid_request",
+2 -2
View File
@@ -134,7 +134,7 @@ async def websocket_authenticate(
await ws.send_json({"status": 400, "detail": "Invalid client_id"})
return
oidc_client = db.data().oid_clients.get(client_uuid)
oidc_client = db.data().oidc.clients.get(client_uuid)
if not oidc_client:
await ws.send_json({"status": 400, "detail": "Unknown client_id"})
return
@@ -148,7 +148,7 @@ async def websocket_authenticate(
# Store as the only allowed redirect URI
db.update_oid_client(client_uuid, redirect_uris=[redirect_uri])
# Reload client to get updated redirect_uris
oidc_client = db.data().oid_clients.get(client_uuid)
oidc_client = db.data().oidc.clients.get(client_uuid)
elif redirect_uri not in oidc_client.redirect_uris:
await ws.send_json({"status": 400, "detail": "Invalid redirect_uri"})
return
+1 -1
View File
@@ -43,7 +43,7 @@ def _collect_oidc_sessions(
session = data.sessions.get(key)
if not session or session.client_uuid is None:
continue
client = data.oid_clients.get(session.client_uuid)
client = data.oidc.clients.get(session.client_uuid)
if not client or not client.backchannel_logout_uri:
continue
sid = base64url.enc(hash_secret("oidc", session.key))
+1 -1
View File
@@ -99,7 +99,7 @@ class ApiUserSession(msgspec.Struct, omit_defaults=True):
) -> ApiUserSession:
client_name = None
if s.client_uuid:
c = db.data().oid_clients.get(s.client_uuid)
c = db.data().oidc.clients.get(s.client_uuid)
client_name = c.name if c else str(s.client_uuid)
return cls(
credential_uuid=s.credential_uuid,
+41
View File
@@ -1,5 +1,8 @@
import hashlib
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
def hash_secret(*data) -> bytes:
"""A custom HMAC that securily combines and hashes the given data (context, secrets). The first argument should be a namespacing string."""
@@ -9,3 +12,41 @@ def hash_secret(*data) -> bytes:
d = d.encode()
inner += hashlib.sha256(d).digest()
return hashlib.sha256(inner).digest()[:12]
def secret_key() -> bytes:
"""Generate a new Ed25519 private key and return as 32 raw bytes."""
private_key = Ed25519PrivateKey.generate()
return private_key.private_bytes(
encoding=serialization.Encoding.Raw,
format=serialization.PrivateFormat.Raw,
encryption_algorithm=serialization.NoEncryption(),
)
def public_key_from_secret(secret_key_bytes: bytes) -> Ed25519PrivateKey:
"""Load Ed25519 private key from 32 raw bytes."""
return Ed25519PrivateKey.from_private_bytes(secret_key_bytes)
def get_public_key_der(private_key: Ed25519PrivateKey) -> bytes:
"""Get DER-encoded public key for kid generation."""
public_key = private_key.public_key()
return public_key.public_bytes(
encoding=serialization.Encoding.DER,
format=serialization.PublicFormat.SubjectPublicKeyInfo,
)
def generate_kid(public_key_der: bytes) -> str:
"""Generate key ID from public key DER bytes."""
return hashlib.sha256(public_key_der).hexdigest()[:16]
def get_public_key_raw(private_key: Ed25519PrivateKey) -> bytes:
"""Get raw 32-byte public key for JWKS."""
public_key = private_key.public_key()
return public_key.public_bytes(
encoding=serialization.Encoding.Raw,
format=serialization.PublicFormat.Raw,
)
+18 -30
View File
@@ -3,53 +3,44 @@ OIDC JWT utilities for signing ID tokens and serving JWKS.
"""
import hashlib
import logging
from base64 import urlsafe_b64encode
from datetime import UTC, datetime, timedelta
from pathlib import Path
from uuid import UUID
import jwt
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
_logger = logging.getLogger(__name__)
from paskia import db
from paskia.util.crypto import (
generate_kid,
get_public_key_der,
get_public_key_raw,
public_key_from_secret,
secret_key,
)
# JWT signing key (loaded on first use)
_private_key = None
_public_key = None
_kid: str | None = None
# Key file location (same directory as database)
_KEY_FILE = Path("oidc_key.pem")
def _load_or_generate_key() -> None:
"""Load existing Ed25519 key or generate a new one."""
global _private_key, _public_key, _kid
if _KEY_FILE.exists():
_logger.info("Loading OIDC signing key from %s", _KEY_FILE)
pem_data = _KEY_FILE.read_bytes()
_private_key = serialization.load_pem_private_key(pem_data, password=None)
data = db.data()
if data.oidc.key is not None:
_private_key = public_key_from_secret(data.oidc.key)
else:
_logger.info("Generating new OIDC signing key")
_private_key = Ed25519PrivateKey.generate()
pem_data = _private_key.private_bytes(
encoding=serialization.Encoding.PEM,
format=serialization.PrivateFormat.PKCS8,
encryption_algorithm=serialization.NoEncryption(),
)
_KEY_FILE.write_bytes(pem_data)
_logger.info("Saved OIDC signing key to %s", _KEY_FILE)
raw_key = secret_key()
with data.transaction("oidc_key"):
data.oidc.key = raw_key
_private_key = public_key_from_secret(raw_key)
_public_key = _private_key.public_key()
# Generate kid from public key fingerprint
pub_der = _public_key.public_bytes(
encoding=serialization.Encoding.DER,
format=serialization.PublicFormat.SubjectPublicKeyInfo,
)
_kid = hashlib.sha256(pub_der).hexdigest()[:16]
pub_der = get_public_key_der(_private_key)
_kid = generate_kid(pub_der)
def _ensure_key() -> None:
@@ -63,10 +54,7 @@ def get_jwks() -> dict:
_ensure_key()
assert _public_key is not None
# Ed25519 public key is 32 bytes raw
pub_bytes = _public_key.public_bytes(
encoding=serialization.Encoding.Raw,
format=serialization.PublicFormat.Raw,
)
pub_bytes = get_public_key_raw(_private_key)
return {
"keys": [
{