Compare commits

...
5 Commits
10 changed files with 153 additions and 42 deletions
+1 -1
View File
@@ -5,7 +5,7 @@ dist/
*.lock *.lock
package-lock.json package-lock.json
paskia.sqlite paskia.sqlite
paskia.jsonl *.paskiadb
/paskia/frontend-build /paskia/frontend-build
/paskia/_version.py /paskia/_version.py
coverage-html/ coverage-html/
+9 -6
View File
@@ -51,7 +51,7 @@ uv tool install paskia
## Configuration ## Configuration
All configuration is passed by CLI arguments, of which there are just a few. You will need to specify your main domain to which all passkeys will be tied as rp-id. Use your main domain even if Paskia is not running there. All other options are optional.
```text ```text
paskia [options] paskia [options]
@@ -61,9 +61,12 @@ paskia [options]
|--------|-------------|---------| |--------|-------------|---------|
| -l, --listen *endpoint* | Listen address: *host*:*port*, :*port* (all interfaces), or */path.sock* | **localhost:4401** | | -l, --listen *endpoint* | Listen address: *host*:*port*, :*port* (all interfaces), or */path.sock* | **localhost:4401** |
| --rp-id *domain* | Main/top domain for passkeys | **localhost** | | --rp-id *domain* | Main/top domain for passkeys | **localhost** |
| --rp-name *"text"* | Name shown during passkey registration | Same as rp-id | | --rp-name *"text"* | Branding name for the entire system (passkey auth, login dialog). | Same as rp-id |
| --origin *url* | Restrict allowed origins for WebSocket auth (repeatable) | All under rp-id | | --origin *url* | Only sites listed can login (repeatable) | rp-id and all subdomains |
| --auth-host *url* | Dedicated authentication site, e.g. **auth.example.com** | Use **/auth/** path on each site | | --auth-host *url* | Dedicated authentication site, e.g. **auth.example.com** | Use **/auth/** path on each site |
| --save | Save current options to database | (only --rp-id required on further invocations) |
To clear a stored setting, pass an empty value like `--auth-host=`. The database is stored in `{rp-id}.paskiadb` in current directory. This can be overridden by environment `PASKIA_DB` if needed.
## Tutorial: From Local Testing to Production ## Tutorial: From Local Testing to Production
@@ -84,10 +87,10 @@ This starts the server on [localhost:4401](http://localhost:4401) with passkeys
For a real deployment, configure Paskia with your domain name (rp-id). This enables SSO setup for that domain and any subdomains. For a real deployment, configure Paskia with your domain name (rp-id). This enables SSO setup for that domain and any subdomains.
```fish ```fish
paskia --rp-id example.com --rp-name "Example Corp" paskia --rp-id example.com --rp-name "Example Corp" --save
``` ```
This binds passkeys to `*.example.com`. The `--rp-name` is shown to users during passkey registration. This binds passkeys to `*.example.com`. The `--rp-name` is shown to users during passkey registration. The `--save` option stores these settings in the database, so future runs only need `paskia --rp-id example.com`.
### Step 3: Set Up Caddy ### Step 3: Set Up Caddy
@@ -187,7 +190,7 @@ Description=Paskia Authentication Server
Type=simple Type=simple
User=paskia User=paskia
WorkingDirectory=/srv/paskia WorkingDirectory=/srv/paskia
ExecStart=uvx paskia --rp-id example.com --rp-name "Example Corp" ExecStart=uvx paskia --rp-id=example.com
[Install] [Install]
WantedBy=multi-user.target WantedBy=multi-user.target
+4
View File
@@ -47,6 +47,7 @@ from paskia.db.operations import (
delete_session, delete_session,
delete_sessions_for_user, delete_sessions_for_user,
delete_user, delete_user,
get_config,
get_organization_users, get_organization_users,
get_reset_token, get_reset_token,
get_user_credential_ids, get_user_credential_ids,
@@ -55,6 +56,7 @@ from paskia.db.operations import (
login, login,
remove_permission_from_org, remove_permission_from_org,
remove_permission_from_role, remove_permission_from_role,
set_config,
set_session_host, set_session_host,
update_credential_sign_count, update_credential_sign_count,
update_org_name, update_org_name,
@@ -110,6 +112,7 @@ __all__ = [
"build_session", "build_session",
"build_user", "build_user",
# Read ops # Read ops
"get_config",
"get_organization_users", "get_organization_users",
"get_reset_token", "get_reset_token",
"get_user_credential_ids", "get_user_credential_ids",
@@ -138,6 +141,7 @@ __all__ = [
"login", "login",
"remove_permission_from_org", "remove_permission_from_org",
"remove_permission_from_role", "remove_permission_from_role",
"set_config",
"set_session_host", "set_session_host",
"update_credential_sign_count", "update_credential_sign_count",
"update_org_name", "update_org_name",
+33 -22
View File
@@ -4,6 +4,8 @@ JSONL persistence layer for the database.
import copy import copy
import logging import logging
import os
import signal
from collections import deque from collections import deque
from contextlib import contextmanager from contextlib import contextmanager
from datetime import UTC, datetime from datetime import UTC, datetime
@@ -69,22 +71,25 @@ def create_change_record(
# Actions that are allowed to create a new database file # Actions that are allowed to create a new database file
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap"}) _BOOTSTRAP_ACTIONS = frozenset({"bootstrap"})
# Flag to prevent duplicate error messages on fatal flush failure
_flush_failed = False
async def flush_changes( async def flush_changes(
db_path: Path, db_path: Path,
pending_changes: deque[_ChangeRecord], pending_changes: deque[_ChangeRecord],
) -> bool: ) -> None:
"""Write all pending changes to disk. """Write all pending changes to disk.
Args: Args:
db_path: Path to the JSONL database file db_path: Path to the JSONL database file
pending_changes: Queue of pending change records (will be cleared on success) pending_changes: Queue of pending change records (will be cleared on success)
Returns: On failure, logs an error and sends SIGTERM to trigger graceful shutdown.
True if flush succeeded, False otherwise
""" """
if not pending_changes: global _flush_failed
return True if _flush_failed or not pending_changes:
return
if not db_path.exists(): if not db_path.exists():
first_action = pending_changes[0].a first_action = pending_changes[0].a
@@ -94,26 +99,25 @@ async def flush_changes(
"only bootstrap can create a new database", "only bootstrap can create a new database",
first_action, first_action,
) )
pending_changes.clear() _flush_failed = True
return False os.kill(os.getpid(), signal.SIGTERM)
return
changes_to_write = list(pending_changes) changes_to_write = list(pending_changes)
pending_changes.clear()
try: try:
lines = [_change_encoder.encode(change) for change in changes_to_write] lines = [_change_encoder.encode(change) for change in changes_to_write]
if not lines: if not lines:
return True pending_changes.clear()
return
async with aiofiles.open(db_path, "ab") as f: async with aiofiles.open(db_path, "ab") as f:
await f.write(b"\n".join(lines) + b"\n") await f.write(b"\n".join(lines) + b"\n")
return True pending_changes.clear()
except OSError: except OSError as e:
_logger.exception("Failed to flush database changes") _logger.error("Failed to flush database: %s", e)
# Re-queue the changes on failure _flush_failed = True
for change in reversed(changes_to_write): os.kill(os.getpid(), signal.SIGTERM)
pending_changes.appendleft(change)
return False
class JsonlStore: class JsonlStore:
@@ -130,10 +134,13 @@ class JsonlStore:
self._transaction_snapshot: dict[str, Any] | None = None self._transaction_snapshot: dict[str, Any] | None = None
self._current_version: int = DBVER # Schema version for new databases self._current_version: int = DBVER # Schema version for new databases
async def load(self, db_path: str | None = None) -> None: async def load(
self, db_path: str | None = None, *, rp_id: str = "localhost"
) -> None:
"""Load data from JSONL change log.""" """Load data from JSONL change log."""
if db_path is not None: if db_path is not None:
self.db_path = Path(db_path) self.db_path = Path(db_path)
self._rp_id = rp_id
if not self.db_path.exists(): if not self.db_path.exists():
return return
@@ -152,8 +159,10 @@ class JsonlStore:
self._current_version = change.get("v", 0) self._current_version = change.get("v", 0)
except Exception as e: except Exception as e:
raise ValueError(f"Error parsing line {line_num}: {e}") raise ValueError(f"Error parsing line {line_num}: {e}")
except (OSError, ValueError, msgspec.DecodeError) as e: except OSError as e:
raise ValueError(f"Failed to load database: {e}") raise SystemExit(f"Failed to load database: {e}")
except (ValueError, msgspec.DecodeError) as e:
raise SystemExit(f"Failed to load database: {e}")
if not data_dict: if not data_dict:
return return
@@ -169,7 +178,9 @@ class JsonlStore:
self._queue_change(action, new_version, current) self._queue_change(action, new_version, current)
# Apply schema migrations one at a time # Apply schema migrations one at a time
await apply_all_migrations(data_dict, self._current_version, persist_migration) await apply_all_migrations(
data_dict, self._current_version, persist_migration, rp_id=rp_id
)
# Decode to msgspec struct # Decode to msgspec struct
decoder = msgspec.json.Decoder(DB) decoder = msgspec.json.Decoder(DB)
@@ -277,6 +288,6 @@ class JsonlStore:
self._in_transaction = False self._in_transaction = False
self._transaction_snapshot = None self._transaction_snapshot = None
async def flush(self) -> bool: async def flush(self) -> None:
"""Write all pending changes to disk.""" """Write all pending changes to disk."""
return await flush_changes(self.db_path, self._pending_changes) await flush_changes(self.db_path, self._pending_changes)
+10 -2
View File
@@ -8,12 +8,18 @@ Each migration should be idempotent and only run when needed.
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
def migrate_v1(d: dict) -> None: def migrate_v1(d: dict, **kwargs) -> None:
"""Remove Org.created_at fields.""" """Remove Org.created_at fields."""
for org_data in d["orgs"].values(): for org_data in d["orgs"].values():
org_data.pop("created_at", None) org_data.pop("created_at", None)
def migrate_v2(d: dict, *, rp_id: str = "localhost") -> None:
"""Add config field if missing."""
if "config" not in d:
d["config"] = {"rp_id": rp_id}
migrations = sorted( migrations = sorted(
[f for n, f in globals().items() if n.startswith("migrate_v")], [f for n, f in globals().items() if n.startswith("migrate_v")],
key=lambda f: int(f.__name__.removeprefix("migrate_v")), key=lambda f: int(f.__name__.removeprefix("migrate_v")),
@@ -26,8 +32,10 @@ async def apply_all_migrations(
data_dict: dict, data_dict: dict,
current_version: int, current_version: int,
persist: Callable[[str, int, dict], Awaitable[None]], persist: Callable[[str, int, dict], Awaitable[None]],
*,
rp_id: str = "localhost",
) -> None: ) -> None:
while current_version < DBVER: while current_version < DBVER:
migrations[current_version](data_dict) migrations[current_version](data_dict, rp_id=rp_id)
current_version += 1 current_version += 1
await persist(f"migrate:v{current_version}", current_version, data_dict) await persist(f"migrate:v{current_version}", current_version, data_dict)
+21 -6
View File
@@ -17,11 +17,11 @@ import uuid7
from paskia.config import SESSION_LIFETIME from paskia.config import SESSION_LIFETIME
from paskia.db.jsonl import ( from paskia.db.jsonl import (
DB_PATH_DEFAULT,
JsonlStore, JsonlStore,
) )
from paskia.db.structs import ( from paskia.db.structs import (
DB, DB,
Config,
Credential, Credential,
Org, Org,
Permission, Permission,
@@ -42,16 +42,15 @@ _db._store = _store
_initialized = False _initialized = False
async def init(*args, **kwargs): async def init(rp_id: str = "localhost", *args, **kwargs):
"""Load database from JSONL file.""" """Load database from JSONL file."""
global _db, _initialized global _db, _initialized
if _initialized: if _initialized:
_logger.debug("Database already initialized, skipping reload") _logger.debug("Database already initialized, skipping reload")
return return
db_path = os.environ.get("PASKIA_DB", DB_PATH_DEFAULT) default_path = f"{rp_id}.paskiadb"
if db_path.startswith("json:"): db_path = os.environ.get("PASKIA_DB", default_path)
db_path = db_path[5:] await _store.load(db_path, rp_id=rp_id)
await _store.load(db_path)
_db = _store.db _db = _store.db
_initialized = True _initialized = True
@@ -823,3 +822,19 @@ def bootstrap(
_db.reset_tokens[reset_token.key] = reset_token _db.reset_tokens[reset_token.key] = reset_token
return reset_passphrase return reset_passphrase
# -------------------------------------------------------------------------
# Config operations
# -------------------------------------------------------------------------
def get_config() -> Config:
"""Get the stored configuration."""
return _db.config
async def set_config(config: Config) -> None:
"""Update the stored configuration."""
with _db.transaction("update_config"):
_db.config = config
+12
View File
@@ -7,6 +7,7 @@ from uuid import UUID
import msgspec import msgspec
import uuid7 import uuid7
from msgspec import field
from paskia import db from paskia import db
from paskia.util.hostutil import normalize_host from paskia.util.hostutil import normalize_host
@@ -397,6 +398,16 @@ class SessionContext(msgspec.Struct):
permissions: list[Permission] = [] permissions: list[Permission] = []
class Config(msgspec.Struct, frozen=True, dict=True, omit_defaults=True):
"""Stored configuration for the instance."""
rp_id: str
rp_name: str | None = None
origins: list[str] | None = None
auth_host: str | None = None
listen: str | None = None
# ------------------------------------------------------------------------- # -------------------------------------------------------------------------
# Database storage structure # Database storage structure
# ------------------------------------------------------------------------- # -------------------------------------------------------------------------
@@ -412,6 +423,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
credentials: dict[UUID, Credential] = {} credentials: dict[UUID, Credential] = {}
sessions: dict[str, Session] = {} sessions: dict[str, Session] = {}
reset_tokens: dict[bytes, ResetToken] = {} reset_tokens: dict[bytes, ResetToken] = {}
config: Config = field(default_factory=lambda: Config(rp_id="localhost"))
def __post_init__(self): def __post_init__(self):
# Store reference for persistence (not serialized) # Store reference for persistence (not serialized)
+44 -3
View File
@@ -6,13 +6,17 @@ import os
from urllib.parse import urlparse from urllib.parse import urlparse
from fastapi_vue.hostutil import parse_endpoint from fastapi_vue.hostutil import parse_endpoint
from uvicorn import Config, Server from uvicorn import Config as UvicornConfig
from uvicorn import Server
from uvicorn import run as uvicorn_run from uvicorn import run as uvicorn_run
from paskia import globals as _globals from paskia import globals as _globals
from paskia.bootstrap import bootstrap_if_needed from paskia.bootstrap import bootstrap_if_needed
from paskia.config import PaskiaConfig from paskia.config import PaskiaConfig
from paskia.db import get_config, set_config
from paskia.db import init as db_init
from paskia.db.background import flush from paskia.db.background import flush
from paskia.db.structs import Config
from paskia.util import startupbox from paskia.util import startupbox
from paskia.util.hostutil import normalize_origin from paskia.util.hostutil import normalize_origin
@@ -62,6 +66,11 @@ def add_common_options(p: argparse.ArgumentParser) -> None:
"--auth-host", "--auth-host",
help=("Dedicated authentication site (optionally with scheme/port)"), help=("Dedicated authentication site (optionally with scheme/port)"),
) )
p.add_argument(
"--save",
action="store_true",
help="Save the CLI options to database for future runs.",
)
def main(): def main():
@@ -88,6 +97,28 @@ def main():
args = parser.parse_args() args = parser.parse_args()
# Handle clearing options
if getattr(args, "auth_host", None) == "":
args.auth_host = None
if getattr(args, "rp_name", None) == "":
args.rp_name = None
if getattr(args, "listen", None) == "":
args.listen = None
# Init db and load stored config
asyncio.run(db_init(rp_id=args.rp_id))
stored_config = get_config()
# Apply defaults from stored config
if args.rp_name is None and stored_config.rp_name is not None:
args.rp_name = stored_config.rp_name
if args.origins is None and stored_config.origins is not None:
args.origins = stored_config.origins
if args.auth_host is None and stored_config.auth_host is not None:
args.auth_host = stored_config.auth_host
if args.listen is None and stored_config.listen is not None:
args.listen = stored_config.listen
# Parse endpoint using fastapi_vue.hostutil # Parse endpoint using fastapi_vue.hostutil
endpoints = parse_endpoint(args.listen, DEFAULT_PORT) endpoints = parse_endpoint(args.listen, DEFAULT_PORT)
@@ -168,6 +199,16 @@ def main():
startupbox.print_startup_config(config) startupbox.print_startup_config(config)
if args.save:
new_config = Config(
rp_id=args.rp_id,
rp_name=args.rp_name,
origins=args.origins,
auth_host=args.auth_host,
listen=args.listen,
)
asyncio.run(set_config(new_config))
run_kwargs: dict = { run_kwargs: dict = {
"log_level": "warning", # Suppress startup messages; we use custom logging "log_level": "warning", # Suppress startup messages; we use custom logging
"access_log": False, # We use custom AccessLogMiddleware instead "access_log": False, # We use custom AccessLogMiddleware instead
@@ -196,7 +237,7 @@ def main():
for ep in endpoints: for ep in endpoints:
tg.create_task( tg.create_task(
Server( Server(
Config(app="paskia.fastapi:app", **run_kwargs, **ep) UvicornConfig(app="paskia.fastapi:app", **run_kwargs, **ep)
).serve() ).serve()
) )
elif DEVMODE: elif DEVMODE:
@@ -205,7 +246,7 @@ def main():
uvicorn_run("paskia.fastapi:app", **run_kwargs, **ep) uvicorn_run("paskia.fastapi:app", **run_kwargs, **ep)
else: else:
server = Server( server = Server(
Config(app="paskia.fastapi:app", **run_kwargs, **endpoints[0]) UvicornConfig(app="paskia.fastapi:app", **run_kwargs, **endpoints[0])
) )
await server.serve() await server.serve()
+2 -2
View File
@@ -42,7 +42,7 @@ async def init(
Database configuration: Database configuration:
Set PASKIA_DB environment variable to specify the JSONL database file path. Set PASKIA_DB environment variable to specify the JSONL database file path.
Default: paskia.jsonl Default: {rp_id}.paskiadb
""" """
# Initialize passkey instance with provided parameters # Initialize passkey instance with provided parameters
@@ -53,7 +53,7 @@ async def init(
) )
# Initialize database # Initialize database
await db.init() await db.init(rp_id=rp_id)
# Initialize remote auth manager # Initialize remote auth manager
await remoteauth.init() await remoteauth.init()
+17
View File
@@ -8,6 +8,7 @@ This module provides a unified interface for WebAuthn operations including:
""" """
import json import json
import re
from urllib.parse import urlparse from urllib.parse import urlparse
from uuid import UUID from uuid import UUID
@@ -62,6 +63,7 @@ class Passkey:
ValueError: If any origin domain doesn't match or isn't a subdomain of rp_id. ValueError: If any origin domain doesn't match or isn't a subdomain of rp_id.
""" """
self.rp_id = rp_id self.rp_id = rp_id
self._validate_rp_id(rp_id)
self.rp_name = rp_name or rp_id self.rp_name = rp_name or rp_id
self.allowed_origins: set[str] | None = None self.allowed_origins: set[str] | None = None
if origins: if origins:
@@ -75,6 +77,21 @@ class Passkey:
COSEAlgorithmIdentifier.RSASSA_PKCS1_v1_5_SHA_256, COSEAlgorithmIdentifier.RSASSA_PKCS1_v1_5_SHA_256,
] ]
def _validate_rp_id(self, rp_id: str) -> None:
"""Validate that rp_id is a valid domain name."""
if not rp_id:
raise ValueError("rp_id cannot be empty")
# Allow localhost, or domain-like strings
if rp_id == "localhost":
return
# Regex for valid domain: letters, digits, hyphens, dots, but not starting/ending with hyphen, etc.
# Simplified: alphanumeric, dots, hyphens
if not re.match(
r"^[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?(\.[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?)*$",
rp_id,
):
raise ValueError(f"rp_id '{rp_id}' is not a valid domain name")
def _validate_origin(self, origin: str, rp_id: str) -> None: def _validate_origin(self, origin: str, rp_id: str) -> None:
"""Validate an origin URL against the rp_id.""" """Validate an origin URL against the rp_id."""
hostname = urlparse(origin).hostname hostname = urlparse(origin).hostname