DB cleanup continued: Made the working copy data public in DB class.
This commit is contained in:
+12
-5
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user