diff --git a/paskia/db/jsonl.py b/paskia/db/jsonl.py index 2dd2aa3..e218b59 100644 --- a/paskia/db/jsonl.py +++ b/paskia/db/jsonl.py @@ -19,7 +19,7 @@ import aiofiles import jsondiff import msgspec -from paskia.db.migrations import apply_migrations +from paskia.db.migrations import DBVER, apply_all_migrations from paskia.db.structs import DB, SessionContext _logger = logging.getLogger(__name__) @@ -33,6 +33,7 @@ class _ChangeRecord(msgspec.Struct, omit_defaults=True): ts: datetime a: str # action - describes the operation (e.g., "migrate", "login", "create_user") + v: int # schema version after this change u: str | None = None # user UUID who performed the action (None for system) diff: dict = {} @@ -41,43 +42,6 @@ class _ChangeRecord(msgspec.Struct, omit_defaults=True): _change_encoder = msgspec.json.Encoder() -async def load_jsonl(db_path: Path) -> dict: - """Load data from disk by applying change log. - - Replays all changes from JSONL file using plain dicts (to handle - schema evolution). - - Args: - db_path: Path to the JSONL database file - - Returns: - The final state after applying all changes - - Raises: - ValueError: If file doesn't exist or cannot be loaded - """ - if not db_path.exists(): - raise ValueError(f"Database file not found: {db_path}") - data_dict: dict = {} - try: - # Read entire file at once and split into lines - async with aiofiles.open(db_path, "rb") as f: - content = await f.read() - for line_num, line in enumerate(content.split(b"\n"), 1): - line = line.strip() - if not line: - continue - try: - change = msgspec.json.decode(line) - # Apply the diff to current state (marshal=True for $-prefixed keys) - data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True) - 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}") - return data_dict - - def compute_diff(previous: dict, current: dict) -> dict | None: """Compute JSON diff between two states. @@ -93,12 +57,13 @@ def compute_diff(previous: dict, current: dict) -> dict | None: def create_change_record( - action: str, diff: dict, user: str | None = None + action: str, version: int, diff: dict, user: str | None = None ) -> _ChangeRecord: """Create a change record for persistence.""" return _ChangeRecord( ts=datetime.now(UTC), a=action, + v=version, u=user, diff=diff, ) @@ -166,89 +131,93 @@ class JsonlStore: self._current_user: str | None = None self._in_transaction: bool = False 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: """Load data from JSONL change log.""" if db_path is not None: self.db_path = Path(db_path) + if not self.db_path.exists(): + return + + # Replay change log to reconstruct state + data_dict: dict = {} try: - data_dict = await load_jsonl(self.db_path) - if data_dict: - # Preserve original state before migrations (deep copy for nested dicts) - original_dict = copy.deepcopy(data_dict) - - # Apply schema migrations (modifies data_dict in place) - migrated = apply_migrations(data_dict) - - decoder = msgspec.json.Decoder(DB) - self.db = decoder.decode(msgspec.json.encode(data_dict)) - self.db._store = self - - # Update previous state to migrated data FIRST (to avoid transaction hardening reset) - self._previous_builtins = data_dict - - # Persist version migration by manually computing and queueing the diff - if migrated: - new_version = data_dict.get("v", 1) - diff = compute_diff(original_dict, data_dict) - if diff: - action = f"migrate:v{new_version}" - self._pending_changes.append( - create_change_record(action, diff, user=None) - ) - diff_json = json.dumps(diff, default=str) - print(f"{action}: {diff_json}", file=sys.stderr) - await self.flush() - # Update original_dict to reflect persisted state - original_dict = copy.deepcopy(data_dict) - - # Normalize via msgspec round-trip (handles omit_defaults etc.) - # This ensures _previous_builtins matches what msgspec would produce - normalized_dict = msgspec.to_builtins(self.db) - diff = compute_diff(original_dict, normalized_dict) - if diff: - action = "migrate:msgspec" - self._pending_changes.append( - create_change_record(action, diff, user=None) - ) - diff_json = json.dumps(diff, default=str) - print(f"{action}: {diff_json}", file=sys.stderr) - await self.flush() - # Update _previous_builtins to normalized state - self._previous_builtins = normalized_dict - else: - # No data loaded - _previous_builtins stays as empty dict - pass - except ValueError: - if self.db_path.exists(): - raise - - def _queue_change(self) -> None: - current = msgspec.to_builtins(self.db) - diff = compute_diff(self._previous_builtins, current) - if diff: - self._pending_changes.append( - create_change_record(self._current_action, diff, self._current_user) - ) - self._previous_builtins = current - # Log the change with user display name if available - user_display = None - if self._current_user: + async with aiofiles.open(self.db_path, "rb") as f: + content = await f.read() + for line_num, line in enumerate(content.split(b"\n"), 1): + line = line.strip() + if not line: + continue try: - user_uuid = UUID(self._current_user) - if user_uuid in self.db.users: - user_display = self.db.users[user_uuid].display_name - except (ValueError, KeyError): - user_display = self._current_user + change = msgspec.json.decode(line) + data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True) + 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}") - diff_json = json.dumps(diff, default=str) - if user_display: - print( - f"{self._current_action} by {user_display}: {diff_json}", - file=sys.stderr, - ) - else: - print(f"{self._current_action}: {diff_json}", file=sys.stderr) + if not data_dict: + return + + # Set previous state for diffing (will be updated by _queue_change) + self._previous_builtins = copy.deepcopy(data_dict) + + # Callback to persist each migration + async def persist_migration( + action: str, new_version: int, current: dict + ) -> None: + self._current_version = new_version + 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) + + # Decode to msgspec struct + decoder = msgspec.json.Decoder(DB) + self.db = decoder.decode(msgspec.json.encode(data_dict)) + self.db._store = self + + # Normalize via msgspec round-trip (handles omit_defaults etc.) + # This ensures _previous_builtins matches what msgspec would produce + normalized_dict = msgspec.to_builtins(self.db) + await persist_migration( + "migrate:msgspec", self._current_version, normalized_dict + ) + + def _queue_change( + self, action: str, version: int, current: dict, user: str | None = None + ) -> None: + """Queue a change record and log it. + + Args: + action: The action name for the change record + version: The schema version for the change record + current: The current state as a plain dict + user: Optional user UUID who performed the action + """ + diff = compute_diff(self._previous_builtins, current) + if not diff: + return + self._pending_changes.append(create_change_record(action, version, diff, user)) + self._previous_builtins = copy.deepcopy(current) + + # Log the change with user display name if available + user_display = None + if user: + try: + user_uuid = UUID(user) + if user_uuid in self.db.users: + user_display = self.db.users[user_uuid].display_name + except (ValueError, KeyError): + user_display = user + + diff_json = json.dumps(diff, default=str) + if user_display: + print(f"{action} by {user_display}: {diff_json}", file=sys.stderr) + else: + print(f"{action}: {diff_json}", file=sys.stderr) @contextmanager def transaction( @@ -295,7 +264,10 @@ class JsonlStore: try: yield - self._queue_change() + current = msgspec.to_builtins(self.db) + self._queue_change( + self._current_action, self._current_version, current, self._current_user + ) except Exception: # Rollback on error: restore from snapshot _logger.warning("Transaction '%s' failed, rolling back changes", action) diff --git a/paskia/db/migrations.py b/paskia/db/migrations.py index 6eddec6..f5689da 100644 --- a/paskia/db/migrations.py +++ b/paskia/db/migrations.py @@ -5,25 +5,29 @@ Migrations are applied during database load based on the version field. Each migration should be idempotent and only run when needed. """ +from collections.abc import Awaitable, Callable -def apply_migrations(data_dict: dict) -> bool: - """Apply any pending schema migrations to the database dictionary. - Args: - data_dict: The raw database dictionary loaded from JSONL +def migrate_v1(d: dict) -> None: + """Remove Org.created_at fields.""" + for org_data in d["orgs"].values(): + org_data.pop("created_at", None) - Returns: - True if any migrations were applied, False otherwise - """ - db_version = data_dict.get("v", 0) - migrated = False - if db_version == 0: - # Migration v0 -> v1: Remove created_at from orgs (field removed from schema) - if "orgs" in data_dict: - for org_data in data_dict["orgs"].values(): - org_data.pop("created_at", None) - data_dict["v"] = 1 - migrated = True +migrations = sorted( + [f for n, f in globals().items() if n.startswith("migrate_v")], + key=lambda f: int(f.__name__.removeprefix("migrate_v")), +) - return migrated +DBVER = len(migrations) # Used by bootstrap and migrate:sql to set initial version + + +async def apply_all_migrations( + data_dict: dict, + current_version: int, + persist: Callable[[str, int, dict], Awaitable[None]], +) -> None: + while current_version < DBVER: + migrations[current_version](data_dict) + current_version += 1 + await persist(f"migrate:v{current_version}", current_version, data_dict) diff --git a/paskia/db/structs.py b/paskia/db/structs.py index 82eb2aa..1831ad2 100644 --- a/paskia/db/structs.py +++ b/paskia/db/structs.py @@ -377,7 +377,6 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False): credentials: dict[UUID, Credential] = {} sessions: dict[str, Session] = {} reset_tokens: dict[bytes, ResetToken] = {} - v: int = 0 def __post_init__(self): # Store reference for persistence (not serialized)