Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0cf551cb28 | ||
|
|
632230d05c |
@@ -47,6 +47,7 @@ from paskia.db.operations import (
|
||||
delete_session,
|
||||
delete_sessions_for_user,
|
||||
delete_user,
|
||||
get_config,
|
||||
get_organization_users,
|
||||
get_reset_token,
|
||||
get_user_credential_ids,
|
||||
@@ -55,6 +56,7 @@ from paskia.db.operations import (
|
||||
login,
|
||||
remove_permission_from_org,
|
||||
remove_permission_from_role,
|
||||
set_config,
|
||||
set_session_host,
|
||||
update_credential_sign_count,
|
||||
update_org_name,
|
||||
@@ -110,6 +112,7 @@ __all__ = [
|
||||
"build_session",
|
||||
"build_user",
|
||||
# Read ops
|
||||
"get_config",
|
||||
"get_organization_users",
|
||||
"get_reset_token",
|
||||
"get_user_credential_ids",
|
||||
@@ -138,6 +141,7 @@ __all__ = [
|
||||
"login",
|
||||
"remove_permission_from_org",
|
||||
"remove_permission_from_role",
|
||||
"set_config",
|
||||
"set_session_host",
|
||||
"update_credential_sign_count",
|
||||
"update_org_name",
|
||||
|
||||
+33
-22
@@ -4,6 +4,8 @@ JSONL persistence layer for the database.
|
||||
|
||||
import copy
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
from collections import deque
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
@@ -69,22 +71,25 @@ def create_change_record(
|
||||
# Actions that are allowed to create a new database file
|
||||
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap"})
|
||||
|
||||
# Flag to prevent duplicate error messages on fatal flush failure
|
||||
_flush_failed = False
|
||||
|
||||
|
||||
async def flush_changes(
|
||||
db_path: Path,
|
||||
pending_changes: deque[_ChangeRecord],
|
||||
) -> bool:
|
||||
) -> None:
|
||||
"""Write all pending changes to disk.
|
||||
|
||||
Args:
|
||||
db_path: Path to the JSONL database file
|
||||
pending_changes: Queue of pending change records (will be cleared on success)
|
||||
|
||||
Returns:
|
||||
True if flush succeeded, False otherwise
|
||||
On failure, logs an error and sends SIGTERM to trigger graceful shutdown.
|
||||
"""
|
||||
if not pending_changes:
|
||||
return True
|
||||
global _flush_failed
|
||||
if _flush_failed or not pending_changes:
|
||||
return
|
||||
|
||||
if not db_path.exists():
|
||||
first_action = pending_changes[0].a
|
||||
@@ -94,26 +99,25 @@ async def flush_changes(
|
||||
"only bootstrap can create a new database",
|
||||
first_action,
|
||||
)
|
||||
pending_changes.clear()
|
||||
return False
|
||||
_flush_failed = True
|
||||
os.kill(os.getpid(), signal.SIGTERM)
|
||||
return
|
||||
|
||||
changes_to_write = list(pending_changes)
|
||||
pending_changes.clear()
|
||||
|
||||
try:
|
||||
lines = [_change_encoder.encode(change) for change in changes_to_write]
|
||||
if not lines:
|
||||
return True
|
||||
pending_changes.clear()
|
||||
return
|
||||
|
||||
async with aiofiles.open(db_path, "ab") as f:
|
||||
await f.write(b"\n".join(lines) + b"\n")
|
||||
return True
|
||||
except OSError:
|
||||
_logger.exception("Failed to flush database changes")
|
||||
# Re-queue the changes on failure
|
||||
for change in reversed(changes_to_write):
|
||||
pending_changes.appendleft(change)
|
||||
return False
|
||||
pending_changes.clear()
|
||||
except OSError as e:
|
||||
_logger.error("Failed to flush database: %s", e)
|
||||
_flush_failed = True
|
||||
os.kill(os.getpid(), signal.SIGTERM)
|
||||
|
||||
|
||||
class JsonlStore:
|
||||
@@ -130,10 +134,13 @@ class JsonlStore:
|
||||
self._transaction_snapshot: dict[str, Any] | None = None
|
||||
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."""
|
||||
if db_path is not None:
|
||||
self.db_path = Path(db_path)
|
||||
self._rp_id = rp_id
|
||||
if not self.db_path.exists():
|
||||
return
|
||||
|
||||
@@ -152,8 +159,10 @@ class JsonlStore:
|
||||
self._current_version = change.get("v", 0)
|
||||
except Exception as e:
|
||||
raise ValueError(f"Error parsing line {line_num}: {e}")
|
||||
except (OSError, ValueError, msgspec.DecodeError) as e:
|
||||
raise ValueError(f"Failed to load database: {e}")
|
||||
except OSError as 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:
|
||||
return
|
||||
@@ -169,7 +178,9 @@ class JsonlStore:
|
||||
self._queue_change(action, new_version, current)
|
||||
|
||||
# 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
|
||||
decoder = msgspec.json.Decoder(DB)
|
||||
@@ -277,6 +288,6 @@ class JsonlStore:
|
||||
self._in_transaction = False
|
||||
self._transaction_snapshot = None
|
||||
|
||||
async def flush(self) -> bool:
|
||||
async def flush(self) -> None:
|
||||
"""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
@@ -8,12 +8,18 @@ Each migration should be idempotent and only run when needed.
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
|
||||
def migrate_v1(d: dict) -> None:
|
||||
def migrate_v1(d: dict, **kwargs) -> None:
|
||||
"""Remove Org.created_at fields."""
|
||||
for org_data in d["orgs"].values():
|
||||
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(
|
||||
[f for n, f in globals().items() if n.startswith("migrate_v")],
|
||||
key=lambda f: int(f.__name__.removeprefix("migrate_v")),
|
||||
@@ -26,8 +32,10 @@ async def apply_all_migrations(
|
||||
data_dict: dict,
|
||||
current_version: int,
|
||||
persist: Callable[[str, int, dict], Awaitable[None]],
|
||||
*,
|
||||
rp_id: str = "localhost",
|
||||
) -> None:
|
||||
while current_version < DBVER:
|
||||
migrations[current_version](data_dict)
|
||||
migrations[current_version](data_dict, rp_id=rp_id)
|
||||
current_version += 1
|
||||
await persist(f"migrate:v{current_version}", current_version, data_dict)
|
||||
|
||||
@@ -21,6 +21,7 @@ from paskia.db.jsonl import (
|
||||
)
|
||||
from paskia.db.structs import (
|
||||
DB,
|
||||
Config,
|
||||
Credential,
|
||||
Org,
|
||||
Permission,
|
||||
@@ -49,7 +50,7 @@ async def init(rp_id: str = "localhost", *args, **kwargs):
|
||||
return
|
||||
default_path = f"{rp_id}.paskiadb"
|
||||
db_path = os.environ.get("PASKIA_DB", default_path)
|
||||
await _store.load(db_path)
|
||||
await _store.load(db_path, rp_id=rp_id)
|
||||
_db = _store.db
|
||||
_initialized = True
|
||||
|
||||
@@ -835,5 +836,5 @@ def get_config() -> Config:
|
||||
|
||||
async def set_config(config: Config) -> None:
|
||||
"""Update the stored configuration."""
|
||||
async with _db.transaction("update_config"):
|
||||
with _db.transaction("update_config"):
|
||||
_db.config = config
|
||||
|
||||
@@ -7,6 +7,7 @@ from uuid import UUID
|
||||
|
||||
import msgspec
|
||||
import uuid7
|
||||
from msgspec import field
|
||||
|
||||
from paskia import db
|
||||
from paskia.util.hostutil import normalize_host
|
||||
@@ -397,10 +398,10 @@ class SessionContext(msgspec.Struct):
|
||||
permissions: list[Permission] = []
|
||||
|
||||
|
||||
class Config(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
class Config(msgspec.Struct, frozen=True, dict=True, omit_defaults=True):
|
||||
"""Stored configuration for the instance."""
|
||||
|
||||
rp_id: str | None = None
|
||||
rp_id: str
|
||||
rp_name: str | None = None
|
||||
origins: list[str] | None = None
|
||||
auth_host: str | None = None
|
||||
@@ -422,7 +423,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
credentials: dict[UUID, Credential] = {}
|
||||
sessions: dict[str, Session] = {}
|
||||
reset_tokens: dict[bytes, ResetToken] = {}
|
||||
config: Config = Config()
|
||||
config: Config = field(default_factory=lambda: Config(rp_id="localhost"))
|
||||
|
||||
def __post_init__(self):
|
||||
# Store reference for persistence (not serialized)
|
||||
|
||||
@@ -6,7 +6,8 @@ import os
|
||||
from urllib.parse import urlparse
|
||||
|
||||
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 paskia import globals as _globals
|
||||
@@ -236,7 +237,7 @@ def main():
|
||||
for ep in endpoints:
|
||||
tg.create_task(
|
||||
Server(
|
||||
Config(app="paskia.fastapi:app", **run_kwargs, **ep)
|
||||
UvicornConfig(app="paskia.fastapi:app", **run_kwargs, **ep)
|
||||
).serve()
|
||||
)
|
||||
elif DEVMODE:
|
||||
@@ -245,7 +246,7 @@ def main():
|
||||
uvicorn_run("paskia.fastapi:app", **run_kwargs, **ep)
|
||||
else:
|
||||
server = Server(
|
||||
Config(app="paskia.fastapi:app", **run_kwargs, **endpoints[0])
|
||||
UvicornConfig(app="paskia.fastapi:app", **run_kwargs, **endpoints[0])
|
||||
)
|
||||
await server.serve()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user