From 9b282503912f818e1c323dd2172869d3e56ef573 Mon Sep 17 00:00:00 2001 From: Leo Vasanko Date: Sat, 13 Jun 2026 21:59:18 +0000 Subject: [PATCH] Upgrade to kanta 0.4.0: - Make use of its new features and cleanup our interfacing and init/shutdown processes and migrations - Clean up circular deps, simplify app init - Add specific pytest for CLI main to cover the changes --- paskia/__main__.py | 45 +++++++- paskia/bootstrap.py | 59 +++------- paskia/db/__init__.py | 16 --- paskia/db/background.py | 31 +----- paskia/db/bootstrap.py | 155 ++++++++++++++++---------- paskia/db/jsonl.py | 52 --------- paskia/db/lifecycle.py | 140 ++++++++++++++++++----- paskia/db/migrations.py | 21 ++-- paskia/db/operations.py | 90 +++++++++------ paskia/db/structs.py | 23 +--- paskia/fastapi/__init__.py | 3 - paskia/fastapi/admin/permissions.py | 2 +- paskia/fastapi/admin/server_config.py | 4 +- paskia/fastapi/api.py | 2 +- paskia/fastapi/mainapp.py | 60 +++++----- paskia/fastapi/ws.py | 2 +- paskia/fastapi/wschat.py | 18 ++- paskia/fastapi/wsutil.py | 2 +- paskia/globals.py | 83 +++----------- paskia/oidc_notify.py | 6 +- paskia/util/constants.py | 6 + paskia/util/hostutil.py | 7 +- paskia/util/oidjwt.py | 5 +- paskia/util/runtime.py | 14 ++- paskia/util/startupbox.py | 6 +- pyproject.toml | 2 +- tests/conftest.py | 59 +++++++--- tests/test_api.py | 5 +- tests/test_cli.py | 154 +++++++++++++++++++++++++ 29 files changed, 620 insertions(+), 452 deletions(-) delete mode 100644 paskia/db/jsonl.py create mode 100644 paskia/util/constants.py create mode 100644 tests/test_cli.py diff --git a/paskia/__main__.py b/paskia/__main__.py index cfcf827..a66b1b0 100644 --- a/paskia/__main__.py +++ b/paskia/__main__.py @@ -1,16 +1,20 @@ import argparse +import asyncio import logging import os import sys +from pathlib import Path import msgspec from fastapi_vue import server from fastapi_vue.hostutil import parse_endpoints +from kanta import Kanta from paskia._version import __version__ -from paskia.db.jsonl import load_readonly from paskia.db.paths import db_file_path +from paskia.db.structs import DB, Config from paskia.util import startupbox +from paskia.util.constants import DEFAULT_PORT, DEVMODE from paskia.util.hostutil import ( normalize_auth_host_and_origins, normalize_origin, @@ -18,9 +22,6 @@ from paskia.util.hostutil import ( ) from paskia.util.runtime import RuntimeConfig -DEFAULT_PORT = 4401 -DEVMODE = os.getenv("PASKIA_DEV") == "1" - EPILOG = """\ Example: paskia --rp-id example.com --rp-name "Example Corporation" --auth-host auth.example.com @@ -50,6 +51,36 @@ def add_common_options(p: argparse.ArgumentParser) -> None: ) +def _load_stored_config(db_path: Path, *, rp_id: str) -> Config: + """Load the stored Config from disk using Kanta in read-only mode. + + This must not depend on PASKIA_CONFIG or the global lifecycle Kanta. + If the database file does not exist, a default config is returned. + """ + if not db_path.exists(): + return Config(rp_id=rp_id) + + kanta = Kanta( + str(db_path), + DB(config=Config(rp_id=rp_id)), + migrations="paskia.db.migrations", + ) + kanta.ctx.rp_id = rp_id + + async def _read() -> Config: + await kanta.open(readonly=True) + try: + return kanta.data.config + finally: + await kanta.close() + + try: + return asyncio.run(_read()) + except Exception as e: + logging.exception("Failed to load database") + raise SystemExit(f"{e}") from e + + def main(): # Configure logging to remove the "ERROR:root:" prefix logging.basicConfig(level=logging.INFO, format="%(message)s", force=True) @@ -75,10 +106,12 @@ def main(): args = parser.parse_args() - # Load stored config (read-only, no writes, no global state) + # Load stored config using a local read-only Kanta instance. + # This happens before PASKIA_CONFIG is set, so we must not import + # modules that initialize the global database lifecycle. db_path = db_file_path(rp_id=args.rp_id, create_root=True) try: - config = load_readonly(str(db_path), rp_id=args.rp_id).config + config = _load_stored_config(db_path, rp_id=args.rp_id) except SystemExit as e: print(f"🛑 Paskia {__version__} could not load") sys.exit(str(e)) diff --git a/paskia/bootstrap.py b/paskia/bootstrap.py index 6eb9590..9c05234 100644 --- a/paskia/bootstrap.py +++ b/paskia/bootstrap.py @@ -4,13 +4,17 @@ Bootstrap module for passkey authentication system. This module handles initial system setup when a new database is created, including creating default admin user, organization, permissions, and generating a reset link for initial admin setup. + +The actual database seeding is performed by the module-level kanta bootstrap +callback defined in :mod:`paskia.db.bootstrap` and registered during +:func:`paskia.db.lifecycle.init`. """ import logging from paskia import authsession, db +from paskia.db.bootstrap import log_reset_link from paskia.db.structs import Config -from paskia.util import hostutil logger = logging.getLogger(__name__) @@ -27,37 +31,10 @@ def _configure_logger() -> None: _configure_logger() -# Shared log message template for admin reset links -ADMIN_RESET_MESSAGE = """ -👤 Admin %s - - Use this link to register a Passkey for the admin user! -""" - def _log_reset_link(passphrase: str, message: str | None = None) -> str: """Log a reset link message and return the URL.""" - reset_link = hostutil.reset_link_url(passphrase) - if message: - logger.info(message) - logger.info(ADMIN_RESET_MESSAGE, reset_link) - return reset_link - - -async def bootstrap_system(config: Config | None = None) -> None: - """ - Bootstrap the entire system with default data. - - Uses db.bootstrap() which performs all operations in a single transaction. - The transaction log will show a single "bootstrap" action with all changes. - - Args: - config: Configuration to store (rp_id, rp_name, origins, etc.) - """ - # Call the single-transaction bootstrap function - reset_passphrase = db.bootstrap(config=config) - - # Log the reset link (this is separate from the transaction log) - _log_reset_link(reset_passphrase, "✅ Bootstrap completed!") + return log_reset_link(passphrase, message) async def check_admin_credentials() -> bool: @@ -114,22 +91,18 @@ async def check_admin_credentials() -> bool: async def bootstrap_if_needed(config: Config | None = None) -> bool: """ - Check if system needs bootstrapping and perform it if necessary. + Check if admin needs credentials and create a reset link if needed. + + Database bootstrapping itself is now handled automatically during + ``db.init()`` via the registered kanta bootstrap callback. This function + remains as a post-init hook for credential checks. Args: - config: Configuration to store during bootstrap (rp_id, rp_name, origins, etc.) + config: Kept for backwards compatibility; config is now applied during + ``db.init()``. Returns: - bool: True if bootstrapping was performed, False if system was already set up + bool: Always returns False (bootstrapping is performed during init). """ - # Check if the admin permission exists - if it does, system is already bootstrapped - if any(p.scope == "auth:admin" for p in db.data().permissions.values()): - # Permission exists, system is already bootstrapped - # Check if admin needs credentials (only for already-bootstrapped systems) - await check_admin_credentials() - return False - - # No admin permission found, need to bootstrap - # Bootstrap creates the admin user AND the reset link, so no need to check credentials after - await bootstrap_system(config=config) - return True + await check_admin_credentials() + return False diff --git a/paskia/db/__init__.py b/paskia/db/__init__.py index 0f0cc21..cdafc6a 100644 --- a/paskia/db/__init__.py +++ b/paskia/db/__init__.py @@ -19,15 +19,7 @@ Usage: """ import paskia.db.operations as operations -from paskia.db.background import ( - start_background, - start_cleanup, - stop_background, - stop_cleanup, -) from paskia.db.bootstrap import bootstrap -from paskia.db.jsonl import load_readonly -from paskia.db.lifecycle import cleanup_expired, init from paskia.db.operations import ( add_permission_to_org, add_permission_to_role, @@ -101,19 +93,11 @@ __all__ = [ "User", # Instance "data", - "init", - "load_readonly", - # Background - "start_background", - "stop_background", - "start_cleanup", - "stop_cleanup", # Read ops # Write ops "add_permission_to_org", "add_permission_to_role", "bootstrap", - "cleanup_expired", "create_credential", "create_credential_session", "create_org", diff --git a/paskia/db/background.py b/paskia/db/background.py index bb75215..b819fa5 100644 --- a/paskia/db/background.py +++ b/paskia/db/background.py @@ -7,12 +7,7 @@ companion task that periodically cleans up expired sessions/tokens. import asyncio import logging -import os -import signal -from kanta.exceptions import DatabaseError - -import paskia.db.operations as _ops from paskia.db.lifecycle import cleanup_expired CLEANUP_INTERVAL = 1 # Expired item cleanup @@ -21,24 +16,6 @@ _logger = logging.getLogger(__name__) _background_task: asyncio.Task | None = None -def _sigterm_on_error(error: DatabaseError) -> None: - """Exit the server when a database write fails.""" - _logger.error("Fatal database error: %s", error) - os.kill(os.getpid(), signal.SIGTERM) - - -async def flush() -> None: - """Write all pending database changes to disk.""" - store = _ops._store - if store is None: - _logger.warning("flush() called but _store is None") - return - try: - await store.flush() - except DatabaseError as e: - _sigterm_on_error(e) - - async def _background_loop(): """Background task that periodically cleans up expired items.""" # Run cleanup immediately on startup to clear old expired items @@ -87,7 +64,7 @@ async def start_background(): async def stop_background(): - """Stop the background cleanup task and close kanta.""" + """Stop the background cleanup task.""" global _background_task if _background_task: _background_task.cancel() @@ -96,12 +73,6 @@ async def stop_background(): except asyncio.CancelledError: pass _background_task = None - store = _ops._store - if store is not None: - try: - await store.close() - except DatabaseError as e: - _sigterm_on_error(e) # Aliases for backwards compatibility diff --git a/paskia/db/bootstrap.py b/paskia/db/bootstrap.py index 2135eec..74f9b45 100644 --- a/paskia/db/bootstrap.py +++ b/paskia/db/bootstrap.py @@ -2,24 +2,60 @@ Bootstrap operations for initial system setup. """ +import logging +import sys from datetime import UTC, datetime import uuid7 -import paskia.db.operations as _ops from paskia.authsession import reset_expires -from paskia.db.structs import Config, Org, Permission, ResetToken, Role, User +from paskia.db.structs import DB, Config, Org, Permission, ResetToken, Role, User from paskia.util.crypto import secret_key +from paskia.util.hostutil import reset_link_url + +_reset_link_logger = logging.getLogger("paskia.reset_link") + + +def _configure_reset_link_logger() -> None: + if _reset_link_logger.handlers: + return + handler = logging.StreamHandler(sys.stderr) + handler.setFormatter(logging.Formatter("%(message)s")) + _reset_link_logger.addHandler(handler) + _reset_link_logger.setLevel(logging.INFO) + _reset_link_logger.propagate = False + + +_configure_reset_link_logger() + +ADMIN_RESET_MESSAGE = """ +👤 Admin %s + - Use this link to register a Passkey for the admin user! +""" + + +def log_reset_link(passphrase: str, message: str | None = None) -> str: + """Log a reset link message and return the URL.""" + reset_link = reset_link_url(passphrase) + if message: + _reset_link_logger.info(message) + _reset_link_logger.info(ADMIN_RESET_MESSAGE, reset_link) + return reset_link def bootstrap( + data: "DB", org_name: str = "Organization", admin_name: str = "Admin", reset_passphrase: str | None = None, reset_expiry: datetime | None = None, config: Config | None = None, ) -> str: - """Bootstrap the entire system in a single transaction. + """Bootstrap the entire system by seeding an empty database. + + This is intended to be called from a ``@kanta.bootstrap`` callback during + ``kanta.open()``. It mutates the provided root ``data`` object directly; + kanta queues the resulting state as the initial "bootstrap" change record. Creates: - auth:admin permission (Master Admin) @@ -29,10 +65,8 @@ def bootstrap( - Reset token for admin registration - Config (if provided) - This is the only way to create a new database file. - All data is created atomically - if any step fails, nothing is written. - Args: + data: The live root database object (usually a ``DB`` instance). org_name: Display name for the organization (default: "Organization") admin_name: Display name for the admin user (default: "Admin") reset_passphrase: Passphrase for the reset token (generated if not provided) @@ -44,7 +78,7 @@ def bootstrap( """ # Check if system is already bootstrapped - for p in _ops._db.permissions.values(): + for p in data.permissions.values(): if p.scope == "auth:admin": raise ValueError( "System already bootstrapped (auth:admin permission exists)" @@ -62,65 +96,66 @@ def bootstrap( if reset_expiry is None: reset_expiry = reset_expires() - with _ops._db.transaction("bootstrap"): - # Create auth:admin permission - perm_admin = Permission( - scope="auth:admin", - display_name="Master Admin", - orgs={org_uuid: True}, # Grant to org - ) - perm_admin.uuid = perm_admin_uuid - perm_admin.store() + # Create auth:admin permission + perm_admin = Permission( + scope="auth:admin", + display_name="Master Admin", + orgs={org_uuid: True}, # Grant to org + ) + perm_admin.uuid = perm_admin_uuid - # Create auth:org:admin permission - perm_org_admin = Permission( - scope="auth:org:admin", - display_name="Org Admin", - orgs={org_uuid: True}, # Grant to org - ) - perm_org_admin.uuid = perm_org_admin_uuid - perm_org_admin.store() + # Create auth:org:admin permission + perm_org_admin = Permission( + scope="auth:org:admin", + display_name="Org Admin", + orgs={org_uuid: True}, # Grant to org + ) + perm_org_admin.uuid = perm_org_admin_uuid - # Create organization - new_org = Org.create(display_name=org_name) - new_org.uuid = org_uuid - new_org.store() + # Create organization + new_org = Org.create(display_name=org_name) + new_org.uuid = org_uuid - # Create Administration role with both permissions - admin_role = Role( - org_uuid=org_uuid, - display_name="Administration", - permissions={perm_admin_uuid: True, perm_org_admin_uuid: True}, - ) - admin_role.uuid = role_uuid - admin_role.store() + # Create Administration role with both permissions + admin_role = Role( + org_uuid=org_uuid, + display_name="Administration", + permissions={perm_admin_uuid: True, perm_org_admin_uuid: True}, + ) + admin_role.uuid = role_uuid - # Create admin user - admin_user = User( - display_name=admin_name, - role_uuid=role_uuid, - created_at=now, - last_seen=None, - visits=0, - theme="", - ) - admin_user.uuid = user_uuid - admin_user.store() + # Create admin user + admin_user = User( + display_name=admin_name, + role_uuid=role_uuid, + created_at=now, + last_seen=None, + visits=0, + theme="", + ) + admin_user.uuid = user_uuid - # Create reset token - reset_token, reset_passphrase = ResetToken.create( - user=user_uuid, - expiry=reset_expiry, - token_type="admin bootstrap", - passphrase=reset_passphrase, - ) - reset_token.store() + # Create reset token + reset_token, reset_passphrase = ResetToken.create( + user=user_uuid, + expiry=reset_expiry, + token_type="admin bootstrap", + passphrase=reset_passphrase, + ) - # Set config if provided - if config is not None: - _ops._db.config = config + # Set config if provided + if config is not None: + data.config = config - # Generate OIDC signing key - _ops._db.oidc.key = secret_key() + # Generate OIDC signing key + data.oidc.key = secret_key() + + # Store all bootstrapped objects in the live data object + data.permissions[perm_admin_uuid] = perm_admin + data.permissions[perm_org_admin_uuid] = perm_org_admin + data.orgs[org_uuid] = new_org + data.roles[role_uuid] = admin_role + data.users[user_uuid] = admin_user + data.reset_tokens[reset_token.key] = reset_token return reset_passphrase diff --git a/paskia/db/jsonl.py b/paskia/db/jsonl.py deleted file mode 100644 index a51e755..0000000 --- a/paskia/db/jsonl.py +++ /dev/null @@ -1,52 +0,0 @@ -""" -JSONL read-only loader using kanta. -""" - -import logging -from pathlib import Path - -import msgspec -from kanta import replay as replay_jsonl -from kanta.migrate import MigrationRegistry - -from paskia.db.migrations import MigrationCtx -from paskia.db.structs import DB, Config - -_logger = logging.getLogger(__name__) - - -def load_readonly(db_path: str, *, rp_id: str = "localhost") -> DB: - """Replay JSONL and apply migrations to produce a DB, without writing anything. - - This is suitable for reading settings before the server starts. - Migrations are applied in-memory only; nothing is queued or flushed. - """ - path = Path(db_path) - if not path.exists(): - return DB(config=Config(rp_id=rp_id)) - - try: - content = path.read_bytes() - rr = replay_jsonl(content) - data_dict = rr.state - version = rr.version - - if not data_dict: - return DB(config=Config(rp_id=rp_id)) - - # Apply migrations in-memory (no persistence) - registry = MigrationRegistry.from_module("paskia.db.migrations") - version = registry.apply( - data_dict, version, MigrationCtx(rp_id=rp_id), silent=True - ) - - # Decode to msgspec struct - return msgspec.json.decode(msgspec.json.encode(data_dict), type=DB) - except OSError as e: - _logger.exception("Failed to load database") - raise SystemExit(f"{e}") - except (ValueError, msgspec.DecodeError) as e: - raise SystemExit(f"{e}") - except Exception as e: - _logger.exception("Unexpected error loading database") - raise SystemExit(f"{e}") diff --git a/paskia/db/lifecycle.py b/paskia/db/lifecycle.py index c1474e6..103ce2a 100644 --- a/paskia/db/lifecycle.py +++ b/paskia/db/lifecycle.py @@ -2,10 +2,14 @@ Database lifecycle: initialization and maintenance. """ +import asyncio import logging import os import signal from datetime import UTC, datetime +from pathlib import Path +from typing import Annotated, Any, Optional +from uuid import UUID from kanta import Kanta from kanta.exceptions import DatabaseError @@ -13,58 +17,136 @@ from kanta.exceptions import DatabaseError import paskia.db.operations as _ops from paskia import oidc_notify from paskia.authsession import EXPIRES -from paskia.db.migrations import MigrationCtx +from paskia.db.bootstrap import bootstrap, log_reset_link from paskia.db.paths import db_file_path from paskia.db.structs import DB +from paskia.util.runtime import config as runtime_config -_logger = logging.getLogger(__name__) +logger = logging.getLogger(__name__) -def _fatal_error(error: DatabaseError) -> None: +runtime = runtime_config() +if runtime is None: + raise RuntimeError("PASKIA_CONFIG must be defined before importing db.lifecycle") + +kanta = Kanta( + str(db_file_path(rp_id=runtime.config.rp_id, create_root=False)), + _ops._db, + migrations="paskia.db.migrations", +) +kanta.ctx.rp_id = runtime.config.rp_id +_ops._db._store = kanta + + +def _lookup_uuid_in_state(state: dict | None, uuid_str: str) -> str | None: + """Resolve UUID to label from serialized state dict.""" + if not state: + return None + + # Display-name based entities. + for bucket in ("users", "orgs", "roles", "permissions"): + entity = state.get(bucket, {}).get(uuid_str) + if isinstance(entity, dict): + display_name = entity.get("display_name") + if isinstance(display_name, str) and display_name: + return display_name + + # OIDC clients use "name" instead of "display_name". + client = state.get("oidc", {}).get("clients", {}).get(uuid_str) + if isinstance(client, dict): + name = client.get("name") + if isinstance(name, str) and name: + return name + + return None + + +def _resolve_uuid_label( + uuid_str: str, + *, + previous: dict | None = None, + current: dict | None = None, +) -> str | None: + """Resolve known entity UUIDs to human-readable labels.""" + # Prefer previous state so deletions/renames still show a useful label. + label = _lookup_uuid_in_state(previous, uuid_str) + if label: + return label + label = _lookup_uuid_in_state(current, uuid_str) + if label: + return label + + try: + uid = UUID(uuid_str) + except ValueError: + return None + + if uid in _ops._db.users: + return _ops._db.users[uid].display_name + if uid in _ops._db.orgs: + return _ops._db.orgs[uid].display_name + if uid in _ops._db.roles: + return _ops._db.roles[uid].display_name + if uid in _ops._db.permissions: + return _ops._db.permissions[uid].display_name + if uid in _ops._db.oidc.clients: + return _ops._db.oidc.clients[uid].name + return None + + +@kanta.logfmt +def format_log_uuid( + value: Any, + path: str, + previous: Annotated[dict, "pre"] | None = None, + current: Annotated[dict, "post"] | None = None, +) -> Optional[str]: # noqa: UP045 + """Format UUID values/keys/actor labels in transaction logs.""" + if not isinstance(value, str): + return None + + # Works for transaction actor metadata ($user), values, and path components. + return _resolve_uuid_label(value, previous=previous, current=current) + + +@kanta.fatal_error +def terminate(error: DatabaseError) -> None: """Fatal error callback: terminate the process on background write failures.""" - _logger.error("Fatal database error: %s", error) + logger.error("Fatal database error: %s", error) os.kill(os.getpid(), signal.SIGTERM) -async def init(rp_id: str, *args, **kwargs): - """Load database from JSONL file using kanta.""" - if _ops._store is not None: - _logger.debug("Database already initialized, skipping reload") - return - db_path = db_file_path(rp_id=rp_id, create_root=True) - db = DB() - kanta = Kanta( - str(db_path), - db, - migrations="paskia.db.migrations", - migration_ctx=MigrationCtx(rp_id=rp_id), - fatal_error=_fatal_error, - ) +@kanta.bootstrap +def bootstrap_db(data: DB) -> None: + reset_passphrase = bootstrap(data, config=runtime.config) + log_reset_link(reset_passphrase, "✅ Bootstrap completed!") + + +async def init(): + """Load database from JSONL file using kanta. + + If the database file is empty, the configured bootstrap callback seeds it + with default permissions, organization, role, admin user and a reset token. + """ + rootpath = Path(kanta.filename).parent try: + await asyncio.to_thread(rootpath.mkdir, parents=True, exist_ok=True) await kanta.open() - except DatabaseError as e: + except Exception as e: raise SystemExit(f"{e}") from e - _ops._store = kanta - _ops._db = db - _ops._db._store = kanta - # Request a snapshot after successful startup - kanta.request_snapshot() def cleanup_expired() -> int: """Remove expired sessions and reset tokens. Returns count removed.""" now = datetime.now(UTC) - count = 0 limit = now - EXPIRES expired_sessions = [k for k, s in _ops._db.sessions.items() if s.validated < limit] if expired_sessions: oidc_notify.schedule_notifications(expired_sessions) - with _ops._db.transaction("expiry"): + with kanta.transaction("expiry"): for k in expired_sessions: del _ops._db.sessions[k] - count += 1 expired_tokens = [k for k, t in _ops._db.reset_tokens.items() if t.expiry < now] for k in expired_tokens: del _ops._db.reset_tokens[k] - count += 1 - return count + return len(expired_sessions) + len(expired_tokens) diff --git a/paskia/db/migrations.py b/paskia/db/migrations.py index 8b1ca61..ff2ce33 100644 --- a/paskia/db/migrations.py +++ b/paskia/db/migrations.py @@ -7,35 +7,30 @@ Each migration should be idempotent and only run when needed. import base64 +from kanta import Kanta + from paskia.util.crypto import secret_key -class MigrationCtx: - """Context passed to each migration function.""" - - def __init__(self, rp_id: str): - self.rp_id = rp_id - - -def migrate_v1(d: dict, ctx: MigrationCtx) -> None: +def migrate_v1(d: dict) -> None: """Remove Org.created_at fields.""" for org_data in d["orgs"].values(): org_data.pop("created_at", None) -def migrate_v2(d: dict, ctx: MigrationCtx) -> None: +def migrate_v2(d: dict, kanta: Kanta) -> None: """Add config field if missing.""" if "config" not in d: - d["config"] = {"rp_id": ctx.rp_id} + d["config"] = {"rp_id": kanta.ctx.rp_id} -def migrate_v3(d: dict, ctx: MigrationCtx) -> None: +def migrate_v3(d: dict) -> None: """Ensure all users have visits field.""" for user_data in d["users"].values(): user_data.setdefault("visits", 0) -def migrate_v4(d: dict, ctx: MigrationCtx) -> None: +def migrate_v4(d: dict) -> None: """OpenID Connect support and hardened session keys.""" # Session keys changed to hashes, drop old sessions d["sessions"] = {} @@ -46,7 +41,7 @@ def migrate_v4(d: dict, ctx: MigrationCtx) -> None: } -def migrate_v5(d: dict, ctx: MigrationCtx) -> None: +def migrate_v5(d: dict) -> None: """Convert config.listen from str to list[str] if needed.""" listen = d["config"].get("listen") if listen and isinstance(listen, str): diff --git a/paskia/db/operations.py b/paskia/db/operations.py index 69fd38e..6853ed3 100644 --- a/paskia/db/operations.py +++ b/paskia/db/operations.py @@ -12,7 +12,6 @@ from datetime import UTC, datetime, timedelta from uuid import UUID import uuid7 -from kanta import Kanta from paskia import oidc_notify from paskia.config import SESSION_LIFETIME @@ -39,7 +38,26 @@ _UNSET = object() # Global database instance (empty until init() loads data) _db = DB(config=Config(rp_id="uninitialized.invalid")) -_store: Kanta[DB] | None = None + + +def _store(): + """Return active Kanta instance for the current DB object.""" + store = _db._store + if store is None: + raise RuntimeError("Kanta store is not initialized") + return store + + +def _transaction( + action: str, + ctx: SessionContext | None = None, + *, + user: str | None = None, + mtime: bool | datetime = True, +): + """Create a Kanta transaction with minimal metadata mapping.""" + user_id = str(ctx.user.uuid) if ctx else user + return _store().transaction(action, user=user_id, mtime=mtime) def is_username_taken(username: str, exclude_uuid: UUID | None = None) -> bool: @@ -60,7 +78,7 @@ def is_username_taken(username: str, exclude_uuid: UUID | None = None) -> bool: def update_config(config: Config) -> None: """Update the stored configuration.""" - with _db.transaction("update_config"): + with _transaction("update_config"): _db.config = config @@ -68,7 +86,7 @@ def create_permission(perm: Permission, *, ctx: SessionContext | None = None) -> """Create a new permission.""" if perm.uuid in _db.permissions: raise ValueError(f"Permission {perm.uuid} already exists") - with _db.transaction("admin:create_permission", ctx): + with _transaction("admin:create_permission", ctx): perm.store() @@ -86,7 +104,7 @@ def update_permission( """ if uuid not in _db.permissions: raise ValueError(f"Permission {uuid} not found") - with _db.transaction("admin:update_permission", ctx): + with _transaction("admin:update_permission", ctx): _db.permissions[uuid].scope = scope _db.permissions[uuid].display_name = display_name _db.permissions[uuid].domain = domain @@ -96,7 +114,7 @@ def delete_permission(uuid: UUID, *, ctx: SessionContext | None = None) -> None: """Delete a permission and remove it from all roles.""" if uuid not in _db.permissions: raise ValueError(f"Permission {uuid} not found") - with _db.transaction("admin:delete_permission", ctx): + with _transaction("admin:delete_permission", ctx): _db.permissions[uuid].delete() @@ -108,7 +126,7 @@ def create_org(org: Org, *, ctx: SessionContext | None = None) -> None: if org.uuid in _db.orgs: raise ValueError(f"Organization {org.uuid} already exists") now = datetime.now(UTC) - with _db.transaction("admin:create_org", ctx): + with _transaction("admin:create_org", ctx): new_org = Org.create(display_name=org.display_name, created_at=now) new_org.uuid = org.uuid new_org.store() @@ -140,7 +158,7 @@ def update_org_name( """Update organization display name.""" if uuid not in _db.orgs: raise ValueError(f"Organization {uuid} not found") - with _db.transaction("admin:update_org_name", ctx): + with _transaction("admin:update_org_name", ctx): _db.orgs[uuid].display_name = display_name @@ -148,7 +166,7 @@ def delete_org(uuid: UUID, *, ctx: SessionContext | None = None) -> None: """Delete organization and all its roles/users.""" if uuid not in _db.orgs: raise ValueError(f"Organization {uuid} not found") - with _db.transaction("admin:delete_org", ctx): + with _transaction("admin:delete_org", ctx): _db.orgs[uuid].delete() @@ -165,7 +183,7 @@ def add_permission_to_org( if permission_uuid not in _db.permissions: raise ValueError(f"Permission {permission_uuid} not found") - with _db.transaction("admin:add_permission_to_org", ctx): + with _transaction("admin:add_permission_to_org", ctx): _db.permissions[permission_uuid].orgs[org_uuid] = True @@ -182,7 +200,7 @@ def remove_permission_from_org( if permission_uuid not in _db.permissions: return # Permission not found, silently return - with _db.transaction("admin:remove_permission_from_org", ctx): + with _transaction("admin:remove_permission_from_org", ctx): _db.permissions[permission_uuid].orgs.pop(org_uuid, None) @@ -192,7 +210,7 @@ def create_role(role: Role, *, ctx: SessionContext | None = None) -> None: raise ValueError(f"Role {role.uuid} already exists") if role.org_uuid not in _db.orgs: raise ValueError(f"Organization {role.org_uuid} not found") - with _db.transaction("admin:create_role", ctx): + with _transaction("admin:create_role", ctx): role.store() @@ -205,7 +223,7 @@ def update_role_name( """Update role display name.""" if uuid not in _db.roles: raise ValueError(f"Role {uuid} not found") - with _db.transaction("admin:update_role_name", ctx): + with _transaction("admin:update_role_name", ctx): _db.roles[uuid].display_name = display_name @@ -220,7 +238,7 @@ def add_permission_to_role( raise ValueError(f"Role {role_uuid} not found") if permission_uuid not in _db.permissions: raise ValueError(f"Permission {permission_uuid} not found") - with _db.transaction("admin:add_permission_to_role", ctx): + with _transaction("admin:add_permission_to_role", ctx): _db.roles[role_uuid].permissions[permission_uuid] = True @@ -233,7 +251,7 @@ def remove_permission_from_role( """Remove permission from role by UUID.""" if role_uuid not in _db.roles: raise ValueError(f"Role {role_uuid} not found") - with _db.transaction("admin:remove_permission_from_role", ctx): + with _transaction("admin:remove_permission_from_role", ctx): _db.roles[role_uuid].permissions.pop(permission_uuid, None) @@ -245,7 +263,7 @@ def delete_role(uuid: UUID, *, ctx: SessionContext | None = None) -> None: role = _db.roles[uuid] if role.users: raise ValueError(f"Cannot delete role {uuid}: users still assigned") - with _db.transaction("admin:delete_role", ctx): + with _transaction("admin:delete_role", ctx): _db.roles[uuid].delete() @@ -255,7 +273,7 @@ def create_user(new_user: User, *, ctx: SessionContext | None = None) -> None: raise ValueError(f"User {new_user.uuid} already exists") if new_user.role_uuid not in _db.roles: raise ValueError(f"Role {new_user.role_uuid} not found") - with _db.transaction("admin:create_user", ctx): + with _transaction("admin:create_user", ctx): new_user.store() @@ -282,7 +300,7 @@ def update_user_display_name( if not display_name: raise ValueError("Display name cannot be empty") user = _db.users[uuid] - with _db.transaction("update_user_display_name", ctx): + with _transaction("update_user_display_name", ctx): user.display_name = display_name # Auto-fill preferred_username if not already set if user.preferred_username is None: @@ -356,7 +374,7 @@ def update_user_info( elif len(telephone) > 32: raise ValueError("telephone too long") - with _db.transaction("update_user_info", ctx): + with _transaction("update_user_info", ctx): if display_name is not _UNSET: user.display_name = display_name if theme is not _UNSET: @@ -380,7 +398,7 @@ def update_user_role( raise ValueError(f"User {uuid} not found") if role_uuid not in _db.roles: raise ValueError(f"Role {role_uuid} not found") - with _db.transaction("admin:update_user_role", ctx): + with _transaction("admin:update_user_role", ctx): _db.users[uuid].role_uuid = role_uuid @@ -388,7 +406,7 @@ def delete_user(uuid: UUID, *, ctx: SessionContext | None = None) -> None: """Delete user and their credentials/sessions.""" if uuid not in _db.users: raise ValueError(f"User {uuid} not found") - with _db.transaction("admin:delete_user", ctx): + with _transaction("admin:delete_user", ctx): _db.users[uuid].delete() @@ -398,7 +416,7 @@ def create_credential(cred: Credential, *, ctx: SessionContext | None = None) -> raise ValueError(f"Credential {cred.uuid} already exists") if cred.user_uuid not in _db.users: raise ValueError(f"User {cred.user_uuid} not found") - with _db.transaction("create_credential", ctx): + with _transaction("create_credential", ctx): cred.store() @@ -412,7 +430,7 @@ def update_credential_sign_count( """Update credential sign count and last_used.""" if uuid not in _db.credentials: raise ValueError(f"Credential {uuid} not found") - with _db.transaction("update_credential_sign_count", ctx): + with _transaction("update_credential_sign_count", ctx): _db.credentials[uuid].sign_count = sign_count if last_used: _db.credentials[uuid].last_used = last_used @@ -434,7 +452,7 @@ def delete_credential( if user_uuid is not None: if cred.user_uuid != user_uuid: raise ValueError(f"Credential {uuid} does not belong to user {user_uuid}") - with _db.transaction("delete_credential", ctx): + with _transaction("delete_credential", ctx): cred.delete() @@ -450,7 +468,7 @@ def update_session( """Update session metadata.""" if key not in _db.sessions: raise ValueError("Session not found") - with _db.transaction("update_session", ctx): + with _transaction("update_session", ctx): s = _db.sessions[key] if host is not None: s.host = host @@ -480,7 +498,7 @@ def delete_session( raise ValueError("Session not found") oidc_notify.schedule_notifications([key]) - with _db.transaction(action, ctx): + with _transaction(action, ctx): _db.sessions[key].delete() @@ -499,7 +517,7 @@ def delete_sessions_for_user( keys = [s.key for s in user.sessions] oidc_notify.schedule_notifications(keys) - with _db.transaction("admin:delete_sessions_for_user", ctx): + with _transaction("admin:delete_sessions_for_user", ctx): for sess in user.sessions: sess.delete() @@ -530,7 +548,7 @@ def create_reset_token( ) if token.key in _db.reset_tokens: raise ValueError("Reset token already exists") - with _db.transaction("create_reset_token", ctx, user=user): + with _transaction("create_reset_token", ctx, user=user): token.store() return passphrase @@ -539,7 +557,7 @@ def delete_reset_token(key: bytes, *, ctx: SessionContext | None = None) -> None """Delete a reset token.""" if key not in _db.reset_tokens: raise ValueError("Reset token not found") - with _db.transaction("delete_reset_token", ctx): + with _transaction("delete_reset_token", ctx): _db.reset_tokens[key].delete() @@ -588,7 +606,7 @@ def login( validated=now, ) user_str = str(user_uuid) - with _db.transaction("login", user=user_str): + with _transaction("login", user=user_str): session.store(now) # Update credential _db.credentials[credential_uuid].sign_count = sign_count @@ -615,7 +633,7 @@ def oidc_login( """ now = datetime.now(UTC) user_str = str(session.user_uuid) - with _db.transaction("oidc_login", user=user_str): + with _transaction("oidc_login", user=user_str): session.store(now) # Update credential _db.credentials[credential_uuid].sign_count = sign_count @@ -661,7 +679,7 @@ def create_credential_session( validated=now, ) user_str = str(user_uuid) - with _db.transaction("create_credential_session", user=user_str): + with _transaction("create_credential_session", user=user_str): # Update display name if provided if display_name: _db.users[user_uuid].display_name = display_name @@ -694,7 +712,7 @@ def create_oid_client(client: Client, *, ctx: SessionContext | None = None) -> N """Create a new OIDC client.""" if client.uuid in _db.oidc.clients: raise ValueError(f"OIDC client {client.uuid} already exists") - with _db.transaction("admin:create_oid_client", ctx): + with _transaction("admin:create_oid_client", ctx): _db.oidc.clients[client.uuid] = client @@ -735,7 +753,7 @@ def update_oid_client( else client.backchannel_logout_uri ) - with _db.transaction("admin:update_oid_client", ctx): + with _transaction("admin:update_oid_client", ctx): # Create updated client with new values updated_client = Client( client_secret_hash=secret_hash @@ -761,7 +779,7 @@ def reset_oid_client_secret( if client_uuid not in _db.oidc.clients: raise ValueError(f"OIDC client {client_uuid} not found") client = _db.oidc.clients[client_uuid] - with _db.transaction("admin:reset_oid_client_secret", ctx): + with _transaction("admin:reset_oid_client_secret", ctx): updated = Client( client_secret_hash=new_secret_hash, name=client.name, @@ -776,5 +794,5 @@ def delete_oid_client(client_uuid: UUID, *, ctx: SessionContext | None = None) - """Delete an OIDC client.""" 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): + with _transaction("admin:delete_oid_client", ctx): del _db.oidc.clients[client_uuid] diff --git a/paskia/db/structs.py b/paskia/db/structs.py index 03f4833..d9e11a0 100644 --- a/paskia/db/structs.py +++ b/paskia/db/structs.py @@ -3,13 +3,13 @@ from __future__ import annotations import hashlib import secrets from datetime import UTC, datetime +from typing import Any from uuid import UUID import msgspec import uuid7 from paskia import db -from paskia.db.logging import UuidResolver from paskia.util import passphrase as passphrase_util from paskia.util.crypto import hash_secret @@ -631,8 +631,8 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False): oidc: OIDC = msgspec.field(default_factory=lambda: OIDC()) def __post_init__(self): - # Store reference for persistence (not serialized) - self._store = None + # Optional store reference for non-global DB instances (e.g. tests). + self._store: Any | None = None # Set the key fields on all stored objects for uuid, perm in self.permissions.items(): perm.uuid = uuid @@ -652,23 +652,6 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False): for uuid, client in self.oidc.clients.items(): client.uuid = uuid - def transaction(self, action, ctx=None, *, user=None): - """Wrap writes in transaction. Delegates to Kanta.""" - user_id = str(ctx.user.uuid) if ctx else user - user_display = None - if user_id: - try: - user_uuid = UUID(user_id) - if user_uuid in self.users: - user_display = self.users[user_uuid].display_name - except (ValueError, KeyError): - user_display = user_id - previous_state = msgspec.to_builtins(self) - resolver = UuidResolver(self, previous_state).resolve - return self._store.transaction( - action, user=user_id, user_display=user_display, resolver=resolver - ) - def session_ctx( self, session_secret: str, host: str | None = None ) -> SessionContext | None: diff --git a/paskia/fastapi/__init__.py b/paskia/fastapi/__init__.py index 552e4a1..e69de29 100644 --- a/paskia/fastapi/__init__.py +++ b/paskia/fastapi/__init__.py @@ -1,3 +0,0 @@ -from paskia.fastapi.mainapp import app - -__all__ = ["app"] diff --git a/paskia/fastapi/admin/permissions.py b/paskia/fastapi/admin/permissions.py index 9d6cee5..434b34e 100644 --- a/paskia/fastapi/admin/permissions.py +++ b/paskia/fastapi/admin/permissions.py @@ -28,7 +28,7 @@ def _validate_permission_domain(domain: str | None) -> None: except ValueError: pass - rp_id = passkey.instance.rp_id + rp_id = passkey.rp_id if domain == rp_id or domain.endswith(f".{rp_id}"): return raise ValueError( diff --git a/paskia/fastapi/admin/server_config.py b/paskia/fastapi/admin/server_config.py index 4f0683a..79007a3 100644 --- a/paskia/fastapi/admin/server_config.py +++ b/paskia/fastapi/admin/server_config.py @@ -22,7 +22,7 @@ async def admin_get_server_config( ): """Get current server configuration (master admin only).""" await authz.verify(auth, ["auth:admin"], host=request.headers.get("host")) - pk = passkey.instance + pk = passkey config = db.data().config return { "rp_name": pk.rp_name, @@ -46,7 +46,7 @@ async def admin_update_server_config( auth, ["auth:admin"], host=request.headers.get("host"), max_age="5m" ) config = db.data().config - pk = passkey.instance + pk = passkey rp_name = payload.get("rp_name", "").strip() or None auth_host = payload.get("auth_host", "").strip() or None diff --git a/paskia/fastapi/api.py b/paskia/fastapi/api.py index dc4f2bc..e0d923d 100644 --- a/paskia/fastapi/api.py +++ b/paskia/fastapi/api.py @@ -256,7 +256,7 @@ async def forward_authentication( @app.get("/settings") async def get_settings(): - pk = global_passkey.instance + pk = global_passkey base_path = hostutil.ui_base_path() return MsgspecResponse( ApiSettings( diff --git a/paskia/fastapi/mainapp.py b/paskia/fastapi/mainapp.py index 257ab52..3c4e45c 100644 --- a/paskia/fastapi/mainapp.py +++ b/paskia/fastapi/mainapp.py @@ -1,3 +1,4 @@ +import asyncio import logging import os from contextlib import asynccontextmanager @@ -7,11 +8,10 @@ import msgspec from fastapi import FastAPI, HTTPException, Request, Response from fastapi.responses import FileResponse, RedirectResponse -from paskia import authcode, db, globals -from paskia.__main__ import DEVMODE +from paskia import authcode, db, remoteauth from paskia.bootstrap import bootstrap_if_needed -from paskia.db import start_background, stop_background -from paskia.db.background import flush +from paskia.db.background import start_background, stop_background +from paskia.db.lifecycle import kanta from paskia.db.logging import configure_db_logging from paskia.fastapi import admin, api, auth_host, oid, ws from paskia.fastapi.admin.adminapp import adminapp @@ -21,6 +21,7 @@ from paskia.fastapi.front import frontend from paskia.fastapi.logging import AccessLogMiddleware, configure_access_logging from paskia.fastapi.session import AUTH_COOKIE from paskia.util import hostutil, passphrase, vitedev +from paskia.util.constants import DEVMODE from paskia.util.runtime import RuntimeConfig # Configure custom logging @@ -43,34 +44,33 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path """ runtime = msgspec.json.decode(os.environ["PASKIA_CONFIG"], type=RuntimeConfig) - try: - await globals.init( - rp_id=runtime.config.rp_id, - rp_name=runtime.config.rp_name, - origins=runtime.config.origins, - bootstrap=False, - ) - except ValueError as e: - logging.error(f"⚠️ {e}") - # Re-raise to fail fast - raise + await asyncio.to_thread( + Path(kanta.filename).parent.mkdir, parents=True, exist_ok=True + ) + async with kanta: + try: + await remoteauth.init() + await authcode.start() + except ValueError as e: + logging.error(f"⚠️ {e}") + # Re-raise to fail fast + raise - # Bootstrap and persist config now that the full DB is loaded - await bootstrap_if_needed(config=runtime.config) - if runtime.save: - db.update_config(runtime.config) - await flush() + # Bootstrap and persist config now that the full DB is loaded + await bootstrap_if_needed(config=runtime.config) + if runtime.save: + db.update_config(runtime.config) - # Restore uvicorn info logging (suppressed during startup in dev mode) - # Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages - if app.debug: - logging.getLogger("uvicorn").setLevel(logging.INFO) - logging.getLogger("uvicorn.error").setLevel(logging.WARNING) - await frontend.load() - await start_background() - yield - await stop_background() - await authcode.stop() + # Restore uvicorn info logging (suppressed during startup in dev mode) + # Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages + if app.debug: + logging.getLogger("uvicorn").setLevel(logging.INFO) + logging.getLogger("uvicorn.error").setLevel(logging.WARNING) + await frontend.load() + await start_background() + yield + await stop_background() + await authcode.stop() app = FastAPI( diff --git a/paskia/fastapi/ws.py b/paskia/fastapi/ws.py index 94e7907..7fe010a 100644 --- a/paskia/fastapi/ws.py +++ b/paskia/fastapi/ws.py @@ -58,7 +58,7 @@ async def websocket_register_add( if reset is not None: if not passphrase.is_well_formed(reset): raise ValueError( - f"The reset link for {passkey.instance.rp_name} is invalid or has expired" + f"The reset link for {passkey.rp_name} is invalid or has expired" ) s = get_reset(reset) user_uuid = s.user_uuid diff --git a/paskia/fastapi/wschat.py b/paskia/fastapi/wschat.py index 50bbf2c..1f9ab04 100644 --- a/paskia/fastapi/wschat.py +++ b/paskia/fastapi/wschat.py @@ -23,14 +23,14 @@ async def register_chat( credential_ids: list[bytes] | None = None, ): """Run WebAuthn registration flow and return the verified credential.""" - options, challenge = passkey.instance.reg_generate_options( + options, challenge = passkey.reg_generate_options( user_id=user_uuid, user_name=user_name, credential_ids=credential_ids, ) await ws.send_json({"optionsJSON": options}) response = await ws.receive_json() - return passkey.instance.reg_verify(response, challenge, user_uuid, origin=origin) + return passkey.reg_verify(response, challenge, user_uuid, origin=origin) async def authenticate_chat( @@ -43,11 +43,9 @@ async def authenticate_chat( tuple of (credential, new_sign_count) where new_sign_count comes from WebAuthn verification """ origin = validate_origin(ws) - options, challenge = passkey.instance.auth_generate_options( - credential_ids=credential_ids - ) + options, challenge = passkey.auth_generate_options(credential_ids=credential_ids) await ws.send_json({"optionsJSON": options}) - authcred = passkey.instance.auth_parse(await ws.receive_json()) + authcred = passkey.auth_parse(await ws.receive_json()) cred = next( ( @@ -58,11 +56,9 @@ async def authenticate_chat( None, ) if not cred: - raise ValueError( - f"This passkey is no longer registered with {passkey.instance.rp_name}" - ) + raise ValueError(f"This passkey is no longer registered with {passkey.rp_name}") - verification = passkey.instance.auth_verify(authcred, challenge, cred, origin) + verification = passkey.auth_verify(authcred, challenge, cred, origin) return cred, verification.new_sign_count @@ -94,7 +90,7 @@ async def authenticate_and_login( if not normalized_host: raise ValueError("Host required for session creation") hostname = normalized_host.split(":")[0] - rp_id = passkey.instance.rp_id + rp_id = passkey.rp_id if not (hostname == rp_id or hostname.endswith(f".{rp_id}")): raise ValueError(f"Host must be the same as or a subdomain of {rp_id}") metadata = infodict(ws, "auth") diff --git a/paskia/fastapi/wsutil.py b/paskia/fastapi/wsutil.py index 158f867..d2ebec4 100644 --- a/paskia/fastapi/wsutil.py +++ b/paskia/fastapi/wsutil.py @@ -96,4 +96,4 @@ def validate_origin(ws: WebSocket) -> str: origin = ws.headers.get("origin") if not origin: raise ValueError("Origin header is required for WebSocket connections") - return passkey.instance.validate_origin(origin) + return passkey.validate_origin(origin) diff --git a/paskia/globals.py b/paskia/globals.py index d4ce0cb..c852698 100644 --- a/paskia/globals.py +++ b/paskia/globals.py @@ -1,71 +1,20 @@ -from typing import Generic, TypeVar +"""Global Passkey instance configured from PASKIA_CONFIG. + +The Passkey instance is created at import time using the runtime configuration +passed via the ``PASKIA_CONFIG`` environment variable. Other runtime setup +(remote auth, auth codes, bootstrap checks) is performed explicitly by the +FastAPI lifespan once the database is open. +""" -from paskia import authcode, db, remoteauth -from paskia.bootstrap import bootstrap_if_needed from paskia.sansio import Passkey +from paskia.util import runtime -T = TypeVar("T") +runtime = runtime.config() +if runtime is None: + raise RuntimeError("PASKIA_CONFIG must be defined before importing paskia.globals") - -class Manager(Generic[T]): - """Generic manager for global instances.""" - - def __init__(self, name: str): - self._instance: T | None = None - self._name = name - - @property - def instance(self) -> T: - if self._instance is None: - raise RuntimeError( - f"{self._name} not initialized. Call globals.init() first." - ) - return self._instance - - @instance.setter - def instance(self, instance: T) -> None: - self._instance = instance - - -async def init( - rp_id: str = "localhost", - rp_name: str | None = None, - origins: list[str] | None = None, - *, - bootstrap: bool = True, -) -> None: - """Initialize global passkey + database. - - If bootstrap=True (default) the system bootstrap_if_needed() will be invoked. - In FastAPI lifespan we call with bootstrap=False to avoid duplicate bootstrapping - since the CLI performs it once before servers start. - - Database configuration: - Set PASKIA_DB environment variable to specify the JSONL database file path. - Default: {rp_id}.paskiadb - """ - - # Initialize passkey instance with provided parameters - passkey.instance = Passkey( - rp_id=rp_id, - rp_name=rp_name or rp_id, - origins=origins, - ) - - # Initialize database - await db.init(rp_id=rp_id) - - # Initialize remote auth manager - await remoteauth.init() - - # Initialize auth code manager - await authcode.start() - - if bootstrap: - # Bootstrap system if needed - - await bootstrap_if_needed() - - -# Global instances -passkey = Manager[Passkey]("Passkey") +passkey = Passkey( + rp_id=runtime.config.rp_id, + rp_name=runtime.config.rp_name, + origins=runtime.config.origins, +) diff --git a/paskia/oidc_notify.py b/paskia/oidc_notify.py index 3b1d320..1c77fad 100644 --- a/paskia/oidc_notify.py +++ b/paskia/oidc_notify.py @@ -13,7 +13,7 @@ import httpx from paskia import db from paskia.util import oidjwt -from paskia.util.runtime import _load_config +from paskia.util.runtime import config as runtime_config _logger = logging.getLogger(__name__) @@ -23,8 +23,8 @@ _TIMEOUT = httpx.Timeout(10.0, connect=5.0) def _issuer() -> str: """Derive issuer URL from config (same base as discovery document).""" - cfg = _load_config() - return cfg.get("site_url", "https://localhost") + cfg = runtime_config() + return cfg.site_url if cfg else "https://localhost" def _collect_oidc_sessions( diff --git a/paskia/util/constants.py b/paskia/util/constants.py new file mode 100644 index 0000000..d7c9f9d --- /dev/null +++ b/paskia/util/constants.py @@ -0,0 +1,6 @@ +"""Small, dependency-free constants shared by CLI and server modules.""" + +import os + +DEFAULT_PORT = 4401 +DEVMODE = os.getenv("PASKIA_DEV") == "1" diff --git a/paskia/util/hostutil.py b/paskia/util/hostutil.py index c8c5498..5577e42 100644 --- a/paskia/util/hostutil.py +++ b/paskia/util/hostutil.py @@ -2,11 +2,12 @@ from urllib.parse import urlparse, urlsplit -from paskia.util.runtime import _load_config +from paskia.util.runtime import clear_config_cache +from paskia.util.runtime import config as runtime_config def _cfg(): - return _load_config() + return runtime_config() def is_root_mode() -> bool: @@ -105,7 +106,7 @@ def normalize_auth_host_and_origins( def reload_config() -> None: - _load_config.cache_clear() + clear_config_cache() def normalize_host(raw_host: str | None) -> str | None: diff --git a/paskia/util/oidjwt.py b/paskia/util/oidjwt.py index 7c3de54..1a0e55a 100644 --- a/paskia/util/oidjwt.py +++ b/paskia/util/oidjwt.py @@ -29,11 +29,14 @@ def _load_or_generate_key() -> None: global _private_key, _public_key, _kid data = db.data() + store = data._store + if store is None: + raise RuntimeError("Kanta store is not initialized") if data.oidc.key is not None: _private_key = public_key_from_secret(data.oidc.key) else: raw_key = secret_key() - with data.transaction("oidc_key"): + with store.transaction("oidc_key"): data.oidc.key = raw_key _private_key = public_key_from_secret(raw_key) diff --git a/paskia/util/runtime.py b/paskia/util/runtime.py index d338b03..e58945d 100644 --- a/paskia/util/runtime.py +++ b/paskia/util/runtime.py @@ -31,9 +31,19 @@ def _load_config() -> "RuntimeConfig | None": return msgspec.json.decode(config_json.encode(), type=RuntimeConfig) +def config() -> "RuntimeConfig | None": + """Return cached runtime config loaded from PASKIA_CONFIG.""" + return _load_config() + + +def clear_config_cache() -> None: + """Clear cached runtime config; next config() call reloads from env.""" + _load_config.cache_clear() + + def update_runtime_config(new_config: Config) -> None: """Update the runtime configuration with a new Config and refresh the cache.""" - current_runtime = _load_config() + current_runtime = config() if not current_runtime: return # No runtime config to update @@ -56,4 +66,4 @@ def update_runtime_config(new_config: Config) -> None: os.environ["PASKIA_CONFIG"] = msgspec.json.encode(new_runtime).decode() # Clear the cache so next access loads the updated config - _load_config.cache_clear() + clear_config_cache() diff --git a/paskia/util/startupbox.py b/paskia/util/startupbox.py index b1f4136..a0f45de 100644 --- a/paskia/util/startupbox.py +++ b/paskia/util/startupbox.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING from fastapi_vue.hostutil import parse_endpoints from paskia._version import __version__ +from paskia.util.constants import DEFAULT_PORT, DEVMODE from paskia.util.hostutil import format_endpoint if TYPE_CHECKING: @@ -73,16 +74,13 @@ def print_startup_config(runtime: RuntimeConfig) -> None: if runtime.config.auth_host: lines.append(line(f"Auth Host: {runtime.config.auth_host}")) - from paskia.__main__ import DEFAULT_PORT as P # noqa: PLC0415 - circular - from paskia.__main__ import DEVMODE # noqa: PLC0415 - circular - # Show frontend URL if in dev mode if DEVMODE: lines.append(line(f"Dev Frontend: {os.environ.get('PASKIA_VITE_URL')}")) # Format listen endpoints (dev mode only uses the first endpoint) - endpoints = list(parse_endpoints(runtime.config.listen, P)) + endpoints = list(parse_endpoints(runtime.config.listen, DEFAULT_PORT)) if DEVMODE: endpoints = endpoints[:1] # server.run reload=True uses only one parts = [format_endpoint(ep) for ep in endpoints] diff --git a/pyproject.toml b/pyproject.toml index 6ddc3d1..ea92a45 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,7 +23,7 @@ dependencies = [ "msgspec>=0.20.0", "fastapi-vue>=1.1.0", "ua-parser[regex]>=1.0.1", - "kanta>=0.1.1", + "kanta>=0.4.0", ] [dependency-groups] dev = [ diff --git a/tests/conftest.py b/tests/conftest.py index 01bf2a7..8db64dc 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -12,6 +12,7 @@ in the database to test authenticated endpoints. from __future__ import annotations import asyncio +import json import os import secrets import tempfile @@ -24,6 +25,20 @@ import pytest import pytest_asyncio from kanta import Kanta +# Keep runtime initialization invariant aligned with production: +# db.lifecycle requires PASKIA_CONFIG at import time. +os.environ.setdefault( + "PASKIA_CONFIG", + json.dumps( + { + "config": {"rp_id": "localhost", "rp_name": "localhost"}, + "site_url": "http://localhost:4401", + "site_path": "/auth/", + "save": False, + } + ), +) + import paskia.db.operations as ops_db from paskia import globals as paskia_globals from paskia.authsession import reset_expires @@ -34,13 +49,12 @@ from paskia.db import ( Permission, Role, User, - bootstrap, create_credential, create_reset_token, create_role, create_user, ) -from paskia.db.migrations import MigrationCtx +from paskia.db.bootstrap import bootstrap from paskia.db.operations import DB from paskia.db.structs import Session from paskia.fastapi.mainapp import app @@ -61,7 +75,7 @@ def event_loop(): async def test_db() -> AsyncGenerator[DB, None]: """Create a temporary JSONL database for testing using kanta. - Uses bootstrap() to properly initialize the database with: + Uses a kanta bootstrap callback to properly initialize the database with: - auth:admin and auth:org:admin permissions - A default organization with Administration role - An admin user with the Administration role @@ -72,34 +86,46 @@ async def test_db() -> AsyncGenerator[DB, None]: f.name, db, migrations="paskia.db.migrations", - migration_ctx=MigrationCtx(rp_id="test.example.com"), ) + kanta.ctx.rp_id = "test.example.com" + + # Register bootstrap callback so kanta seeds the empty DB during open() + @kanta.bootstrap(action="bootstrap") + def bootstrap_test_db(data: DB) -> None: + bootstrap( + data, + org_name="Test Organization", + admin_name="Test Admin", + ) + await kanta.open() - ops_db._store = kanta ops_db._db = db ops_db._db._store = kanta - # Bootstrap creates the initial permissions, org, role, and admin user - bootstrap( - org_name="Test Organization", - admin_name="Test Admin", - ) yield ops_db._db await kanta.close() ops_db._db = None - ops_db._store = None @pytest_asyncio.fixture(scope="function") async def passkey_instance() -> Passkey: - """Initialize a passkey instance for testing.""" + """Override the module-level passkey instance for testing.""" pk = Passkey( rp_id="localhost", rp_name="Test RP", origins=["http://localhost:4401"], ) - paskia_globals.passkey._instance = pk + original = { + "rp_id": paskia_globals.passkey.rp_id, + "rp_name": paskia_globals.passkey.rp_name, + "allowed_origins": paskia_globals.passkey.allowed_origins, + } + paskia_globals.passkey.rp_id = pk.rp_id + paskia_globals.passkey.rp_name = pk.rp_name + paskia_globals.passkey.allowed_origins = pk.allowed_origins yield pk - paskia_globals.passkey._instance = None + paskia_globals.passkey.rp_id = original["rp_id"] + paskia_globals.passkey.rp_name = original["rp_name"] + paskia_globals.passkey.allowed_origins = original["allowed_origins"] @pytest_asyncio.fixture(scope="function") @@ -286,7 +312,10 @@ def create_test_session( ) if session.key in ops_db._db.sessions: raise ValueError("Session already exists") - with ops_db._db.transaction("create_test_session"): + store = ops_db._db._store + if store is None: + raise RuntimeError("Test DB store is not initialized") + with store.transaction("create_test_session"): session.store(now) return session.key, token diff --git a/tests/test_api.py b/tests/test_api.py index 7f6a819..9f3af3c 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -426,7 +426,10 @@ class TestOidcUserInfoEndpoint: redirect_uris=["https://client.example/callback"], client_secret="topsecret", ) - with test_db.transaction("create_test_oidc_client"): + store = test_db._store + if store is None: + raise RuntimeError("Test DB store is not initialized") + with store.transaction("create_test_oidc_client"): test_db.oidc.clients[oidc_client.uuid] = oidc_client access_token = oidjwt.create_access_token( diff --git a/tests/test_cli.py b/tests/test_cli.py new file mode 100644 index 0000000..ac8ac93 --- /dev/null +++ b/tests/test_cli.py @@ -0,0 +1,154 @@ +"""Tests for the CLI entry point in paskia/__main__.py.""" + +from __future__ import annotations + +import asyncio +import os +import subprocess +import sys +import tempfile +from pathlib import Path +from typing import Any + +import pytest +from kanta import Kanta + +from paskia.__main__ import main +from paskia.db.structs import DB, Config +from paskia.util.runtime import clear_config_cache +from paskia.util.runtime import config as runtime_config + + +@pytest.fixture +def cli_run(monkeypatch): + """Run the CLI main() with the given args and return the RuntimeConfig.""" + + def _run(*args: str, db_root: str | None = None) -> Any: + env = os.environ.copy() + if db_root is not None: + env["PASKIA_DB"] = db_root + monkeypatch.setattr(os, "environ", env) + + monkeypatch.setattr(sys, "argv", ["paskia", *args]) + monkeypatch.setattr("fastapi_vue.server.run", lambda *_args, **_kw: None) + monkeypatch.setattr( + "paskia.util.startupbox.print_startup_config", lambda _rt: None + ) + monkeypatch.setattr("logging.basicConfig", lambda **_kw: None) + + clear_config_cache() + main() + runtime = runtime_config() + clear_config_cache() + return runtime + + return _run + + +async def _write_config(db_path: Path, config: Config) -> None: + """Write a Config into a JSONL database file using Kanta. + + The initial root uses a different rp_id so the stored diff includes the + target rp_id (required because Config omits defaults when diffing). + """ + kanta = Kanta( + str(db_path), + DB(config=Config(rp_id="uninitialized.invalid")), + migrations="paskia.db.migrations", + ) + kanta.ctx.rp_id = config.rp_id + await kanta.open() + with kanta.transaction("test:write_config"): + kanta.data.config = config + await kanta.close() + + +def write_config(db_path: Path, config: Config) -> None: + """Synchronous wrapper for _write_config.""" + asyncio.run(_write_config(db_path, config)) + + +def test_cli_defaults(cli_run): + with tempfile.TemporaryDirectory() as tmp: + runtime = cli_run("--rp-id", "localhost", db_root=tmp) + + assert runtime.config.rp_id == "localhost" + assert runtime.config.rp_name is None + assert runtime.config.auth_host is None + assert runtime.config.origins is None + assert runtime.site_url == "http://localhost:4401" + assert runtime.site_path == "/auth/" + assert runtime.save is False + + +def test_cli_explicit_options(cli_run): + runtime = cli_run( + "--rp-id", + "example.com", + "--rp-name", + "Example Corp", + "--auth-host", + "auth.example.com", + "--origin", + "https://app.example.com", + ) + + assert runtime.config.rp_id == "example.com" + assert runtime.config.rp_name == "Example Corp" + assert runtime.config.auth_host == "https://auth.example.com" + assert runtime.config.origins == [ + "https://auth.example.com", + "https://app.example.com", + ] + assert runtime.site_url == "https://auth.example.com" + assert runtime.site_path == "/" + + +def test_cli_loads_stored_config(cli_run): + with tempfile.TemporaryDirectory() as tmp: + db_path = Path(tmp) / "main.db" + write_config( + db_path, + Config( + rp_id="example.com", + rp_name="Stored Name", + origins=["https://stored.example.com"], + ), + ) + runtime = cli_run("--rp-id", "example.com", db_root=tmp) + + assert runtime.config.rp_name == "Stored Name" + assert runtime.config.origins == ["https://stored.example.com"] + assert runtime.site_url == "https://stored.example.com" + + +def test_cli_overrides_stored_config(cli_run): + with tempfile.TemporaryDirectory() as tmp: + db_path = Path(tmp) / "main.db" + write_config(db_path, Config(rp_id="example.com", rp_name="Stored Name")) + runtime = cli_run( + "--rp-id", "example.com", "--rp-name", "Overridden", db_root=tmp + ) + + assert runtime.config.rp_name == "Overridden" + + +def test_cli_save_flag(cli_run): + runtime = cli_run("--save") + assert runtime.save is True + + +def test_cli_invalid_auth_host(cli_run): + with pytest.raises(SystemExit): + cli_run("--rp-id", "example.com", "--auth-host", "notsub.example.org") + + +def test_cli_help(): + result = subprocess.run( + [sys.executable, "-m", "paskia", "--help"], + capture_output=True, + text=True, + check=False, + ) + assert result.returncode == 0 + assert "Paskia authentication server" in result.stdout