DB cleanup continued: Made the working copy data public in DB class.

This commit is contained in:
Leo Vasanko
2026-01-27 18:02:16 +00:00
parent 7efd7e4e98
commit 1c4fda6aa2
7 changed files with 335 additions and 351 deletions
+12 -5
View File
@@ -1,7 +1,7 @@
"""
Database module for WebAuthn passkey authentication.
Read: Access _db._data directly, use build_* to convert to public structs.
Read: Access db() directly, use build_* to convert to public structs.
CTX: get_session_context(key) returns SessionContext with effective permissions.
Write: Functions validate and commit, or raise ValueError.
@@ -9,7 +9,7 @@ Usage:
from paskia import db
# Read (after init)
user_data = db._db._data.users[user_uuid]
user_data = db.db().users[user_uuid]
user = db.build_user(user_uuid)
# Context
@@ -26,8 +26,6 @@ from paskia.db.background import (
stop_cleanup,
)
from paskia.db.operations import (
DB,
_db,
add_permission_to_organization,
add_permission_to_role,
bootstrap,
@@ -82,6 +80,7 @@ from paskia.db.operations import (
update_user_role_in_organization,
)
from paskia.db.structs import (
DB,
Credential,
Org,
Permission,
@@ -92,6 +91,14 @@ from paskia.db.structs import (
User,
)
def db() -> DB:
"""Get the database instance for direct read access."""
from paskia.db.operations import _db
return _db
__all__ = [
# Types
"Credential",
@@ -104,7 +111,7 @@ __all__ = [
"SessionContext",
"User",
# Instance
"_db",
"db",
"init",
# Background
"start_background",
+9 -11
View File
@@ -8,8 +8,6 @@ import asyncio
import logging
from datetime import datetime, timezone
from paskia.db.jsonl import flush_changes
# Flush changes to disk every N seconds
FLUSH_INTERVAL = 1
# Cleanup expired items every N seconds (cheap when nothing to remove)
@@ -24,7 +22,7 @@ def cleanup() -> None:
"""Remove expired sessions and reset tokens from the database."""
from paskia.db.operations import _db
if _db is None or _db._data is None:
if _db is None:
return
with _db.transaction("expiry"):
@@ -32,27 +30,27 @@ def cleanup() -> None:
# Clean expired sessions
to_delete_sessions = [
k for k, s in _db._data.sessions.items() if s.expiry < current_time
k for k, s in _db.sessions.items() if s.expiry < current_time
]
for k in to_delete_sessions:
del _db._data.sessions[k]
del _db.sessions[k]
# Clean expired reset tokens
to_delete_tokens = [
k for k, t in _db._data.reset_tokens.items() if t.expiry < current_time
k for k, t in _db.reset_tokens.items() if t.expiry < current_time
]
for k in to_delete_tokens:
del _db._data.reset_tokens[k]
del _db.reset_tokens[k]
async def flush() -> None:
"""Write all pending database changes to disk."""
from paskia.db.operations import _db
from paskia.db.operations import _store
if _db is None:
_logger.warning("flush() called but _db is None")
if _store is None:
_logger.warning("flush() called but _store is None")
return
await flush_changes(_db.db_path, _db._pending_changes)
await _store.flush()
async def _background_loop():
+94 -3
View File
@@ -1,19 +1,25 @@
"""
JSONL persistence layer for the database.
Handles file I/O, JSON diffs, and persistence. Works with plain JSON/dict data.
Uses aiofiles for async I/O operations.
"""
from __future__ import annotations
import json
import logging
import sys
from collections import deque
from contextlib import contextmanager
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from uuid import UUID
import aiofiles
import jsondiff
import msgspec
from paskia.db.structs import DB, SessionContext
_logger = logging.getLogger(__name__)
# Default database path
@@ -144,3 +150,88 @@ async def flush_changes(
for change in reversed(changes_to_write):
pending_changes.appendleft(change)
return False
class JsonlStore:
"""JSONL persistence layer for a DB instance."""
def __init__(self, db: DB, db_path: str = DB_PATH_DEFAULT):
self.db: DB = db
self.db_path = Path(db_path)
self._previous_builtins: dict[str, Any] = {}
self._pending_changes: deque[_ChangeRecord] = deque()
self._current_action: str = "system"
self._current_user: str | None = None
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)
try:
data_dict = await load_jsonl(self.db_path)
if data_dict:
decoder = msgspec.json.Decoder(DB)
self.db = decoder.decode(msgspec.json.encode(data_dict))
self.db._store = self
self._previous_builtins = data_dict
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:
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
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)
@contextmanager
def transaction(
self,
action: str,
ctx: SessionContext | None = None,
*,
user: str | None = None,
):
"""Wrap writes in transaction. Queues change on successful exit.
Args:
action: Describes the operation (e.g., "Created user", "Login")
ctx: Session context of user performing the action (None for system operations)
user: User UUID string (alternative to ctx when full context unavailable)
"""
old_action = self._current_action
old_user = self._current_user
self._current_action = action
# Prefer ctx.user.uuid if ctx provided, otherwise use user param
self._current_user = str(ctx.user.uuid) if ctx else user
try:
yield
self._queue_change()
finally:
self._current_action = old_action
self._current_user = old_user
async def flush(self) -> bool:
"""Write all pending changes to disk."""
return await flush_changes(self.db_path, self._pending_changes)
+166 -284
View File
@@ -1,36 +1,25 @@
"""
Database for WebAuthn passkey authentication.
Read operations: Access _db._data directly, use build_* helpers to get public structs.
Read operations: Access _db directly, use build_* helpers to get public structs.
Context lookup: get_session_context() returns full SessionContext with effective permissions.
Write operations: Functions that validate and commit, or raise ValueError.
"""
import hashlib
import json
import logging
import os
import secrets
import sys
from collections import deque
from contextlib import contextmanager
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from uuid import UUID
import msgspec
from paskia.db.jsonl import (
DB_PATH_DEFAULT,
_ChangeRecord,
compute_diff,
create_change_record,
load_jsonl,
JsonlStore,
)
from paskia.db.structs import (
DB,
Credential,
DatabaseData,
Org,
Permission,
ResetToken,
@@ -43,119 +32,20 @@ from paskia.util.passphrase import is_well_formed as _is_passphrase
_logger = logging.getLogger(__name__)
# msgspec encoder/decoder
_json_encoder = msgspec.json.Encoder()
_json_decoder = msgspec.json.Decoder(DatabaseData)
class DB:
"""In-memory database with JSONL persistence.
Access data directly via _data for reads.
Use transaction() context manager for writes.
"""
def __init__(self, db_path: str = DB_PATH_DEFAULT):
self.db_path = Path(db_path)
self._data = DatabaseData(
permissions={},
orgs={},
roles={},
users={},
credentials={},
sessions={},
reset_tokens={},
)
self._previous_builtins: dict[str, Any] = {}
self._pending_changes: deque[_ChangeRecord] = deque()
self._current_action: str = "system"
self._current_user: str | None = None
async def load(self, db_path: str | None = None) -> None:
"""Load data from JSONL change log.
If file doesn't exist or is empty, keeps the initialized empty structure and
sets _previous_builtins to {} for creating a new database.
"""
if db_path is not None:
self.db_path = Path(db_path)
try:
data_dict = await load_jsonl(self.db_path)
if data_dict: # Only decode if we have data
self._data = _json_decoder.decode(_json_encoder.encode(data_dict))
# Track the JSONL file state directly - this is what we diff against
self._previous_builtins = data_dict
# If data_dict is empty, keep initialized _data and _previous_builtins = {}
except ValueError:
if self.db_path.exists():
raise # File exists but failed to load - re-raise
# File doesn't exist: keep initialized _data, _previous_builtins stays {}
def _queue_change(self) -> None:
current = msgspec.to_builtins(self._data)
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:
try:
user_uuid = UUID(self._current_user)
if user_uuid in self._data.users:
user_display = self._data.users[user_uuid].display_name
except (ValueError, KeyError):
user_display = self._current_user
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)
@contextmanager
def transaction(
self,
action: str,
ctx: SessionContext | None = None,
*,
user: str | None = None,
):
"""Wrap writes in transaction. Queues change on successful exit.
Args:
action: Describes the operation (e.g., "Created user", "Login")
ctx: Session context of user performing the action (None for system operations)
user: User UUID string (alternative to ctx when full context unavailable)
"""
old_action = self._current_action
old_user = self._current_user
self._current_action = action
# Prefer ctx.user.uuid if ctx provided, otherwise use user param
self._current_user = str(ctx.user.uuid) if ctx else user
try:
yield
self._queue_change()
finally:
self._current_action = old_action
self._current_user = old_user
# Global instance, always available (empty until init() loads data)
# Global database instance (empty until init() loads data)
_db = DB()
_store = JsonlStore(_db)
_db._store = _store
async def init(*args, **kwargs):
"""Load database from JSONL file."""
global _db
db_path = os.environ.get("PASKIA_DB", DB_PATH_DEFAULT)
if db_path.startswith("json:"):
db_path = db_path[5:]
await _db.load(db_path)
await _store.load(db_path)
_db = _store.db
# -------------------------------------------------------------------------
@@ -164,12 +54,10 @@ async def init(*args, **kwargs):
def build_org(uuid: UUID, include_roles: bool = False) -> Org:
o = _db._data.orgs[uuid]
o.permissions = {pid for pid, p in _db._data.permissions.items() if uuid in p.orgs}
o = _db.orgs[uuid]
o.permissions = {pid for pid, p in _db.permissions.items() if uuid in p.orgs}
if include_roles:
o.roles = [
_db._data.roles[rid] for rid, r in _db._data.roles.items() if r.org == uuid
]
o.roles = [_db.roles[rid] for rid, r in _db.roles.items() if r.org == uuid]
return o
@@ -191,7 +79,7 @@ def get_permission(uuid: UUID) -> Permission | None:
- Get permission for renaming its scope (admin.py:1031)
- Get permission to check scope before deleting (admin.py:1071)
"""
return _db._data.permissions.get(uuid)
return _db.permissions.get(uuid)
def get_permission_by_scope(scope: str) -> Permission | None:
@@ -200,7 +88,7 @@ def get_permission_by_scope(scope: str) -> Permission | None:
Call sites:
- Check if system is already bootstrapped by looking for auth:admin permission (bootstrap.py:113)
"""
for p in _db._data.permissions.values():
for p in _db.permissions.values():
if p.scope == scope:
return p
return None
@@ -212,7 +100,7 @@ def get_permissions_by_scope(scope: str) -> list[Permission]:
Since scopes are not unique, this returns all matching permissions.
Use this for scope-based permission checking.
"""
return [p for p in _db._data.permissions.values() if p.scope == scope]
return [p for p in _db.permissions.values() if p.scope == scope]
def list_permissions() -> list[Permission]:
@@ -225,7 +113,7 @@ def list_permissions() -> list[Permission]:
- List permissions to check admin permissions when deleting permission (admin.py:882)
- Admin API endpoint to list permissions (admin.py:914)
"""
return list(_db._data.permissions.values())
return list(_db.permissions.values())
def get_permission_organizations(scope: str) -> list[Org]:
@@ -235,7 +123,7 @@ def get_permission_organizations(scope: str) -> list[Org]:
- Get organizations that can grant auth:admin to find admin users (bootstrap.py:67)
- Get organizations with auth:admin permission to find admin users for reset targets (reset.py:29,40,55)
"""
for p in _db._data.permissions.values():
for p in _db.permissions.values():
if p.scope == scope:
return [build_org(org_uuid) for org_uuid in p.orgs]
return []
@@ -248,7 +136,7 @@ def get_organization(uuid: UUID) -> Org | None:
- Get organization when creating a role to check grantable permissions (admin.py:271)
- Get organization when adding permission to role to check if org can grant it (admin.py:352)
"""
return build_org(uuid, include_roles=True) if uuid in _db._data.orgs else None
return build_org(uuid, include_roles=True) if uuid in _db.orgs else None
def list_organizations() -> list[Org]:
@@ -258,7 +146,7 @@ def list_organizations() -> list[Org]:
- List organizations during migration (migrate/__init__.py:131)
- Admin API endpoint to list organizations (admin.py:94)
"""
return [build_org(uuid, include_roles=True) for uuid in _db._data.orgs]
return [build_org(uuid, include_roles=True) for uuid in _db.orgs]
def get_organization_users(org_uuid: UUID) -> list[tuple[User, str]]:
@@ -270,13 +158,9 @@ def get_organization_users(org_uuid: UUID) -> list[tuple[User, str]]:
- Get users from organization to check if admin has credentials (bootstrap.py:73)
"""
role_map = {
rid: r.display_name for rid, r in _db._data.roles.items() if r.org == org_uuid
rid: r.display_name for rid, r in _db.roles.items() if r.org == org_uuid
}
return [
(u, role_map[u.role])
for u in _db._data.users.values()
if u.role in role_map
]
return [(u, role_map[u.role]) for u in _db.users.values() if u.role in role_map]
def get_role(uuid: UUID) -> Role | None:
@@ -288,7 +172,7 @@ def get_role(uuid: UUID) -> Role | None:
- Get role to remove permission from it (admin.py:380)
- Get role to delete it (admin.py:421)
"""
return _db._data.roles.get(uuid)
return _db.roles.get(uuid)
def get_roles_by_organization(org_uuid: UUID) -> list[Role]:
@@ -298,7 +182,7 @@ def get_roles_by_organization(org_uuid: UUID) -> list[Role]:
- Get roles by organization when creating a user to find the role by name (admin.py:459)
- Get roles by organization when updating user role to validate the new role name (admin.py:498)
"""
return [r for r in _db._data.roles.values() if r.org == org_uuid]
return [r for r in _db.roles.values() if r.org == org_uuid]
def get_user_by_uuid(uuid: UUID) -> User | None:
@@ -309,7 +193,7 @@ def get_user_by_uuid(uuid: UUID) -> User | None:
- Get user from reset token for registration info (api.py:127)
- Get user for listing user credentials in admin API (admin.py:594)
"""
return _db._data.users.get(uuid)
return _db.users.get(uuid)
def get_user_organization(user_uuid: UUID) -> tuple[Org, str]:
@@ -325,12 +209,12 @@ def get_user_organization(user_uuid: UUID) -> tuple[Org, str]:
- Get user's organization for deleting user credential (admin.py:754)
- Get user's organization for deleting user session (admin.py:783)
"""
if user_uuid not in _db._data.users:
if user_uuid not in _db.users:
raise ValueError(f"User {user_uuid} not found")
role_uuid = _db._data.users[user_uuid].role
if role_uuid not in _db._data.roles:
role_uuid = _db.users[user_uuid].role
if role_uuid not in _db.roles:
raise ValueError(f"Role {role_uuid} not found")
role_data = _db._data.roles[role_uuid]
role_data = _db.roles[role_uuid]
org_uuid = role_data.org
return build_org(org_uuid, include_roles=True), role_data.display_name
@@ -342,7 +226,7 @@ def get_credential_by_id(credential_id: bytes) -> Credential | None:
- Get credential by ID for WebAuthn authentication (ws.py:132)
- Get credential by ID for remote authentication (remote.py:325)
"""
for c in _db._data.credentials.values():
for c in _db.credentials.values():
if c.credential_id == credential_id:
return c
return None
@@ -359,7 +243,7 @@ def get_credentials_by_user_uuid(user_uuid: UUID) -> list[Credential]:
- Get credentials to check if admin user has credentials (bootstrap.py:81)
- Get credentials for user info formatting (userinfo.py:51)
"""
return [c for c in _db._data.credentials.values() if c.user == user_uuid]
return [c for c in _db.credentials.values() if c.user == user_uuid]
def get_session(key: str) -> Session | None:
@@ -371,7 +255,7 @@ def get_session(key: str) -> Session | None:
- Get session to refresh it (authsession.py:59)
- Get session to delete it in user API (user.py:94)
"""
return _db._data.sessions.get(key)
return _db.sessions.get(key)
def list_sessions_for_user(user_uuid: UUID) -> list[Session]:
@@ -381,7 +265,7 @@ def list_sessions_for_user(user_uuid: UUID) -> list[Session]:
- List sessions for user info (userinfo.py:75)
- List sessions for user details API (admin.py:651)
"""
return [s for s in _db._data.sessions.values() if s.user == user_uuid]
return [s for s in _db.sessions.values() if s.user == user_uuid]
def _reset_key(passphrase: str) -> bytes:
@@ -402,7 +286,7 @@ def get_reset_token(passphrase: str) -> ResetToken | None:
- Get reset token to validate it (authsession.py:34)
"""
key = _reset_key(passphrase)
return _db._data.reset_tokens.get(key)
return _db.reset_tokens.get(key)
# -------------------------------------------------------------------------
@@ -428,10 +312,10 @@ def get_session_context(
"""
from paskia.util.hostutil import normalize_host
if session_key not in _db._data.sessions:
if session_key not in _db.sessions:
return None
s = _db._data.sessions[session_key]
s = _db.sessions[session_key]
if s.expiry < datetime.now(timezone.utc):
return None
@@ -446,28 +330,28 @@ def get_session_context(
return None
# Validate user exists
if s.user not in _db._data.users:
if s.user not in _db.users:
return None
# Validate role exists
role_uuid = _db._data.users[s.user].role
if role_uuid not in _db._data.roles:
role_uuid = _db.users[s.user].role
if role_uuid not in _db.roles:
return None
# Validate org exists
org_uuid = _db._data.roles[role_uuid].org
if org_uuid not in _db._data.orgs:
org_uuid = _db.roles[role_uuid].org
if org_uuid not in _db.orgs:
return None
session = _db._data.sessions[session_key]
user = _db._data.users[s.user]
role = _db._data.roles[role_uuid]
session = _db.sessions[session_key]
user = _db.users[s.user]
role = _db.roles[role_uuid]
org = build_org(org_uuid)
# Credential must exist (sessions are cascade-deleted when credential is deleted)
if s.credential not in _db._data.credentials:
if s.credential not in _db.credentials:
return None
credential = _db._data.credentials[s.credential]
credential = _db.credentials[s.credential]
# Effective permissions: role's permissions that the org can grant
# Also filter by domain if host is provided
@@ -479,13 +363,13 @@ def get_session_context(
for perm_uuid in role.permission_set:
if perm_uuid not in org_perm_uuids:
continue
if perm_uuid not in _db._data.permissions:
if perm_uuid not in _db.permissions:
continue
p = _db._data.permissions[perm_uuid]
p = _db.permissions[perm_uuid]
# Check domain restriction
if p.domain is not None and p.domain != host_without_port:
continue
effective_perms.append(_db._data.permissions[perm_uuid])
effective_perms.append(_db.permissions[perm_uuid])
return SessionContext(
session=session,
@@ -504,20 +388,20 @@ def get_session_context(
def create_permission(perm: Permission, *, ctx: SessionContext | None = None) -> None:
"""Create a new permission."""
if perm.uuid in _db._data.permissions:
if perm.uuid in _db.permissions:
raise ValueError(f"Permission {perm.uuid} already exists")
with _db.transaction("Created permission", ctx):
_db._data.permissions[perm.uuid] = perm
_db.permissions[perm.uuid] = perm
def update_permission(perm: Permission, *, ctx: SessionContext | None = None) -> None:
"""Update a permission's scope, display_name, and domain."""
if perm.uuid not in _db._data.permissions:
if perm.uuid not in _db.permissions:
raise ValueError(f"Permission {perm.uuid} not found")
with _db.transaction("Updated permission", ctx):
_db._data.permissions[perm.uuid].scope = perm.scope
_db._data.permissions[perm.uuid].display_name = perm.display_name
_db._data.permissions[perm.uuid].domain = perm.domain
_db.permissions[perm.uuid].scope = perm.scope
_db.permissions[perm.uuid].display_name = perm.display_name
_db.permissions[perm.uuid].domain = perm.domain
def rename_permission(
@@ -533,25 +417,25 @@ def rename_permission(
Since roles reference permissions by UUID, no role updates are needed.
Note: Scopes do not need to be unique (same scope with different domains is valid).
"""
if uuid not in _db._data.permissions:
if uuid not in _db.permissions:
raise ValueError(f"Permission {uuid} not found")
with _db.transaction("Renamed permission", ctx):
# Update the permission
_db._data.permissions[uuid].scope = new_scope
_db._data.permissions[uuid].display_name = display_name
_db._data.permissions[uuid].domain = domain
_db.permissions[uuid].scope = new_scope
_db.permissions[uuid].display_name = display_name
_db.permissions[uuid].domain = domain
def delete_permission(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
"""Delete a permission and remove it from all roles."""
if uuid not in _db._data.permissions:
if uuid not in _db.permissions:
raise ValueError(f"Permission {uuid} not found")
with _db.transaction("Deleted permission", ctx):
# Remove this permission from all roles
for role in _db._data.roles.values():
for role in _db.roles.values():
role.permissions.pop(uuid, None)
del _db._data.permissions[uuid]
del _db.permissions[uuid]
def create_organization(org: Org, *, ctx: SessionContext | None = None) -> None:
@@ -559,13 +443,13 @@ def create_organization(org: Org, *, ctx: SessionContext | None = None) -> None:
Automatically creates an 'Administration' role with auth:org:admin permission.
"""
if org.uuid in _db._data.orgs:
if org.uuid in _db.orgs:
raise ValueError(f"Organization {org.uuid} already exists")
with _db.transaction("Created organization", ctx):
new_org = Org(
display_name=org.display_name, created_at=datetime.now(timezone.utc)
)
_db._data.orgs[org.uuid] = new_org
_db.orgs[org.uuid] = new_org
new_org.uuid = org.uuid
# Create Administration role with org admin permission
import uuid7
@@ -573,7 +457,7 @@ def create_organization(org: Org, *, ctx: SessionContext | None = None) -> None:
admin_role_uuid = uuid7.create()
# Find the auth:org:admin permission UUID
org_admin_perm_uuid = None
for pid, p in _db._data.permissions.items():
for pid, p in _db.permissions.items():
if p.scope == "auth:org:admin":
org_admin_perm_uuid = pid
break
@@ -584,7 +468,7 @@ def create_organization(org: Org, *, ctx: SessionContext | None = None) -> None:
permissions=role_permissions,
)
admin_role.uuid = admin_role_uuid
_db._data.roles[admin_role_uuid] = admin_role
_db.roles[admin_role_uuid] = admin_role
def update_organization_name(
@@ -594,29 +478,29 @@ def update_organization_name(
ctx: SessionContext | None = None,
) -> None:
"""Update organization display name."""
if uuid not in _db._data.orgs:
if uuid not in _db.orgs:
raise ValueError(f"Organization {uuid} not found")
with _db.transaction("Renamed organization", ctx):
_db._data.orgs[uuid].display_name = display_name
_db.orgs[uuid].display_name = display_name
def delete_organization(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
"""Delete organization and all its roles/users."""
if uuid not in _db._data.orgs:
if uuid not in _db.orgs:
raise ValueError(f"Organization {uuid} not found")
with _db.transaction("Deleted organization", ctx):
# Remove org from all permissions
for p in _db._data.permissions.values():
for p in _db.permissions.values():
p.orgs.pop(uuid, None)
# Delete roles in this org
role_uuids = [rid for rid, r in _db._data.roles.items() if r.org == uuid]
role_uuids = [rid for rid, r in _db.roles.items() if r.org == uuid]
for rid in role_uuids:
del _db._data.roles[rid]
del _db.roles[rid]
# Delete users with those roles
user_uuids = [uid for uid, u in _db._data.users.items() if u.role in role_uuids]
user_uuids = [uid for uid, u in _db.users.items() if u.role in role_uuids]
for uid in user_uuids:
del _db._data.users[uid]
del _db._data.orgs[uuid]
del _db.users[uid]
del _db.orgs[uuid]
def add_permission_to_organization(
@@ -626,14 +510,14 @@ def add_permission_to_organization(
ctx: SessionContext | None = None,
) -> None:
"""Grant a permission to an organization by UUID."""
if org_uuid not in _db._data.orgs:
if org_uuid not in _db.orgs:
raise ValueError(f"Organization {org_uuid} not found")
if permission_uuid not in _db._data.permissions:
if permission_uuid not in _db.permissions:
raise ValueError(f"Permission {permission_uuid} not found")
with _db.transaction("Granted org permission", ctx):
_db._data.permissions[permission_uuid].orgs[org_uuid] = True
_db.permissions[permission_uuid].orgs[org_uuid] = True
def remove_permission_from_organization(
@@ -643,24 +527,24 @@ def remove_permission_from_organization(
ctx: SessionContext | None = None,
) -> None:
"""Remove a permission from an organization by UUID."""
if org_uuid not in _db._data.orgs:
if org_uuid not in _db.orgs:
raise ValueError(f"Organization {org_uuid} not found")
if permission_uuid not in _db._data.permissions:
if permission_uuid not in _db.permissions:
return # Permission not found, silently return
with _db.transaction("Revoked org permission", ctx):
_db._data.permissions[permission_uuid].orgs.pop(org_uuid, None)
_db.permissions[permission_uuid].orgs.pop(org_uuid, None)
def create_role(role: Role, *, ctx: SessionContext | None = None) -> None:
"""Create a new role."""
if role.uuid in _db._data.roles:
if role.uuid in _db.roles:
raise ValueError(f"Role {role.uuid} already exists")
if role.org not in _db._data.orgs:
if role.org not in _db.orgs:
raise ValueError(f"Organization {role.org} not found")
with _db.transaction("Created role", ctx):
_db._data.roles[role.uuid] = role
_db.roles[role.uuid] = role
def update_role_name(
@@ -670,10 +554,10 @@ def update_role_name(
ctx: SessionContext | None = None,
) -> None:
"""Update role display name."""
if uuid not in _db._data.roles:
if uuid not in _db.roles:
raise ValueError(f"Role {uuid} not found")
with _db.transaction("Renamed role", ctx):
_db._data.roles[uuid].display_name = display_name
_db.roles[uuid].display_name = display_name
def add_permission_to_role(
@@ -683,12 +567,12 @@ def add_permission_to_role(
ctx: SessionContext | None = None,
) -> None:
"""Add permission to role by UUID."""
if role_uuid not in _db._data.roles:
if role_uuid not in _db.roles:
raise ValueError(f"Role {role_uuid} not found")
if permission_uuid not in _db._data.permissions:
if permission_uuid not in _db.permissions:
raise ValueError(f"Permission {permission_uuid} not found")
with _db.transaction("Granted role permission", ctx):
_db._data.roles[role_uuid].permissions[permission_uuid] = True
_db.roles[role_uuid].permissions[permission_uuid] = True
def remove_permission_from_role(
@@ -698,31 +582,31 @@ def remove_permission_from_role(
ctx: SessionContext | None = None,
) -> None:
"""Remove permission from role by UUID."""
if role_uuid not in _db._data.roles:
if role_uuid not in _db.roles:
raise ValueError(f"Role {role_uuid} not found")
with _db.transaction("Revoked role permission", ctx):
_db._data.roles[role_uuid].permissions.pop(permission_uuid, None)
_db.roles[role_uuid].permissions.pop(permission_uuid, None)
def delete_role(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
"""Delete a role."""
if uuid not in _db._data.roles:
if uuid not in _db.roles:
raise ValueError(f"Role {uuid} not found")
# Check no users have this role
if any(u.role == uuid for u in _db._data.users.values()):
if any(u.role == uuid for u in _db.users.values()):
raise ValueError(f"Cannot delete role {uuid}: users still assigned")
with _db.transaction("Deleted role", ctx):
del _db._data.roles[uuid]
del _db.roles[uuid]
def create_user(new_user: User, *, ctx: SessionContext | None = None) -> None:
"""Create a new user."""
if new_user.uuid in _db._data.users:
if new_user.uuid in _db.users:
raise ValueError(f"User {new_user.uuid} already exists")
if new_user.role not in _db._data.roles:
if new_user.role not in _db.roles:
raise ValueError(f"Role {new_user.role} not found")
with _db.transaction("Created user", ctx):
_db._data.users[new_user.uuid] = new_user
_db.users[new_user.uuid] = new_user
def update_user_display_name(
@@ -738,12 +622,12 @@ def update_user_display_name(
"""
if isinstance(uuid, str):
uuid = UUID(uuid)
if uuid not in _db._data.users:
if uuid not in _db.users:
raise ValueError(f"User {uuid} not found")
# For self-service, derive user from the uuid being modified
user_str = str(uuid) if not ctx else None
with _db.transaction("Renamed user", ctx, user=user_str):
_db._data.users[uuid].display_name = display_name
_db.users[uuid].display_name = display_name
def update_user_role(
@@ -753,12 +637,12 @@ def update_user_role(
ctx: SessionContext | None = None,
) -> None:
"""Update user's role."""
if uuid not in _db._data.users:
if uuid not in _db.users:
raise ValueError(f"User {uuid} not found")
if role_uuid not in _db._data.roles:
if role_uuid not in _db.roles:
raise ValueError(f"Role {role_uuid} not found")
with _db.transaction("Changed user role", ctx):
_db._data.users[uuid].role = role_uuid
_db.users[uuid].role = role_uuid
def update_user_role_in_organization(
@@ -768,52 +652,52 @@ def update_user_role_in_organization(
ctx: SessionContext | None = None,
) -> None:
"""Update user's role by role name within their current organization."""
if user_uuid not in _db._data.users:
if user_uuid not in _db.users:
raise ValueError(f"User {user_uuid} not found")
current_role_uuid = _db._data.users[user_uuid].role
if current_role_uuid not in _db._data.roles:
current_role_uuid = _db.users[user_uuid].role
if current_role_uuid not in _db.roles:
raise ValueError("Current role not found")
org_uuid = _db._data.roles[current_role_uuid].org
org_uuid = _db.roles[current_role_uuid].org
# Find role by name in the same org
new_role_uuid = None
for rid, r in _db._data.roles.items():
for rid, r in _db.roles.items():
if r.org == org_uuid and r.display_name == role_name:
new_role_uuid = rid
break
if new_role_uuid is None:
raise ValueError(f"Role '{role_name}' not found in organization")
with _db.transaction("Changed user role", ctx):
_db._data.users[user_uuid].role = new_role_uuid
_db.users[user_uuid].role = new_role_uuid
def delete_user(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
"""Delete user and their credentials/sessions."""
if uuid not in _db._data.users:
if uuid not in _db.users:
raise ValueError(f"User {uuid} not found")
with _db.transaction("Deleted user", ctx):
# Delete credentials
cred_uuids = [cid for cid, c in _db._data.credentials.items() if c.user == uuid]
cred_uuids = [cid for cid, c in _db.credentials.items() if c.user == uuid]
for cid in cred_uuids:
del _db._data.credentials[cid]
del _db.credentials[cid]
# Delete sessions
sess_keys = [k for k, s in _db._data.sessions.items() if s.user == uuid]
sess_keys = [k for k, s in _db.sessions.items() if s.user == uuid]
for k in sess_keys:
del _db._data.sessions[k]
del _db.sessions[k]
# Delete reset tokens
token_keys = [k for k, t in _db._data.reset_tokens.items() if t.user == uuid]
token_keys = [k for k, t in _db.reset_tokens.items() if t.user == uuid]
for k in token_keys:
del _db._data.reset_tokens[k]
del _db._data.users[uuid]
del _db.reset_tokens[k]
del _db.users[uuid]
def create_credential(cred: Credential, *, ctx: SessionContext | None = None) -> None:
"""Create a new credential."""
if cred.uuid in _db._data.credentials:
if cred.uuid in _db.credentials:
raise ValueError(f"Credential {cred.uuid} already exists")
if cred.user not in _db._data.users:
if cred.user not in _db.users:
raise ValueError(f"User {cred.user} not found")
with _db.transaction("Added credential", ctx):
_db._data.credentials[cred.uuid] = cred
_db.credentials[cred.uuid] = cred
def update_credential_sign_count(
@@ -824,12 +708,12 @@ def update_credential_sign_count(
ctx: SessionContext | None = None,
) -> None:
"""Update credential sign count and last_used."""
if uuid not in _db._data.credentials:
if uuid not in _db.credentials:
raise ValueError(f"Credential {uuid} not found")
with _db.transaction("Updated credential", ctx):
_db._data.credentials[uuid].sign_count = sign_count
_db.credentials[uuid].sign_count = sign_count
if last_used:
_db._data.credentials[uuid].last_used = last_used
_db.credentials[uuid].last_used = last_used
def delete_credential(
@@ -842,18 +726,18 @@ def delete_credential(
If user_uuid is provided, validates that the credential belongs to that user.
"""
if uuid not in _db._data.credentials:
if uuid not in _db.credentials:
raise ValueError(f"Credential {uuid} not found")
if user_uuid is not None:
cred_user = _db._data.credentials[uuid].user
cred_user = _db.credentials[uuid].user
if cred_user != user_uuid:
raise ValueError(f"Credential {uuid} does not belong to user {user_uuid}")
with _db.transaction("Deleted credential", ctx):
# Delete all sessions using this credential
keys = [k for k, s in _db._data.sessions.items() if s.credential == uuid]
keys = [k for k, s in _db.sessions.items() if s.credential == uuid]
for k in keys:
del _db._data.sessions[k]
del _db._data.credentials[uuid]
del _db.sessions[k]
del _db.credentials[uuid]
def create_session(
@@ -868,14 +752,14 @@ def create_session(
ctx: SessionContext | None = None,
) -> None:
"""Create a new session."""
if key in _db._data.sessions:
if key in _db.sessions:
raise ValueError("Session already exists")
if user_uuid not in _db._data.users:
if user_uuid not in _db.users:
raise ValueError(f"User {user_uuid} not found")
if credential_uuid not in _db._data.credentials:
if credential_uuid not in _db.credentials:
raise ValueError(f"Credential {credential_uuid} not found")
with _db.transaction("Created session", ctx):
_db._data.sessions[key] = Session(
_db.sessions[key] = Session(
user=user_uuid,
credential=credential_uuid,
host=host,
@@ -895,10 +779,10 @@ def update_session(
ctx: SessionContext | None = None,
) -> None:
"""Update session metadata."""
if key not in _db._data.sessions:
if key not in _db.sessions:
raise ValueError("Session not found")
with _db.transaction("Updated session", ctx):
s = _db._data.sessions[key]
s = _db.sessions[key]
if host is not None:
s.host = host
if ip is not None:
@@ -920,12 +804,12 @@ def delete_session(key: str, *, ctx: SessionContext | None = None) -> None:
For logout (user deleting own session), ctx can be None and user is derived from session.
For admin operations, ctx should be provided.
"""
if key not in _db._data.sessions:
if key not in _db.sessions:
raise ValueError("Session not found")
# For self-service logout, derive user from the session being deleted
user_str = str(_db._data.sessions[key].user) if not ctx else None
user_str = str(_db.sessions[key].user) if not ctx else None
with _db.transaction("Deleted session", ctx, user=user_str):
del _db._data.sessions[key]
del _db.sessions[key]
def delete_sessions_for_user(
@@ -939,9 +823,9 @@ def delete_sessions_for_user(
# For self-service, derive user from the user_uuid param
user_str = str(user_uuid) if not ctx else None
with _db.transaction("Deleted user sessions", ctx, user=user_str):
keys = [k for k, s in _db._data.sessions.items() if s.user == user_uuid]
keys = [k for k, s in _db.sessions.items() if s.user == user_uuid]
for k in keys:
del _db._data.sessions[k]
del _db.sessions[k]
def create_reset_token(
@@ -958,24 +842,24 @@ def create_reset_token(
For admin operations, ctx should be provided.
"""
key = _reset_key(passphrase)
if key in _db._data.reset_tokens:
if key in _db.reset_tokens:
raise ValueError("Reset token already exists")
if user_uuid not in _db._data.users:
if user_uuid not in _db.users:
raise ValueError(f"User {user_uuid} not found")
# For self-service, derive user from the user_uuid param
user_str = str(user_uuid) if not ctx else None
with _db.transaction("Created reset token", ctx, user=user_str):
_db._data.reset_tokens[key] = ResetToken(
_db.reset_tokens[key] = ResetToken(
user=user_uuid, expiry=expiry, token_type=token_type
)
def delete_reset_token(key: bytes, *, ctx: SessionContext | None = None) -> None:
"""Delete a reset token."""
if key not in _db._data.reset_tokens:
if key not in _db.reset_tokens:
raise ValueError("Reset token not found")
with _db.transaction("Deleted reset token", ctx):
del _db._data.reset_tokens[key]
del _db.reset_tokens[key]
# -------------------------------------------------------------------------
@@ -988,15 +872,13 @@ def cleanup_expired() -> int:
now = datetime.now(timezone.utc)
count = 0
with _db.transaction("Cleaned up expired"):
expired_sessions = [k for k, s in _db._data.sessions.items() if s.expiry < now]
expired_sessions = [k for k, s in _db.sessions.items() if s.expiry < now]
for k in expired_sessions:
del _db._data.sessions[k]
del _db.sessions[k]
count += 1
expired_tokens = [
k for k, t in _db._data.reset_tokens.items() if t.expiry < now
]
expired_tokens = [k for k, t in _db.reset_tokens.items() if t.expiry < now]
for k in expired_tokens:
del _db._data.reset_tokens[k]
del _db.reset_tokens[k]
count += 1
return count
@@ -1032,22 +914,22 @@ def login(
if isinstance(user_uuid, str):
user_uuid = UUID(user_uuid)
now = datetime.now(timezone.utc)
if user_uuid not in _db._data.users:
if user_uuid not in _db.users:
raise ValueError(f"User {user_uuid} not found")
if credential.uuid not in _db._data.credentials:
if credential.uuid not in _db.credentials:
raise ValueError(f"Credential {credential.uuid} not found")
session_key = _create_token()
user_str = str(user_uuid)
with _db.transaction("User logged in", user=user_str):
# Update user
_db._data.users[user_uuid].last_seen = now
_db._data.users[user_uuid].visits += 1
_db.users[user_uuid].last_seen = now
_db.users[user_uuid].visits += 1
# Update credential
_db._data.credentials[credential.uuid].sign_count = credential.sign_count
_db._data.credentials[credential.uuid].last_used = now
_db.credentials[credential.uuid].sign_count = credential.sign_count
_db.credentials[credential.uuid].last_used = now
# Create session
_db._data.sessions[session_key] = Session(
_db.sessions[session_key] = Session(
user=user_uuid,
credential=credential.uuid,
host=host,
@@ -1083,20 +965,20 @@ def create_credential_session(
expiry = now + SESSION_LIFETIME
session_key = _create_token()
if user_uuid not in _db._data.users:
if user_uuid not in _db.users:
raise ValueError(f"User {user_uuid} not found")
user_str = str(user_uuid)
with _db.transaction("Registered credential", user=user_str):
# Update display name if provided
if display_name:
_db._data.users[user_uuid].display_name = display_name
_db.users[user_uuid].display_name = display_name
# Create credential
_db._data.credentials[credential.uuid] = credential
_db.credentials[credential.uuid] = credential
# Create session
_db._data.sessions[session_key] = Session(
_db.sessions[session_key] = Session(
user=user_uuid,
credential=credential.uuid,
host=host,
@@ -1107,8 +989,8 @@ def create_credential_session(
# Delete reset token if provided
if reset_key:
if reset_key in _db._data.reset_tokens:
del _db._data.reset_tokens[reset_key]
if reset_key in _db.reset_tokens:
del _db.reset_tokens[reset_key]
return session_key
@@ -1150,7 +1032,7 @@ def bootstrap(
from paskia.util.passphrase import generate as generate_passphrase
# Check if system is already bootstrapped
for p in _db._data.permissions.values():
for p in _db.permissions.values():
if p.scope == "auth:admin":
raise ValueError(
"System already bootstrapped (auth:admin permission exists)"
@@ -1180,7 +1062,7 @@ def bootstrap(
orgs={org_uuid: True}, # Grant to org
)
perm_admin.uuid = perm_admin_uuid
_db._data.permissions[perm_admin_uuid] = perm_admin
_db.permissions[perm_admin_uuid] = perm_admin
# Create auth:org:admin permission
perm_org_admin = Permission(
@@ -1189,7 +1071,7 @@ def bootstrap(
orgs={org_uuid: True}, # Grant to org
)
perm_org_admin.uuid = perm_org_admin_uuid
_db._data.permissions[perm_org_admin_uuid] = perm_org_admin
_db.permissions[perm_org_admin_uuid] = perm_org_admin
# Create organization
new_org = Org(
@@ -1197,7 +1079,7 @@ def bootstrap(
created_at=now,
)
new_org.uuid = org_uuid
_db._data.orgs[org_uuid] = new_org
_db.orgs[org_uuid] = new_org
# Create Administration role with both permissions
admin_role = Role(
@@ -1206,7 +1088,7 @@ def bootstrap(
permissions={perm_admin_uuid: True, perm_org_admin_uuid: True},
)
admin_role.uuid = role_uuid
_db._data.roles[role_uuid] = admin_role
_db.roles[role_uuid] = admin_role
# Create admin user
admin_user = User(
@@ -1217,10 +1099,10 @@ def bootstrap(
visits=0,
)
admin_user.uuid = user_uuid
_db._data.users[user_uuid] = admin_user
_db.users[user_uuid] = admin_user
# Create reset token
_db._data.reset_tokens[reset_key] = ResetToken(
_db.reset_tokens[reset_key] = ResetToken(
user=user_uuid,
expiry=reset_expiry,
token_type="admin bootstrap",
+16 -8
View File
@@ -199,17 +199,21 @@ class SessionContext(msgspec.Struct):
# -------------------------------------------------------------------------
class DatabaseData(msgspec.Struct, omit_defaults=True):
permissions: dict[UUID, Permission]
orgs: dict[UUID, Org]
roles: dict[UUID, Role]
users: dict[UUID, User]
credentials: dict[UUID, Credential]
sessions: dict[str, Session]
reset_tokens: dict[bytes, ResetToken]
class DB(msgspec.Struct, dict=True, omit_defaults=False):
"""In-memory database. Access fields directly for reads."""
permissions: dict[UUID, Permission] = {}
orgs: dict[UUID, Org] = {}
roles: dict[UUID, Role] = {}
users: dict[UUID, User] = {}
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)
self._store = None
# Set the key fields on all stored objects
for uuid, perm in self.permissions.items():
perm.uuid = uuid
@@ -225,3 +229,7 @@ class DatabaseData(msgspec.Struct, omit_defaults=True):
session.key = key
for key, token in self.reset_tokens.items():
token.key = key
def transaction(self, action, ctx=None, *, user=None):
"""Wrap writes in transaction. Delegates to JsonlStore."""
return self._store.transaction(action, ctx, user=user)
+30 -33
View File
@@ -12,12 +12,26 @@ Or via the CLI entry point (if installed):
"""
import asyncio
import re
from datetime import datetime, timezone
from uuid import UUID
import base64url
import uuid7
from sqlalchemy import select
from paskia.authsession import EXPIRES
from paskia.db.jsonl import JsonlStore
from paskia.db.structs import (
DB,
Credential,
Org,
Permission,
ResetToken,
Role,
Session,
User,
)
from .sql import (
DB as SQLDB,
@@ -47,30 +61,14 @@ async def migrate_from_sql(
sql_db_path: SQLAlchemy connection string for the source SQL database
json_db_path: Path for the destination JSONL file
"""
# Import here to avoid circular imports and to not require JSON db at import time
import re
import uuid7
from sqlalchemy import select
from paskia.db.operations import DB as JSONDB
from paskia.db.structs import (
Credential,
Org,
Permission,
ResetToken,
Role,
Session,
User,
)
# Initialize source SQL database
sql_db = SQLDB(sql_db_path)
await sql_db.init_db()
# Initialize destination JSON database (fresh, don't load existing)
json_db = JSONDB(json_db_path)
# Don't call json_db.load() - we want a fresh database, not to load existing
db = DB()
store = JsonlStore(db, json_db_path)
db._store = store
print(f"Migrating from {sql_db_path} to {json_db_path}...")
@@ -96,7 +94,7 @@ async def migrate_from_sql(
orgs={},
)
org_admin_perm.uuid = org_admin_perm_uuid
json_db._data.permissions[org_admin_perm_uuid] = org_admin_perm
db.permissions[org_admin_perm_uuid] = org_admin_perm
# Mapping from old permission ID to new permission UUID
perm_id_to_uuid: dict[str, UUID] = {}
@@ -121,7 +119,7 @@ async def migrate_from_sql(
orgs={},
)
new_perm.uuid = perm_uuid
json_db._data.permissions[perm_uuid] = new_perm
db.permissions[perm_uuid] = new_perm
perm_id_to_uuid[perm.id] = perm_uuid
print(
f" Migrated {len(permissions)} permissions (with {len(org_admin_uuids)} org-specific admins consolidated to auth:org:admin)"
@@ -133,14 +131,14 @@ async def migrate_from_sql(
org_key: UUID = org.uuid
new_org = Org(display_name=org.display_name)
new_org.uuid = org_key
json_db._data.orgs[org_key] = new_org
db.orgs[org_key] = new_org
# Update permissions to allow this org to grant them (by UUID)
for old_perm_id in org.permissions:
perm_uuid = perm_id_to_uuid.get(old_perm_id)
if perm_uuid and perm_uuid in json_db._data.permissions:
json_db._data.permissions[perm_uuid].orgs[org_key] = True
if perm_uuid and perm_uuid in db.permissions:
db.permissions[perm_uuid].orgs[org_key] = True
# Ensure every org can grant auth:org:admin
json_db._data.permissions[org_admin_perm_uuid].orgs[org_key] = True
db.permissions[org_admin_perm_uuid].orgs[org_key] = True
print(f" Migrated {len(orgs)} organizations")
# Migrate roles - convert old permission IDs to UUIDs
@@ -160,7 +158,7 @@ async def migrate_from_sql(
permissions=new_permissions,
)
new_role.uuid = role_key
json_db._data.roles[role_key] = new_role
db.roles[role_key] = new_role
role_count += 1
print(f" Migrated {role_count} roles")
@@ -179,7 +177,7 @@ async def migrate_from_sql(
visits=legacy_user.visits,
)
new_user.uuid = user_key
json_db._data.users[user_key] = new_user
db.users[user_key] = new_user
print(f" Migrated {len(user_models)} users")
# Migrate credentials
@@ -200,7 +198,7 @@ async def migrate_from_sql(
last_verified=legacy_cred.last_verified,
)
new_cred.uuid = cred_key
json_db._data.credentials[cred_key] = new_cred
db.credentials[cred_key] = new_cred
print(f" Migrated {len(cred_models)} credentials")
# Migrate sessions
@@ -217,7 +215,7 @@ async def migrate_from_sql(
else:
# Already in new format or unknown - try to use as-is
session_key = base64url.enc(old_key[:12])
json_db._data.sessions[session_key] = Session(
db.sessions[session_key] = Session(
user=sess.user_uuid,
credential=sess.credential_uuid,
host=sess.host,
@@ -241,7 +239,7 @@ async def migrate_from_sql(
else:
# Already in new format or unknown - truncate to 9 bytes
token_key = old_key[:9]
json_db._data.reset_tokens[token_key] = ResetToken(
db.reset_tokens[token_key] = ResetToken(
user=token.user_uuid,
expiry=token.expiry,
token_type=token.token_type,
@@ -249,11 +247,10 @@ async def migrate_from_sql(
print(f" Migrated {len(token_models)} reset tokens")
# Queue and flush all changes using the transaction mechanism
with json_db.transaction("migrate"):
with db.transaction("migrate"):
pass # All data already added to _data, transaction commits on exit
from paskia.db.jsonl import flush_changes
await flush_changes(json_db.db_path, json_db._pending_changes)
await store.flush()
print("Migration complete!")
+8 -7
View File
@@ -51,19 +51,20 @@ def event_loop():
@pytest_asyncio.fixture(scope="function")
async def test_db() -> AsyncGenerator[DB, None]:
"""Create an in-memory JSON database for testing.
Uses a temp file that gets cleaned up after each test.
"""
"""Create an in-memory JSON database for testing."""
import paskia.db.operations as ops_db
from paskia.db.jsonl import JsonlStore
with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f:
db = DB(f.name)
await db.load()
db = DB()
store = JsonlStore(db, f.name)
db._store = store
await store.load()
ops_db._db = db
ops_db._store = store
yield db
# Clean up
ops_db._db = None
ops_db._store = None
@pytest_asyncio.fixture(scope="function")