Database refactor to separate modules.

This commit is contained in:
2026-01-23 18:27:12 +00:00
parent 2c6a5c72d9
commit f9d23a196c
19 changed files with 1653 additions and 1642 deletions
+141 -39
View File
@@ -1,21 +1,92 @@
""" """
Database module for WebAuthn passkey authentication. Database module for WebAuthn passkey authentication.
This module re-exports the JSONL database types and implementation. Read: Access _db._data directly, use build_* to convert to public structs.
All data types are msgspec Structs for efficient serialization. CTX: get_session_context(key) returns SessionContext with effective permissions.
Database methods are synchronous (no await needed). Write: Functions validate and commit, or raise ValueError.
Usage: Usage:
from paskia import db from paskia import db
# Access the database instance (after init) # Read (after init)
db.create_session(...) user_data = db._db._data.users[user_uuid]
user = db.get_user_by_uuid(uuid) user = db.build_user(user_uuid)
# Context
ctx = db.get_session_context(session_key)
# Write
db.create_user(user)
""" """
from paskia.db.json import ( from paskia.db.background import (
Credential, start_background,
start_cleanup,
stop_background,
stop_cleanup,
)
from paskia.db.operations import (
DB, DB,
_db,
add_permission_to_organization,
add_permission_to_role,
build_credential,
build_org,
build_permission,
build_reset_token,
build_role,
build_session,
build_user,
cleanup_expired,
create_credential,
create_credential_session,
create_organization,
create_permission,
create_reset_token,
create_role,
create_session,
create_user,
delete_credential,
delete_organization,
delete_permission,
delete_reset_token,
delete_role,
delete_session,
delete_sessions_for_user,
delete_user,
get_credential_by_id,
get_credentials_by_user_uuid,
get_organization,
get_organization_users,
get_permission,
get_permission_by_scope,
get_permission_organizations,
get_reset_token,
get_role,
get_roles_by_organization,
get_session,
get_session_context,
get_user_by_uuid,
get_user_organization,
init,
list_organizations,
list_permissions,
list_sessions_for_user,
login,
remove_permission_from_organization,
remove_permission_from_role,
rename_permission,
update_credential_sign_count,
update_organization_name,
update_permission,
update_role_name,
update_session,
update_user_display_name,
update_user_role,
update_user_role_in_organization,
)
from paskia.db.structs import (
Credential,
Org, Org,
Permission, Permission,
ResetToken, ResetToken,
@@ -23,40 +94,10 @@ from paskia.db.json import (
Session, Session,
SessionContext, SessionContext,
User, User,
init,
start_background,
stop_background,
start_cleanup,
stop_cleanup,
) )
import paskia.db.json as _json_module
class _DBProxy:
"""Proxy that forwards attribute access to the global DB instance.
This allows using `db.method()` directly instead of `db.get_db().method()`.
"""
def __getattr__(self, name: str):
db = _json_module._db
if db is None:
raise RuntimeError("Database not initialized. Call init() first.")
return getattr(db, name)
# Module-level proxy for direct access
_proxy = _DBProxy()
def __getattr__(name: str):
"""Module-level __getattr__ to forward DB method calls."""
if name in __all__:
raise AttributeError(name)
return getattr(_proxy, name)
__all__ = [ __all__ = [
# Types
"Credential", "Credential",
"DB", "DB",
"Org", "Org",
@@ -66,9 +107,70 @@ __all__ = [
"Session", "Session",
"SessionContext", "SessionContext",
"User", "User",
# Instance
"_db",
"init", "init",
# Background
"start_background", "start_background",
"stop_background", "stop_background",
"start_cleanup", "start_cleanup",
"stop_cleanup", "stop_cleanup",
# Builders
"build_credential",
"build_org",
"build_permission",
"build_reset_token",
"build_role",
"build_session",
"build_user",
# Read ops
"get_credential_by_id",
"get_credentials_by_user_uuid",
"get_organization",
"get_organization_users",
"get_permission",
"get_permission_by_scope",
"get_permission_organizations",
"get_reset_token",
"get_role",
"get_roles_by_organization",
"get_session",
"get_session_context",
"get_user_by_uuid",
"get_user_organization",
"list_organizations",
"list_permissions",
"list_sessions_for_user",
# Write ops
"add_permission_to_organization",
"add_permission_to_role",
"cleanup_expired",
"create_credential",
"create_credential_session",
"create_organization",
"create_permission",
"create_reset_token",
"create_role",
"create_session",
"create_user",
"delete_credential",
"delete_organization",
"delete_permission",
"delete_reset_token",
"delete_role",
"delete_session",
"delete_sessions_for_user",
"delete_user",
"login",
"remove_permission_from_organization",
"remove_permission_from_role",
"rename_permission",
"update_credential_sign_count",
"update_organization_name",
"update_permission",
"update_role_name",
"update_session",
"update_user_display_name",
"update_user_role",
"update_user_role_in_organization",
] ]
+106
View File
@@ -0,0 +1,106 @@
"""
Background task for database maintenance.
Periodically flushes pending changes to disk and cleans up expired items.
"""
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)
CLEANUP_INTERVAL = 1
_logger = logging.getLogger(__name__)
_background_task: asyncio.Task | None = None
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:
return
with _db.transaction("expiry"):
current_time = datetime.now(timezone.utc)
# Clean expired sessions
to_delete_sessions = [
k for k, s in _db._data.sessions.items() if s.expiry < current_time
]
for k in to_delete_sessions:
del _db._data.sessions[k]
# Clean expired reset tokens
to_delete_tokens = [
k for k, t in _db._data.reset_tokens.items() if t.expiry < current_time
]
for k in to_delete_tokens:
del _db._data.reset_tokens[k]
async def flush() -> None:
"""Write all pending database changes to disk."""
from paskia.db.operations import _db
if _db is None:
return
await flush_changes(_db.db_path, _db._pending_changes)
async def _background_loop():
"""Background task that periodically flushes changes and cleans up."""
# Run cleanup immediately on startup to clear old expired items
cleanup()
await flush()
last_cleanup = datetime.now(timezone.utc)
while True:
try:
await asyncio.sleep(FLUSH_INTERVAL)
# Flush pending changes to disk
await flush()
# Run cleanup less frequently
now = datetime.now(timezone.utc)
if (now - last_cleanup).total_seconds() >= CLEANUP_INTERVAL:
cleanup()
await flush() # Flush cleanup changes
last_cleanup = now
except asyncio.CancelledError:
# Final flush before exit
await flush()
break
except Exception:
_logger.exception("Error in database background loop")
async def start_background():
"""Start the background flush/cleanup task."""
global _background_task
if _background_task is None:
_background_task = asyncio.create_task(_background_loop())
async def stop_background():
"""Stop the background task and flush any pending changes."""
global _background_task
if _background_task:
_background_task.cancel()
try:
await _background_task
except asyncio.CancelledError:
pass
_background_task = None
# Aliases for backwards compatibility
start_cleanup = start_background
stop_cleanup = stop_background
-1517
View File
File diff suppressed because it is too large Load Diff
+125
View File
@@ -0,0 +1,125 @@
"""
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.
"""
import logging
from collections import deque
from datetime import datetime, timezone
from pathlib import Path
import aiofiles
import jsondiff
import msgspec
_logger = logging.getLogger(__name__)
# Default database path
DB_PATH_DEFAULT = "paskia.jsonl"
class _ChangeRecord(msgspec.Struct):
"""A single change record in the JSONL file."""
ts: datetime
actor: str
diff: dict
# msgspec encoder for change records
_change_encoder = msgspec.json.Encoder()
async def load_jsonl(db_path: Path, empty_data: dict) -> dict:
"""Load data from disk by applying change log.
Replays all changes from JSONL file using plain dicts (to handle
schema evolution).
Args:
db_path: Path to the JSONL database file
empty_data: Empty data structure to start with (as dict)
Returns:
The final state after applying all changes
"""
data_dict = empty_data.copy()
if db_path.exists():
try:
# Read entire file at once and split into lines
async with aiofiles.open(db_path, "rb") as f:
content = await f.read()
for line_num, line in enumerate(content.split(b"\n"), 1):
line = line.strip()
if not line:
continue
try:
change = msgspec.json.decode(line)
# Apply the diff to current state (marshal=True for $-prefixed keys)
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
except Exception as e:
raise ValueError(f"Error parsing line {line_num}: {e}")
except (OSError, ValueError, msgspec.DecodeError) as e:
raise ValueError(f"Failed to load database: {e}")
return data_dict
def compute_diff(previous: dict, current: dict) -> dict | None:
"""Compute JSON diff between two states.
Args:
previous: Previous state (JSON-compatible dict)
current: Current state (JSON-compatible dict)
Returns:
The diff, or None if no changes
"""
diff = jsondiff.diff(previous, current, marshal=True)
return diff if diff else None
def create_change_record(actor: str, diff: dict) -> _ChangeRecord:
"""Create a change record for persistence."""
return _ChangeRecord(
ts=datetime.now(timezone.utc),
actor=actor,
diff=diff,
)
async def flush_changes(
db_path: Path,
pending_changes: deque[_ChangeRecord],
) -> bool:
"""Write all pending changes to disk.
Args:
db_path: Path to the JSONL database file
pending_changes: Queue of pending change records (will be cleared on success)
Returns:
True if flush succeeded, False otherwise
"""
if not pending_changes:
return True
# Collect all pending changes
changes_to_write = list(pending_changes)
pending_changes.clear()
try:
# Build lines to append (keep as bytes, join with \n)
lines = [_change_encoder.encode(change) for change in changes_to_write]
# Append all lines in a single write (binary mode for Windows compatibility)
async with aiofiles.open(db_path, "ab") as f:
await f.write(b"\n".join(lines) + b"\n")
return True
except OSError:
_logger.exception("Failed to flush database changes")
# Re-queue the changes on failure
for change in reversed(changes_to_write):
pending_changes.appendleft(change)
return False
File diff suppressed because it is too large Load Diff
+168
View File
@@ -0,0 +1,168 @@
from datetime import datetime
from uuid import UUID
import msgspec
class Permission(msgspec.Struct, omit_defaults=True):
"""A permission that can be granted to roles."""
uuid: UUID # UUID primary key
scope: str # Permission scope identifier (e.g. "auth:admin", "myapp:write")
display_name: str
domain: str | None = None # If set, scopes permission to this domain
class Role(msgspec.Struct):
"""A role within an organization that can be assigned to users."""
uuid: UUID
org_uuid: UUID
display_name: str
permissions: list[str] = [] # permission IDs this role grants
class Org(msgspec.Struct):
"""An organization that contains users and roles."""
uuid: UUID
display_name: str
permissions: list[str] = [] # permission IDs this org can grant
roles: list[Role] = [] # roles belonging to this org
class User(msgspec.Struct):
"""A user in the authentication system."""
uuid: UUID
display_name: str
role_uuid: UUID
created_at: datetime | None = None
last_seen: datetime | None = None
visits: int = 0
class Credential(msgspec.Struct):
"""A WebAuthn credential (passkey) belonging to a user."""
uuid: UUID
credential_id: bytes # Long binary ID from the authenticator
user_uuid: UUID
aaguid: UUID
public_key: bytes
sign_count: int
created_at: datetime
last_used: datetime | None = None
last_verified: datetime | None = None
class Session(msgspec.Struct):
"""An active user session."""
key: bytes
user_uuid: UUID
credential_uuid: UUID
host: str | None
ip: str | None
user_agent: str | None
expiry: datetime
def metadata(self) -> dict:
"""Return session metadata for backwards compatibility."""
return {
"ip": self.ip,
"user_agent": self.user_agent,
"expiry": self.expiry.isoformat(),
}
# -------------------------------------------------------------------------
# Public data types (msgspec Structs)
# -------------------------------------------------------------------------
class ResetToken(msgspec.Struct):
"""A token for password reset or device addition."""
key: bytes
user_uuid: UUID
expiry: datetime
token_type: str
class SessionContext(msgspec.Struct):
"""Complete context for an authenticated session."""
session: Session
user: User
org: Org
role: Role
credential: Credential | None = None
permissions: list[Permission] | None = None
# -------------------------------------------------------------------------
# Internal storage types (different structure for efficient storage)
# -------------------------------------------------------------------------
class _PermissionData(msgspec.Struct, omit_defaults=True):
scope: str # Permission scope identifier
display_name: str
domain: str | None = None
orgs: dict[str, bool] = {} # org_uuid -> True (which orgs can grant this)
class _OrgData(msgspec.Struct):
display_name: str
created_at: datetime | None = None
class _RoleData(msgspec.Struct):
org: str
display_name: str
permissions: dict[str, bool] # permission_id -> True
class _UserData(msgspec.Struct):
display_name: str
role: str
created_at: datetime
last_seen: datetime | None
visits: int
class _CredentialData(msgspec.Struct):
credential_id: bytes # msgspec uses standard base64
user: str
aaguid: str
public_key: bytes # msgspec uses standard base64
sign_count: int
created_at: datetime
last_used: datetime | None
last_verified: datetime | None
class _SessionData(msgspec.Struct):
user: str
credential: str
host: str | None
ip: str | None
user_agent: str | None
expiry: datetime
class _ResetTokenData(msgspec.Struct):
user: str
expiry: datetime
token_type: str
class _DatabaseData(msgspec.Struct):
permissions: dict[str, _PermissionData]
orgs: dict[str, _OrgData]
roles: dict[str, _RoleData]
users: dict[str, _UserData]
credentials: dict[str, _CredentialData]
sessions: dict[str, _SessionData]
reset_tokens: dict[str, _ResetTokenData]
+3 -7
View File
@@ -356,7 +356,7 @@ async def admin_add_role_permission(
db.get_permission(permission_id) db.get_permission(permission_id)
org = db.get_organization(str(org_uuid)) org = db.get_organization(str(org_uuid))
if permission_id not in org.permissions: if permission_id not in org.permissions:
raise ValueError(f"Permission not grantable by organization") raise ValueError("Permission not grantable by organization")
db.add_permission_to_role(role_uuid, permission_id) db.add_permission_to_role(role_uuid, permission_id)
return {"status": "ok"} return {"status": "ok"}
@@ -593,14 +593,10 @@ async def admin_get_user_detail(
status_code=403, detail="Insufficient permissions", mode="forbidden" status_code=403, detail="Insufficient permissions", mode="forbidden"
) )
user = db.get_user_by_uuid(user_uuid) user = db.get_user_by_uuid(user_uuid)
cred_ids = db.get_credentials_by_user_uuid(user_uuid) user_creds = db.get_credentials_by_user_uuid(user_uuid)
creds: list[dict] = [] creds: list[dict] = []
aaguids: set[str] = set() aaguids: set[str] = set()
for cid in cred_ids: for c in user_creds:
try:
c = db.get_credential_by_id(cid)
except ValueError: # pragma: no cover - race condition handling
continue
aaguid_str = str(c.aaguid) aaguid_str = str(c.aaguid)
aaguids.add(aaguid_str) aaguids.add(aaguid_str)
creds.append( creds.append(
+1 -1
View File
@@ -13,6 +13,7 @@ from fastapi import (
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from fastapi.security import HTTPBearer from fastapi.security import HTTPBearer
from paskia import db
from paskia.authsession import ( from paskia.authsession import (
EXPIRES, EXPIRES,
get_reset, get_reset,
@@ -21,7 +22,6 @@ from paskia.authsession import (
) )
from paskia.fastapi import authz, session, user from paskia.fastapi import authz, session, user
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME
from paskia import db
from paskia.globals import passkey as global_passkey from paskia.globals import passkey as global_passkey
from paskia.util import frontend, hostutil, htmlutil, passphrase, userinfo from paskia.util import frontend, hostutil, htmlutil, passphrase, userinfo
from paskia.util.tokens import session_key from paskia.util.tokens import session_key
+1 -2
View File
@@ -15,11 +15,10 @@ from uuid import UUID
import base64url import base64url
from fastapi import FastAPI, WebSocket, WebSocketDisconnect from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from paskia import remoteauth from paskia import db, remoteauth
from paskia.authsession import create_session from paskia.authsession import create_session
from paskia.fastapi.session import infodict from paskia.fastapi.session import infodict
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
from paskia import db
from paskia.globals import passkey from paskia.globals import passkey
from paskia.util import passphrase, pow from paskia.util import passphrase, pow
+8 -8
View File
@@ -27,9 +27,9 @@ async def _resolve_targets(query: str | None):
targets: list[tuple] = [] targets: list[tuple] = []
try: try:
q_uuid = UUID(query) q_uuid = UUID(query)
perm_orgs = await _db.get_permission_organizations("auth:admin") perm_orgs = _db.get_permission_organizations("auth:admin")
for o in perm_orgs: for o in perm_orgs:
users = await _db.get_organization_users(str(o.uuid)) users = _db.get_organization_users(str(o.uuid))
for u, role_name in users: for u, role_name in users:
if u.uuid == q_uuid: if u.uuid == q_uuid:
return [(u, role_name)] return [(u, role_name)]
@@ -38,9 +38,9 @@ async def _resolve_targets(query: str | None):
pass pass
# Substring search # Substring search
needle = query.lower() needle = query.lower()
perm_orgs = await _db.get_permission_organizations("auth:admin") perm_orgs = _db.get_permission_organizations("auth:admin")
for o in perm_orgs: for o in perm_orgs:
users = await _db.get_organization_users(str(o.uuid)) users = _db.get_organization_users(str(o.uuid))
for u, role_name in users: for u, role_name in users:
if needle in (u.display_name or "").lower(): if needle in (u.display_name or "").lower():
targets.append((u, role_name)) targets.append((u, role_name))
@@ -53,10 +53,10 @@ async def _resolve_targets(query: str | None):
deduped.append((u, role_name)) deduped.append((u, role_name))
return deduped return deduped
# No query -> master admin # No query -> master admin
perm_orgs = await _db.get_permission_organizations("auth:admin") perm_orgs = _db.get_permission_organizations("auth:admin")
if not perm_orgs: if not perm_orgs:
return [] return []
users = await _db.get_organization_users(str(perm_orgs[0].uuid)) users = _db.get_organization_users(str(perm_orgs[0].uuid))
admin_users = [pair for pair in users if pair[1] == "Administration"] admin_users = [pair for pair in users if pair[1] == "Administration"]
return admin_users[:1] return admin_users[:1]
@@ -64,9 +64,9 @@ async def _resolve_targets(query: str | None):
async def _create_reset(user, role_name: str): async def _create_reset(user, role_name: str):
token = passphrase.generate() token = passphrase.generate()
expiry = _authsession.reset_expires() expiry = _authsession.reset_expires()
await _db.create_reset_token( _db.create_reset_token(
user_uuid=user.uuid,
key=_tokens.reset_key(token), key=_tokens.reset_key(token),
user_uuid=user.uuid,
expiry=expiry, expiry=expiry,
token_type="manual reset", token_type="manual reset",
) )
+1 -1
View File
@@ -10,6 +10,7 @@ from fastapi import (
) )
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from paskia import db
from paskia.authsession import ( from paskia.authsession import (
delete_credential, delete_credential,
expires, expires,
@@ -17,7 +18,6 @@ from paskia.authsession import (
) )
from paskia.fastapi import authz, session from paskia.fastapi import authz, session
from paskia.fastapi.session import AUTH_COOKIE from paskia.fastapi.session import AUTH_COOKIE
from paskia import db
from paskia.util import hostutil, passphrase, tokens from paskia.util import hostutil, passphrase, tokens
from paskia.util.tokens import decode_session_key, session_key from paskia.util.tokens import decode_session_key, session_key
+1 -1
View File
@@ -2,11 +2,11 @@ from uuid import UUID
from fastapi import FastAPI, WebSocket from fastapi import FastAPI, WebSocket
from paskia import db
from paskia.authsession import create_session, get_reset, get_session from paskia.authsession import create_session, get_reset, get_session
from paskia.fastapi import authz, remote from paskia.fastapi import authz, remote
from paskia.fastapi.session import AUTH_COOKIE, infodict from paskia.fastapi.session import AUTH_COOKIE, infodict
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
from paskia import db
from paskia.globals import passkey from paskia.globals import passkey
from paskia.util import passphrase from paskia.util import passphrase
from paskia.util.tokens import create_token, session_key from paskia.util.tokens import create_token, session_key
+3 -5
View File
@@ -42,8 +42,7 @@ async def init(
Set PASKIA_DB environment variable to specify the JSONL database file path. Set PASKIA_DB environment variable to specify the JSONL database file path.
Default: paskia.jsonl Default: paskia.jsonl
""" """
from . import remoteauth from . import db, remoteauth
from .db import json as json_db
# Initialize passkey instance with provided parameters # Initialize passkey instance with provided parameters
passkey.instance = Passkey( passkey.instance = Passkey(
@@ -52,9 +51,8 @@ async def init(
origins=origins, origins=origins,
) )
# Initialize database if not already done # Initialize database
if json_db._db is None: await db.init()
await json_db.init()
# Initialize remote auth manager # Initialize remote auth manager
await remoteauth.init() await remoteauth.init()
+2 -4
View File
@@ -59,10 +59,8 @@ async def migrate_from_sql(
import uuid7 import uuid7
from sqlalchemy import select from sqlalchemy import select
from paskia.db.json import ( from paskia.db.operations import DB as JSONDB
DB as JSONDB, from paskia.db.structs import (
)
from paskia.db.json import (
_CredentialData, _CredentialData,
_OrgData, _OrgData,
_PermissionData, _PermissionData,
+1 -1
View File
@@ -4,7 +4,7 @@ Legacy SQL database implementation for migration purposes.
This module provides the async SQLAlchemy database layer that was used This module provides the async SQLAlchemy database layer that was used
before the JSONL format. It is kept here for migration purposes only. before the JSONL format. It is kept here for migration purposes only.
DO NOT use this module for new code. Use paskia.db.json instead. DO NOT use this module for new code. Use paskia.db instead.
""" """
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
+2 -7
View File
@@ -44,16 +44,11 @@ async def format_user_info(
ctx = await permutil.session_context(auth, request_host) ctx = await permutil.session_context(auth, request_host)
# Fetch and format credentials # Fetch and format credentials
credential_ids = db.get_credentials_by_user_uuid(user_uuid) user_credentials = db.get_credentials_by_user_uuid(user_uuid)
credentials: list[dict] = [] credentials: list[dict] = []
user_aaguids: set[str] = set() user_aaguids: set[str] = set()
for cred_id in credential_ids: for c in user_credentials:
try:
c = db.get_credential_by_id(cred_id)
except ValueError:
continue
aaguid_str = str(c.aaguid) aaguid_str = str(c.aaguid)
user_aaguids.add(aaguid_str) user_aaguids.add(aaguid_str)
credentials.append( credentials.append(
+34 -21
View File
@@ -18,14 +18,27 @@ from uuid import UUID
import httpx import httpx
import pytest import pytest
from paskia.authsession import expires
import pytest_asyncio import pytest_asyncio
import uuid7 import uuid7
from paskia import globals as paskia_globals from paskia import globals as paskia_globals
from paskia.db import Credential, Org, Permission, Role, User from paskia.authsession import expires
from paskia.db.json import DB from paskia.db import (
Credential,
Org,
Permission,
Role,
User,
add_permission_to_organization,
create_credential,
create_organization,
create_permission,
create_reset_token,
create_role,
create_session,
create_user,
)
from paskia.db.operations import DB
from paskia.fastapi.session import AUTH_COOKIE_NAME from paskia.fastapi.session import AUTH_COOKIE_NAME
from paskia.sansio import Passkey from paskia.sansio import Passkey
from paskia.util.tokens import create_token, session_key from paskia.util.tokens import create_token, session_key
@@ -45,15 +58,15 @@ async def test_db() -> AsyncGenerator[DB, None]:
Uses a temp file that gets cleaned up after each test. Uses a temp file that gets cleaned up after each test.
""" """
import paskia.db.json as json_db import paskia.db.operations as ops_db
with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f: with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f:
db = DB(f.name) db = DB(f.name)
db.load() # Synchronous now await db.load()
json_db._db = db ops_db._db = db
yield db yield db
# Clean up # Clean up
json_db._db = None ops_db._db = None
@pytest_asyncio.fixture(scope="function") @pytest_asyncio.fixture(scope="function")
@@ -77,7 +90,7 @@ async def test_org(test_db: DB, admin_permission: Permission) -> Org:
display_name="Test Organization", display_name="Test Organization",
permissions=["auth:admin"], # Org can grant this permission permissions=["auth:admin"], # Org can grant this permission
) )
test_db.create_organization(org) create_organization(org)
return org return org
@@ -89,7 +102,7 @@ async def admin_permission(test_db: DB) -> Permission:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="auth:admin", display_name="Master Admin" uuid=uuid7.create(), scope="auth:admin", display_name="Master Admin"
) )
test_db.create_permission(perm) create_permission(perm)
return perm return perm
@@ -101,9 +114,9 @@ async def org_admin_permission(test_db: DB, test_org: Org) -> Permission:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="auth:org:admin", display_name="Organization Admin" uuid=uuid7.create(), scope="auth:org:admin", display_name="Organization Admin"
) )
test_db.create_permission(perm) create_permission(perm)
# Make it grantable by the org # Make it grantable by the org
test_db.add_permission_to_organization(str(test_org.uuid), "auth:org:admin") add_permission_to_organization(str(test_org.uuid), "auth:org:admin")
return perm return perm
@@ -121,7 +134,7 @@ async def test_role(
display_name="Test Admin Role", display_name="Test Admin Role",
permissions=["auth:admin", "auth:org:admin"], permissions=["auth:admin", "auth:org:admin"],
) )
test_db.create_role(role) create_role(role)
return role return role
@@ -134,7 +147,7 @@ async def user_role(test_db: DB, test_org: Org) -> Role:
display_name="User Role", display_name="User Role",
permissions=[], permissions=[],
) )
test_db.create_role(role) create_role(role)
return role return role
@@ -148,7 +161,7 @@ async def test_user(test_db: DB, test_role: Role) -> User:
created_at=datetime.now(timezone.utc), created_at=datetime.now(timezone.utc),
visits=0, visits=0,
) )
test_db.create_user(user) create_user(user)
return user return user
@@ -162,7 +175,7 @@ async def regular_user(test_db: DB, user_role: Role) -> User:
created_at=datetime.now(timezone.utc), created_at=datetime.now(timezone.utc),
visits=0, visits=0,
) )
test_db.create_user(user) create_user(user)
return user return user
@@ -180,7 +193,7 @@ async def test_credential(test_db: DB, test_user: User) -> Credential:
last_used=None, last_used=None,
last_verified=None, last_verified=None,
) )
test_db.create_credential(credential) create_credential(credential)
return credential return credential
@@ -198,7 +211,7 @@ async def regular_credential(test_db: DB, regular_user: User) -> Credential:
last_used=None, last_used=None,
last_verified=None, last_verified=None,
) )
test_db.create_credential(credential) create_credential(credential)
return credential return credential
@@ -208,7 +221,7 @@ async def session_token(
) -> str: ) -> str:
"""Create a session for the admin user and return the token.""" """Create a session for the admin user and return the token."""
token = create_token() token = create_token()
test_db.create_session( create_session(
user_uuid=test_user.uuid, user_uuid=test_user.uuid,
credential_uuid=test_credential.uuid, credential_uuid=test_credential.uuid,
key=session_key(token), key=session_key(token),
@@ -226,7 +239,7 @@ async def regular_session_token(
) -> str: ) -> str:
"""Create a session for a regular user and return the token.""" """Create a session for a regular user and return the token."""
token = create_token() token = create_token()
test_db.create_session( create_session(
user_uuid=regular_user.uuid, user_uuid=regular_user.uuid,
credential_uuid=regular_credential.uuid, credential_uuid=regular_credential.uuid,
key=session_key(token), key=session_key(token),
@@ -246,7 +259,7 @@ async def reset_token(test_db: DB, test_user: User, test_credential: Credential)
from paskia.util.tokens import reset_key from paskia.util.tokens import reset_key
token = generate() token = generate()
test_db.create_reset_token( create_reset_token(
user_uuid=test_user.uuid, user_uuid=test_user.uuid,
key=reset_key(token), key=reset_key(token),
expiry=reset_expires(), expiry=reset_expires(),
+39 -26
View File
@@ -20,8 +20,21 @@ import pytest_asyncio
import uuid7 import uuid7
from paskia.authsession import expires from paskia.authsession import expires
from paskia.db import Credential, Org, Permission, Role, User from paskia.db import (
from paskia.db.json import DB Credential,
Org,
Permission,
Role,
User,
add_permission_to_organization,
create_credential,
create_organization,
create_permission,
create_role,
create_session,
create_user,
)
from paskia.db.operations import DB
from paskia.util.tokens import create_token, encode_session_key, session_key from paskia.util.tokens import create_token, encode_session_key, session_key
from tests.conftest import auth_headers from tests.conftest import auth_headers
@@ -36,7 +49,7 @@ async def second_org(test_db: DB) -> Org:
display_name="Second Organization", display_name="Second Organization",
permissions=[], permissions=[],
) )
test_db.create_organization(org) create_organization(org)
return org return org
@@ -51,7 +64,7 @@ async def second_org_role(
display_name="Second Org Admin Role", display_name="Second Org Admin Role",
permissions=["auth:admin"], permissions=["auth:admin"],
) )
test_db.create_role(role) create_role(role)
return role return role
@@ -65,7 +78,7 @@ async def second_org_user(test_db: DB, second_org_role: Role) -> User:
created_at=datetime.now(timezone.utc), created_at=datetime.now(timezone.utc),
visits=0, visits=0,
) )
test_db.create_user(user) create_user(user)
return user return user
@@ -85,7 +98,7 @@ async def second_org_credential(test_db: DB, second_org_user: User) -> Credentia
last_used=datetime.now(timezone.utc), last_used=datetime.now(timezone.utc),
last_verified=datetime.now(timezone.utc), last_verified=datetime.now(timezone.utc),
) )
test_db.create_credential(credential) create_credential(credential)
return credential return credential
@@ -95,7 +108,7 @@ async def second_org_session_token(
) -> str: ) -> str:
"""Create a session for the second org admin user.""" """Create a session for the second org admin user."""
token = create_token() token = create_token()
test_db.create_session( create_session(
user_uuid=second_org_user.uuid, user_uuid=second_org_user.uuid,
credential_uuid=second_org_credential.uuid, credential_uuid=second_org_credential.uuid,
key=session_key(token), key=session_key(token),
@@ -116,7 +129,7 @@ async def org_admin_role(test_db: DB, test_org: Org, org_admin_permission) -> Ro
display_name="Org Admin Role", display_name="Org Admin Role",
permissions=["auth:org:admin"], permissions=["auth:org:admin"],
) )
test_db.create_role(role) create_role(role)
return role return role
@@ -131,7 +144,7 @@ async def org_admin_user(test_db: DB, org_admin_role: Role) -> User:
visits=5, visits=5,
last_seen=datetime.now(timezone.utc), last_seen=datetime.now(timezone.utc),
) )
test_db.create_user(user) create_user(user)
return user return user
@@ -151,7 +164,7 @@ async def org_admin_credential(test_db: DB, org_admin_user: User) -> Credential:
last_used=datetime.now(timezone.utc), last_used=datetime.now(timezone.utc),
last_verified=None, last_verified=None,
) )
test_db.create_credential(credential) create_credential(credential)
return credential return credential
@@ -161,7 +174,7 @@ async def org_admin_session_token(
) -> str: ) -> str:
"""Create a session for the org admin user.""" """Create a session for the org admin user."""
token = create_token() token = create_token()
test_db.create_session( create_session(
user_uuid=org_admin_user.uuid, user_uuid=org_admin_user.uuid,
credential_uuid=org_admin_credential.uuid, credential_uuid=org_admin_credential.uuid,
key=session_key(token), key=session_key(token),
@@ -181,9 +194,9 @@ async def grantable_permission(test_db: DB, test_org: Org) -> Permission:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="test:grantable:perm", display_name="Grantable Perm" uuid=uuid7.create(), scope="test:grantable:perm", display_name="Grantable Perm"
) )
test_db.create_permission(perm) create_permission(perm)
# Add to org's grantable permissions # Add to org's grantable permissions
test_db.add_permission_to_organization(str(test_org.uuid), perm.scope) add_permission_to_organization(str(test_org.uuid), perm.scope)
return perm return perm
@@ -414,7 +427,7 @@ class TestAdminOrganizations:
display_name="Org To Delete", display_name="Org To Delete",
permissions=[], permissions=[],
) )
test_db.create_organization(org_to_delete) create_organization(org_to_delete)
# Create some org-specific permissions to test cleanup # Create some org-specific permissions to test cleanup
org_perm = Permission( org_perm = Permission(
@@ -422,7 +435,7 @@ class TestAdminOrganizations:
scope=f"test:org:{org_to_delete.uuid}:feature", scope=f"test:org:{org_to_delete.uuid}:feature",
display_name="Org Feature", display_name="Org Feature",
) )
test_db.create_permission(org_perm) create_permission(org_perm)
response = await client.delete( response = await client.delete(
f"/auth/api/admin/orgs/{org_to_delete.uuid}", f"/auth/api/admin/orgs/{org_to_delete.uuid}",
@@ -601,7 +614,7 @@ class TestAdminRoles:
scope="test:not:grantable", scope="test:not:grantable",
display_name="Not Grantable", display_name="Not Grantable",
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.post( response = await client.post(
f"/auth/api/admin/orgs/{test_org.uuid}/roles", f"/auth/api/admin/orgs/{test_org.uuid}/roles",
@@ -676,7 +689,7 @@ class TestAdminRoles:
scope="test:not:grantable:update", scope="test:not:grantable:update",
display_name="Not Grantable", display_name="Not Grantable",
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.post( response = await client.post(
f"/auth/api/admin/orgs/{test_org.uuid}/roles/{user_role.uuid}/permissions/test:not:grantable:update", f"/auth/api/admin/orgs/{test_org.uuid}/roles/{user_role.uuid}/permissions/test:not:grantable:update",
@@ -1097,7 +1110,7 @@ class TestAdminUsersInOrg:
created_at=datetime.now(timezone.utc), created_at=datetime.now(timezone.utc),
visits=0, visits=0,
) )
test_db.create_user(user_no_cred) create_user(user_no_cred)
response = await client.post( response = await client.post(
f"/auth/api/admin/orgs/{test_org.uuid}/users/{user_no_cred.uuid}/create-link", f"/auth/api/admin/orgs/{test_org.uuid}/users/{user_no_cred.uuid}/create-link",
@@ -1184,7 +1197,7 @@ class TestAdminSessions:
# Create an additional session to delete # Create an additional session to delete
extra_token = create_token() extra_token = create_token()
extra_key = session_key(extra_token) extra_key = session_key(extra_token)
test_db.create_session( create_session(
user_uuid=test_user.uuid, user_uuid=test_user.uuid,
credential_uuid=test_credential.uuid, credential_uuid=test_credential.uuid,
key=extra_key, key=extra_key,
@@ -1380,7 +1393,7 @@ class TestAdminPermissions:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="test:updateable", display_name="Updateable" uuid=uuid7.create(), scope="test:updateable", display_name="Updateable"
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.patch( response = await client.patch(
"/auth/api/admin/permission?permission_id=test:updateable&display_name=Updated%20Name", "/auth/api/admin/permission?permission_id=test:updateable&display_name=Updated%20Name",
@@ -1401,7 +1414,7 @@ class TestAdminPermissions:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="test:perm", display_name="Test Perm" uuid=uuid7.create(), scope="test:perm", display_name="Test Perm"
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.patch( response = await client.patch(
"/auth/api/admin/permission?permission_id=test:perm&display_name=", "/auth/api/admin/permission?permission_id=test:perm&display_name=",
@@ -1422,7 +1435,7 @@ class TestAdminPermissions:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="test:renameable2", display_name="Renameable" uuid=uuid7.create(), scope="test:renameable2", display_name="Renameable"
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.post( response = await client.post(
"/auth/api/admin/permission/rename", "/auth/api/admin/permission/rename",
@@ -1469,7 +1482,7 @@ class TestAdminPermissions:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="test:rename:withname", display_name="Old Name" uuid=uuid7.create(), scope="test:rename:withname", display_name="Old Name"
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.post( response = await client.post(
"/auth/api/admin/permission/rename", "/auth/api/admin/permission/rename",
@@ -1493,7 +1506,7 @@ class TestAdminPermissions:
perm = Permission( perm = Permission(
uuid=uuid7.create(), scope="test:deleteable", display_name="Deleteable" uuid=uuid7.create(), scope="test:deleteable", display_name="Deleteable"
) )
test_db.create_permission(perm) create_permission(perm)
response = await client.delete( response = await client.delete(
"/auth/api/admin/permission?permission_id=test:deleteable", "/auth/api/admin/permission?permission_id=test:deleteable",
@@ -1529,7 +1542,7 @@ class TestAdminPermissions:
perm2 = Permission( perm2 = Permission(
uuid=uuid7.create(), scope="auth:admin", display_name="Secondary Admin" uuid=uuid7.create(), scope="auth:admin", display_name="Secondary Admin"
) )
test_db.create_permission(perm2) create_permission(perm2)
# Now we can delete the original one # Now we can delete the original one
response = await client.delete( response = await client.delete(
@@ -1556,7 +1569,7 @@ class TestAdminPermissions:
display_name="Other Domain Admin", display_name="Other Domain Admin",
domain="other.example.com", domain="other.example.com",
) )
test_db.create_permission(perm2) create_permission(perm2)
# Cannot delete the original one because the remaining one is not accessible # Cannot delete the original one because the remaining one is not accessible
response = await client.delete( response = await client.delete(
+3 -2
View File
@@ -15,6 +15,7 @@ from datetime import datetime, timezone
import httpx import httpx
import pytest import pytest
from paskia.db import create_session, delete_session
from tests.conftest import auth_headers from tests.conftest import auth_headers
@@ -526,7 +527,7 @@ class TestValidateSessionRefresh:
# Create a session with an old expiry time to trigger refresh # Create a session with an old expiry time to trigger refresh
token = create_token() token = create_token()
old_expiry = datetime.now(timezone.utc) + EXPIRES - timedelta(minutes=10) old_expiry = datetime.now(timezone.utc) + EXPIRES - timedelta(minutes=10)
test_db.create_session( create_session(
user_uuid=test_user.uuid, user_uuid=test_user.uuid,
credential_uuid=test_credential.uuid, credential_uuid=test_credential.uuid,
key=session_key(token), key=session_key(token),
@@ -537,7 +538,7 @@ class TestValidateSessionRefresh:
) )
# Delete the session right before validate tries to refresh # Delete the session right before validate tries to refresh
test_db.delete_session(session_key(token)) delete_session(session_key(token))
response = await client.post( response = await client.post(
"/auth/api/validate", "/auth/api/validate",