Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7c4418e631 | ||
|
|
1648c8641f | ||
|
|
63eb088dbd | ||
|
|
e88cc004dd | ||
|
|
76921e8b31 | ||
|
|
c1b0aab296 | ||
|
|
8f89bb6d4b | ||
|
|
d16d1ed1c2 | ||
|
|
2cfca81672 | ||
|
|
ce300ebdaf | ||
|
|
7329223784 | ||
|
|
c8d659b5ca |
@@ -10,6 +10,7 @@ This module provides functionality to:
|
|||||||
import json
|
import json
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
from importlib.resources import files
|
from importlib.resources import files
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
__ALL__ = ["AAGUID", "filter"]
|
__ALL__ = ["AAGUID", "filter"]
|
||||||
|
|
||||||
@@ -18,15 +19,15 @@ AAGUID_FILE = files("paskia") / "aaguid" / "combined_aaguid.json"
|
|||||||
AAGUID: dict[str, dict] = json.loads(AAGUID_FILE.read_text(encoding="utf-8"))
|
AAGUID: dict[str, dict] = json.loads(AAGUID_FILE.read_text(encoding="utf-8"))
|
||||||
|
|
||||||
|
|
||||||
def filter(aaguids: Iterable[str]) -> dict[str, dict]:
|
def filter(aaguids: Iterable[UUID]) -> dict[str, dict]:
|
||||||
"""
|
"""
|
||||||
Get AAGUID information only for the provided set of AAGUIDs.
|
Get AAGUID information only for the provided set of AAGUIDs.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
aaguids: Set of AAGUID strings that the user has credentials for
|
aaguids: Iterable of AAGUIDs (UUIDs) that the user has credentials for
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary mapping AAGUID to authenticator information for only
|
Dictionary mapping AAGUID string to authenticator information for only
|
||||||
the AAGUIDs that the user has and that we have data for
|
the AAGUIDs that the user has and that we have data for
|
||||||
"""
|
"""
|
||||||
return {aaguid: AAGUID[aaguid] for aaguid in aaguids if aaguid in AAGUID}
|
return {(s := str(a)): AAGUID[s] for a in aaguids if (s := str(a)) in AAGUID}
|
||||||
|
|||||||
+4
-19
@@ -8,7 +8,7 @@ independent of any web framework:
|
|||||||
- Credential management
|
- Credential management
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
@@ -23,11 +23,11 @@ EXPIRES = SESSION_LIFETIME
|
|||||||
|
|
||||||
|
|
||||||
def expires() -> datetime:
|
def expires() -> datetime:
|
||||||
return datetime.now(timezone.utc) + EXPIRES
|
return datetime.now(UTC) + EXPIRES
|
||||||
|
|
||||||
|
|
||||||
def reset_expires() -> datetime:
|
def reset_expires() -> datetime:
|
||||||
return datetime.now(timezone.utc) + RESET_LIFETIME
|
return datetime.now(UTC) + RESET_LIFETIME
|
||||||
|
|
||||||
|
|
||||||
def get_reset(token: str) -> "ResetToken":
|
def get_reset(token: str) -> "ResetToken":
|
||||||
@@ -39,24 +39,9 @@ def get_reset(token: str) -> "ResetToken":
|
|||||||
raise ValueError("This authentication link is no longer valid.")
|
raise ValueError("This authentication link is no longer valid.")
|
||||||
|
|
||||||
|
|
||||||
def refresh_session_token(token: str, *, ip: str, user_agent: str):
|
|
||||||
"""Refresh a session extending its expiry."""
|
|
||||||
session_record = db.data().sessions.get(token)
|
|
||||||
if not session_record:
|
|
||||||
raise ValueError("Session not found or expired")
|
|
||||||
updated = db.update_session(
|
|
||||||
token,
|
|
||||||
ip=ip,
|
|
||||||
user_agent=user_agent,
|
|
||||||
expiry=expires(),
|
|
||||||
)
|
|
||||||
if not updated:
|
|
||||||
raise ValueError("Session not found or expired")
|
|
||||||
|
|
||||||
|
|
||||||
def delete_credential(credential_uuid: UUID, auth: str, host: str | None = None):
|
def delete_credential(credential_uuid: UUID, auth: str, host: str | None = None):
|
||||||
"""Delete a specific credential for the current user."""
|
"""Delete a specific credential for the current user."""
|
||||||
ctx = db.get_session_context(auth, hostutil.normalize_host(host))
|
ctx = db.data().session_ctx(auth, hostutil.normalize_host(host))
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise ValueError("Session expired")
|
raise ValueError("Session expired")
|
||||||
db.delete_credential(credential_uuid, ctx.user.uuid)
|
db.delete_credential(credential_uuid, ctx.user.uuid)
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
Database module for WebAuthn passkey authentication.
|
Database module for WebAuthn passkey authentication.
|
||||||
|
|
||||||
Read: Access data() directly, use build_* to convert to public structs.
|
Read: Access data() directly, use build_* to convert to public structs.
|
||||||
CTX: get_session_context(key) returns SessionContext with effective permissions.
|
CTX: data().session_ctx(key) returns SessionContext with effective permissions.
|
||||||
Write: Functions validate and commit, or raise ValueError.
|
Write: Functions validate and commit, or raise ValueError.
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
@@ -13,7 +13,7 @@ Usage:
|
|||||||
user = db.build_user(user_uuid)
|
user = db.build_user(user_uuid)
|
||||||
|
|
||||||
# Context
|
# Context
|
||||||
ctx = db.get_session_context(session_key)
|
ctx = db.data().session_ctx(session_key)
|
||||||
|
|
||||||
# Write
|
# Write
|
||||||
db.create_user(user)
|
db.create_user(user)
|
||||||
@@ -49,7 +49,6 @@ from paskia.db.operations import (
|
|||||||
delete_user,
|
delete_user,
|
||||||
get_organization_users,
|
get_organization_users,
|
||||||
get_reset_token,
|
get_reset_token,
|
||||||
get_session_context,
|
|
||||||
get_user_credential_ids,
|
get_user_credential_ids,
|
||||||
get_user_organization,
|
get_user_organization,
|
||||||
init,
|
init,
|
||||||
@@ -113,7 +112,6 @@ __all__ = [
|
|||||||
# Read ops
|
# Read ops
|
||||||
"get_organization_users",
|
"get_organization_users",
|
||||||
"get_reset_token",
|
"get_reset_token",
|
||||||
"get_session_context",
|
|
||||||
"get_user_credential_ids",
|
"get_user_credential_ids",
|
||||||
"get_user_organization",
|
"get_user_organization",
|
||||||
# Write ops
|
# Write ops
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ Periodically flushes pending changes to disk and cleans up expired items.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from paskia.db.operations import _store, cleanup_expired
|
from paskia.db.operations import _store, cleanup_expired
|
||||||
|
|
||||||
@@ -33,7 +33,7 @@ async def _background_loop():
|
|||||||
cleanup_expired()
|
cleanup_expired()
|
||||||
await flush()
|
await flush()
|
||||||
|
|
||||||
last_cleanup = datetime.now(timezone.utc)
|
last_cleanup = datetime.now(UTC)
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
@@ -42,7 +42,7 @@ async def _background_loop():
|
|||||||
await flush()
|
await flush()
|
||||||
|
|
||||||
# Run cleanup periodically
|
# Run cleanup periodically
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if (now - last_cleanup).total_seconds() >= CLEANUP_INTERVAL:
|
if (now - last_cleanup).total_seconds() >= CLEANUP_INTERVAL:
|
||||||
cleanup_expired()
|
cleanup_expired()
|
||||||
await flush() # Flush cleanup changes
|
await flush() # Flush cleanup changes
|
||||||
|
|||||||
+100
-112
@@ -2,15 +2,11 @@
|
|||||||
JSONL persistence layer for the database.
|
JSONL persistence layer for the database.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import copy
|
import copy
|
||||||
import json
|
|
||||||
import logging
|
import logging
|
||||||
import sys
|
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
@@ -19,7 +15,8 @@ import aiofiles
|
|||||||
import jsondiff
|
import jsondiff
|
||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
from paskia.db.migrations import apply_migrations
|
from paskia.db.logging import log_change
|
||||||
|
from paskia.db.migrations import DBVER, apply_all_migrations
|
||||||
from paskia.db.structs import DB, SessionContext
|
from paskia.db.structs import DB, SessionContext
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
@@ -33,6 +30,7 @@ class _ChangeRecord(msgspec.Struct, omit_defaults=True):
|
|||||||
|
|
||||||
ts: datetime
|
ts: datetime
|
||||||
a: str # action - describes the operation (e.g., "migrate", "login", "create_user")
|
a: str # action - describes the operation (e.g., "migrate", "login", "create_user")
|
||||||
|
v: int # schema version after this change
|
||||||
u: str | None = None # user UUID who performed the action (None for system)
|
u: str | None = None # user UUID who performed the action (None for system)
|
||||||
diff: dict = {}
|
diff: dict = {}
|
||||||
|
|
||||||
@@ -41,43 +39,6 @@ class _ChangeRecord(msgspec.Struct, omit_defaults=True):
|
|||||||
_change_encoder = msgspec.json.Encoder()
|
_change_encoder = msgspec.json.Encoder()
|
||||||
|
|
||||||
|
|
||||||
async def load_jsonl(db_path: Path) -> dict:
|
|
||||||
"""Load data from disk by applying change log.
|
|
||||||
|
|
||||||
Replays all changes from JSONL file using plain dicts (to handle
|
|
||||||
schema evolution).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
db_path: Path to the JSONL database file
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
The final state after applying all changes
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
ValueError: If file doesn't exist or cannot be loaded
|
|
||||||
"""
|
|
||||||
if not db_path.exists():
|
|
||||||
raise ValueError(f"Database file not found: {db_path}")
|
|
||||||
data_dict: dict = {}
|
|
||||||
try:
|
|
||||||
# Read entire file at once and split into lines
|
|
||||||
async with aiofiles.open(db_path, "rb") as f:
|
|
||||||
content = await f.read()
|
|
||||||
for line_num, line in enumerate(content.split(b"\n"), 1):
|
|
||||||
line = line.strip()
|
|
||||||
if not line:
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
change = msgspec.json.decode(line)
|
|
||||||
# Apply the diff to current state (marshal=True for $-prefixed keys)
|
|
||||||
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
|
|
||||||
except Exception as e:
|
|
||||||
raise ValueError(f"Error parsing line {line_num}: {e}")
|
|
||||||
except (OSError, ValueError, msgspec.DecodeError) as e:
|
|
||||||
raise ValueError(f"Failed to load database: {e}")
|
|
||||||
return data_dict
|
|
||||||
|
|
||||||
|
|
||||||
def compute_diff(previous: dict, current: dict) -> dict | None:
|
def compute_diff(previous: dict, current: dict) -> dict | None:
|
||||||
"""Compute JSON diff between two states.
|
"""Compute JSON diff between two states.
|
||||||
|
|
||||||
@@ -93,19 +54,20 @@ def compute_diff(previous: dict, current: dict) -> dict | None:
|
|||||||
|
|
||||||
|
|
||||||
def create_change_record(
|
def create_change_record(
|
||||||
action: str, diff: dict, user: str | None = None
|
action: str, version: int, diff: dict, user: str | None = None
|
||||||
) -> _ChangeRecord:
|
) -> _ChangeRecord:
|
||||||
"""Create a change record for persistence."""
|
"""Create a change record for persistence."""
|
||||||
return _ChangeRecord(
|
return _ChangeRecord(
|
||||||
ts=datetime.now(timezone.utc),
|
ts=datetime.now(UTC),
|
||||||
a=action,
|
a=action,
|
||||||
|
v=version,
|
||||||
u=user,
|
u=user,
|
||||||
diff=diff,
|
diff=diff,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Actions that are allowed to create a new database file
|
# Actions that are allowed to create a new database file
|
||||||
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap", "migrate"})
|
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap", "migrate:sql"})
|
||||||
|
|
||||||
|
|
||||||
async def flush_changes(
|
async def flush_changes(
|
||||||
@@ -166,66 +128,89 @@ class JsonlStore:
|
|||||||
self._current_user: str | None = None
|
self._current_user: str | None = None
|
||||||
self._in_transaction: bool = False
|
self._in_transaction: bool = False
|
||||||
self._transaction_snapshot: dict[str, Any] | None = None
|
self._transaction_snapshot: dict[str, Any] | None = None
|
||||||
|
self._current_version: int = DBVER # Schema version for new databases
|
||||||
|
|
||||||
async def load(self, db_path: str | None = None) -> None:
|
async def load(self, db_path: str | None = None) -> None:
|
||||||
"""Load data from JSONL change log."""
|
"""Load data from JSONL change log."""
|
||||||
if db_path is not None:
|
if db_path is not None:
|
||||||
self.db_path = Path(db_path)
|
self.db_path = Path(db_path)
|
||||||
|
if not self.db_path.exists():
|
||||||
|
return
|
||||||
|
|
||||||
|
# Replay change log to reconstruct state
|
||||||
|
data_dict: dict = {}
|
||||||
try:
|
try:
|
||||||
data_dict = await load_jsonl(self.db_path)
|
async with aiofiles.open(self.db_path, "rb") as f:
|
||||||
if data_dict:
|
content = await f.read()
|
||||||
# Preserve original state before migrations (deep copy for nested dicts)
|
for line_num, line in enumerate(content.split(b"\n"), 1):
|
||||||
original_dict = copy.deepcopy(data_dict)
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
# Apply schema migrations (modifies data_dict in place)
|
continue
|
||||||
migrated = apply_migrations(data_dict)
|
|
||||||
|
|
||||||
decoder = msgspec.json.Decoder(DB)
|
|
||||||
self.db = decoder.decode(msgspec.json.encode(data_dict))
|
|
||||||
self.db._store = self
|
|
||||||
|
|
||||||
# Update previous state to migrated data FIRST (to avoid transaction hardening reset)
|
|
||||||
self._previous_builtins = data_dict
|
|
||||||
|
|
||||||
# Persist migration by manually computing and queueing the diff
|
|
||||||
if migrated:
|
|
||||||
diff = compute_diff(original_dict, data_dict)
|
|
||||||
if diff:
|
|
||||||
self._pending_changes.append(
|
|
||||||
create_change_record("migrate", diff, user=None)
|
|
||||||
)
|
|
||||||
_logger.info("Queued migration changes for persistence")
|
|
||||||
await self.flush()
|
|
||||||
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:
|
try:
|
||||||
user_uuid = UUID(self._current_user)
|
change = msgspec.json.decode(line)
|
||||||
if user_uuid in self.db.users:
|
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
|
||||||
user_display = self.db.users[user_uuid].display_name
|
self._current_version = change.get("v", 0)
|
||||||
except (ValueError, KeyError):
|
except Exception as e:
|
||||||
user_display = self._current_user
|
raise ValueError(f"Error parsing line {line_num}: {e}")
|
||||||
|
except (OSError, ValueError, msgspec.DecodeError) as e:
|
||||||
|
raise ValueError(f"Failed to load database: {e}")
|
||||||
|
|
||||||
diff_json = json.dumps(diff, default=str)
|
if not data_dict:
|
||||||
if user_display:
|
return
|
||||||
print(
|
|
||||||
f"{self._current_action} by {user_display}: {diff_json}",
|
# Set previous state for diffing (will be updated by _queue_change)
|
||||||
file=sys.stderr,
|
self._previous_builtins = copy.deepcopy(data_dict)
|
||||||
)
|
|
||||||
else:
|
# Callback to persist each migration
|
||||||
print(f"{self._current_action}: {diff_json}", file=sys.stderr)
|
async def persist_migration(
|
||||||
|
action: str, new_version: int, current: dict
|
||||||
|
) -> None:
|
||||||
|
self._current_version = new_version
|
||||||
|
self._queue_change(action, new_version, current)
|
||||||
|
|
||||||
|
# Apply schema migrations one at a time
|
||||||
|
await apply_all_migrations(data_dict, self._current_version, persist_migration)
|
||||||
|
|
||||||
|
# Decode to msgspec struct
|
||||||
|
decoder = msgspec.json.Decoder(DB)
|
||||||
|
self.db = decoder.decode(msgspec.json.encode(data_dict))
|
||||||
|
self.db._store = self
|
||||||
|
|
||||||
|
# Normalize via msgspec round-trip (handles omit_defaults etc.)
|
||||||
|
# This ensures _previous_builtins matches what msgspec would produce
|
||||||
|
normalized_dict = msgspec.to_builtins(self.db)
|
||||||
|
await persist_migration(
|
||||||
|
"migrate:msgspec", self._current_version, normalized_dict
|
||||||
|
)
|
||||||
|
|
||||||
|
def _queue_change(
|
||||||
|
self, action: str, version: int, current: dict, user: str | None = None
|
||||||
|
) -> None:
|
||||||
|
"""Queue a change record and log it.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
action: The action name for the change record
|
||||||
|
version: The schema version for the change record
|
||||||
|
current: The current state as a plain dict
|
||||||
|
user: Optional user UUID who performed the action
|
||||||
|
"""
|
||||||
|
diff = compute_diff(self._previous_builtins, current)
|
||||||
|
if not diff:
|
||||||
|
return
|
||||||
|
self._pending_changes.append(create_change_record(action, version, diff, user))
|
||||||
|
self._previous_builtins = copy.deepcopy(current)
|
||||||
|
|
||||||
|
# Log the change with user display name if available
|
||||||
|
user_display = None
|
||||||
|
if user:
|
||||||
|
try:
|
||||||
|
user_uuid = UUID(user)
|
||||||
|
if user_uuid in self.db.users:
|
||||||
|
user_display = self.db.users[user_uuid].display_name
|
||||||
|
except (ValueError, KeyError):
|
||||||
|
user_display = user
|
||||||
|
|
||||||
|
log_change(action, diff, user_display)
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def transaction(
|
def transaction(
|
||||||
@@ -248,19 +233,19 @@ class JsonlStore:
|
|||||||
# Check for out-of-transaction modifications
|
# Check for out-of-transaction modifications
|
||||||
current_state = msgspec.to_builtins(self.db)
|
current_state = msgspec.to_builtins(self.db)
|
||||||
if current_state != self._previous_builtins:
|
if current_state != self._previous_builtins:
|
||||||
diff = compute_diff(self._previous_builtins, current_state)
|
# Allow bootstrap/migrate to create a new database from empty state
|
||||||
diff_json = json.dumps(diff, default=str, indent=2)
|
is_bootstrap = action in _BOOTSTRAP_ACTIONS or action.startswith("migrate:")
|
||||||
_logger.error(
|
if is_bootstrap and not self._previous_builtins:
|
||||||
"Database state modified outside of transaction! "
|
pass # Expected: creating database from scratch
|
||||||
"This indicates a bug where DB changes occurred without a transaction wrapper. "
|
else:
|
||||||
"Resetting to last known state from JSONL file.\n"
|
diff = compute_diff(self._previous_builtins, current_state)
|
||||||
f"Changes detected:\n{diff_json}"
|
diff_json = msgspec.json.encode(diff).decode()
|
||||||
)
|
_logger.critical(
|
||||||
# Hard reset to last known good state
|
"Database state modified outside of transaction! "
|
||||||
decoder = msgspec.json.Decoder(DB)
|
"This indicates a bug where DB changes occurred without a transaction wrapper.\n"
|
||||||
self.db = decoder.decode(msgspec.json.encode(self._previous_builtins))
|
f"Changes detected:\n{diff_json}"
|
||||||
self.db._store = self
|
)
|
||||||
current_state = self._previous_builtins.copy()
|
raise SystemExit(1)
|
||||||
|
|
||||||
old_action = self._current_action
|
old_action = self._current_action
|
||||||
old_user = self._current_user
|
old_user = self._current_user
|
||||||
@@ -272,7 +257,10 @@ class JsonlStore:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
self._queue_change()
|
current = msgspec.to_builtins(self.db)
|
||||||
|
self._queue_change(
|
||||||
|
self._current_action, self._current_version, current, self._current_user
|
||||||
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
# Rollback on error: restore from snapshot
|
# Rollback on error: restore from snapshot
|
||||||
_logger.warning("Transaction '%s' failed, rolling back changes", action)
|
_logger.warning("Transaction '%s' failed, rolling back changes", action)
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
"""
|
||||||
|
Database change logging with pretty-printed diffs.
|
||||||
|
|
||||||
|
Provides a logger for JSONL database changes that formats diffs
|
||||||
|
in a human-readable path.notation style with color coding.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
logger = logging.getLogger("paskia.db")
|
||||||
|
|
||||||
|
# Pattern to match control characters and bidirectional overrides
|
||||||
|
_UNSAFE_CHARS = re.compile(
|
||||||
|
r"[\x00-\x1f\x7f-\x9f" # C0 and C1 control characters
|
||||||
|
r"\u200e\u200f" # LRM, RLM
|
||||||
|
r"\u202a-\u202e" # LRE, RLE, PDF, LRO, RLO
|
||||||
|
r"\u2066-\u2069" # LRI, RLI, FSI, PDI
|
||||||
|
r"]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# ANSI color codes (matching FastAPI logging style)
|
||||||
|
_RESET = "\033[0m"
|
||||||
|
_DIM = "\033[2m"
|
||||||
|
_PATH_PREFIX = "\033[1;30m" # Dark grey for path prefix (like host in access log)
|
||||||
|
_PATH_FINAL = "\033[0m" # Default for final element (like path in access log)
|
||||||
|
_REPLACE = "\033[0;33m" # Yellow for replacements
|
||||||
|
_DELETE = "\033[0;31m" # Red for deletions
|
||||||
|
_ADD = "\033[0;32m" # Green for additions
|
||||||
|
_ACTION = "\033[1;34m" # Bold blue for action name
|
||||||
|
_USER = "\033[0;34m" # Blue for user display
|
||||||
|
|
||||||
|
|
||||||
|
def _use_color() -> bool:
|
||||||
|
"""Check if we should use color output."""
|
||||||
|
return sys.stderr.isatty()
|
||||||
|
|
||||||
|
|
||||||
|
def _format_value(value: Any, use_color: bool, max_len: int = 60) -> str:
|
||||||
|
"""Format a value for display, truncating if needed."""
|
||||||
|
if value is None:
|
||||||
|
return "null"
|
||||||
|
|
||||||
|
if isinstance(value, bool):
|
||||||
|
return "true" if value else "false"
|
||||||
|
|
||||||
|
if isinstance(value, (int, float)):
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
if isinstance(value, str):
|
||||||
|
# Filter out control characters and bidirectional overrides
|
||||||
|
value = _UNSAFE_CHARS.sub("", value)
|
||||||
|
# Truncate long strings
|
||||||
|
if len(value) > max_len:
|
||||||
|
return value[: max_len - 3] + "..."
|
||||||
|
return value
|
||||||
|
|
||||||
|
if isinstance(value, dict):
|
||||||
|
if not value:
|
||||||
|
return "{}"
|
||||||
|
# For small dicts, show inline
|
||||||
|
if len(value) == 1:
|
||||||
|
k, v = next(iter(value.items()))
|
||||||
|
return "{" + f"{k}: {_format_value(v, use_color, max_len=30)}" + "}"
|
||||||
|
return f"{{...{len(value)} keys}}"
|
||||||
|
|
||||||
|
if isinstance(value, list):
|
||||||
|
if not value:
|
||||||
|
return "[]"
|
||||||
|
if len(value) == 1:
|
||||||
|
return "[" + _format_value(value[0], use_color, max_len=30) + "]"
|
||||||
|
return f"[...{len(value)} items]"
|
||||||
|
|
||||||
|
# Fallback for other types
|
||||||
|
text = str(value)
|
||||||
|
if len(text) > max_len:
|
||||||
|
text = text[: max_len - 3] + "..."
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def _format_path(path: list[str], use_color: bool) -> str:
|
||||||
|
"""Format a path as dot notation with prefix in dark grey, final in default."""
|
||||||
|
if not path:
|
||||||
|
return ""
|
||||||
|
if not use_color:
|
||||||
|
return ".".join(path)
|
||||||
|
if len(path) == 1:
|
||||||
|
return f"{_PATH_FINAL}{path[0]}{_RESET}"
|
||||||
|
prefix = ".".join(path[:-1])
|
||||||
|
final = path[-1]
|
||||||
|
return f"{_PATH_PREFIX}{prefix}.{_RESET}{_PATH_FINAL}{final}{_RESET}"
|
||||||
|
|
||||||
|
|
||||||
|
def _collect_changes(
|
||||||
|
diff: dict, path: list[str], changes: list[tuple[str, list[str], Any, Any | None]]
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Recursively collect changes from a diff into a flat list.
|
||||||
|
|
||||||
|
Each change is a tuple of (change_type, path, new_value, old_value).
|
||||||
|
change_type is one of: 'set', 'replace', 'delete'
|
||||||
|
"""
|
||||||
|
if not isinstance(diff, dict):
|
||||||
|
# Leaf value - this is a set operation
|
||||||
|
changes.append(("set", path, diff, None))
|
||||||
|
return
|
||||||
|
|
||||||
|
for key, value in diff.items():
|
||||||
|
if key == "$delete":
|
||||||
|
# $delete contains a list of keys to delete
|
||||||
|
if isinstance(value, list):
|
||||||
|
for deleted_key in value:
|
||||||
|
changes.append(("delete", path + [str(deleted_key)], None, None))
|
||||||
|
else:
|
||||||
|
changes.append(("delete", path + [str(value)], None, None))
|
||||||
|
|
||||||
|
elif key == "$replace":
|
||||||
|
# $replace contains the new value for this path
|
||||||
|
if isinstance(value, dict):
|
||||||
|
# Replacing with a dict - show each key as a replacement
|
||||||
|
for rkey, rval in value.items():
|
||||||
|
changes.append(("replace", path + [str(rkey)], rval, None))
|
||||||
|
if not value:
|
||||||
|
# Empty replacement - clearing the collection
|
||||||
|
changes.append(("replace", path, {}, None))
|
||||||
|
else:
|
||||||
|
changes.append(("replace", path, value, None))
|
||||||
|
|
||||||
|
elif key.startswith("$"):
|
||||||
|
# Other special operations (future-proofing)
|
||||||
|
changes.append(("set", path, {key: value}, None))
|
||||||
|
|
||||||
|
else:
|
||||||
|
# Regular nested key
|
||||||
|
_collect_changes(value, path + [str(key)], changes)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_change_line(
|
||||||
|
change_type: str, path: list[str], value: Any, use_color: bool
|
||||||
|
) -> str:
|
||||||
|
"""Format a single change as a one-line string."""
|
||||||
|
path_str = _format_path(path, use_color)
|
||||||
|
value_str = _format_value(value, use_color)
|
||||||
|
|
||||||
|
if change_type == "delete":
|
||||||
|
if use_color:
|
||||||
|
return f" ❌ {path_str}"
|
||||||
|
return f" - {path_str}"
|
||||||
|
|
||||||
|
if change_type == "replace":
|
||||||
|
if use_color:
|
||||||
|
return f" {_REPLACE}⟳{_RESET} {path_str} {_DIM}={_RESET} {value_str}"
|
||||||
|
return f" ~ {path_str} = {value_str}"
|
||||||
|
|
||||||
|
# Default: set/add
|
||||||
|
if use_color:
|
||||||
|
return f" {_ADD}+{_RESET} {path_str} {_DIM}={_RESET} {value_str}"
|
||||||
|
return f" + {path_str} = {value_str}"
|
||||||
|
|
||||||
|
|
||||||
|
def format_diff(diff: dict) -> list[str]:
|
||||||
|
"""
|
||||||
|
Format a JSON diff as human-readable lines.
|
||||||
|
|
||||||
|
Returns a list of formatted lines (without newlines).
|
||||||
|
Single changes return one line, multiple changes return multiple lines.
|
||||||
|
"""
|
||||||
|
use_color = _use_color()
|
||||||
|
changes: list[tuple[str, list[str], Any, Any | None]] = []
|
||||||
|
_collect_changes(diff, [], changes)
|
||||||
|
|
||||||
|
if not changes:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Format each change
|
||||||
|
lines = []
|
||||||
|
for change_type, path, value, _ in changes:
|
||||||
|
lines.append(_format_change_line(change_type, path, value, use_color))
|
||||||
|
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def format_action_header(action: str, user_display: str | None = None) -> str:
|
||||||
|
"""Format the action header line."""
|
||||||
|
use_color = _use_color()
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
action_str = f"{_ACTION}{action}{_RESET}"
|
||||||
|
if user_display:
|
||||||
|
user_str = f"{_USER}{user_display}{_RESET}"
|
||||||
|
return f"{action_str} by {user_str}"
|
||||||
|
return action_str
|
||||||
|
else:
|
||||||
|
if user_display:
|
||||||
|
return f"{action} by {user_display}"
|
||||||
|
return action
|
||||||
|
|
||||||
|
|
||||||
|
def log_change(action: str, diff: dict, user_display: str | None = None) -> None:
|
||||||
|
"""
|
||||||
|
Log a database change with pretty-printed diff.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
action: The action name (e.g., "login", "admin:delete_user")
|
||||||
|
diff: The JSON diff dict
|
||||||
|
user_display: Optional display name of the user who performed the action
|
||||||
|
"""
|
||||||
|
header = format_action_header(action, user_display)
|
||||||
|
diff_lines = format_diff(diff)
|
||||||
|
|
||||||
|
if not diff_lines:
|
||||||
|
logger.info(header)
|
||||||
|
return
|
||||||
|
|
||||||
|
if len(diff_lines) == 1:
|
||||||
|
# Single change - combine on one line
|
||||||
|
logger.info(f"{header}{diff_lines[0]}")
|
||||||
|
else:
|
||||||
|
# Multiple changes - header on its own line, then changes
|
||||||
|
logger.info(header)
|
||||||
|
for line in diff_lines:
|
||||||
|
logger.info(line)
|
||||||
|
|
||||||
|
|
||||||
|
def configure_db_logging() -> None:
|
||||||
|
"""Configure the database logger to output to stderr without prefix."""
|
||||||
|
handler = logging.StreamHandler(sys.stderr)
|
||||||
|
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||||
|
logger.addHandler(handler)
|
||||||
|
logger.setLevel(logging.INFO)
|
||||||
|
logger.propagate = False
|
||||||
+20
-21
@@ -5,30 +5,29 @@ Migrations are applied during database load based on the version field.
|
|||||||
Each migration should be idempotent and only run when needed.
|
Each migration should be idempotent and only run when needed.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
from collections.abc import Awaitable, Callable
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
def apply_migrations(data_dict: dict) -> bool:
|
def migrate_v1(d: dict) -> None:
|
||||||
"""Apply any pending schema migrations to the database dictionary.
|
"""Remove Org.created_at fields."""
|
||||||
|
for org_data in d["orgs"].values():
|
||||||
|
org_data.pop("created_at", None)
|
||||||
|
|
||||||
Args:
|
|
||||||
data_dict: The raw database dictionary loaded from JSONL
|
|
||||||
|
|
||||||
Returns:
|
migrations = sorted(
|
||||||
True if any migrations were applied, False otherwise
|
[f for n, f in globals().items() if n.startswith("migrate_v")],
|
||||||
"""
|
key=lambda f: int(f.__name__.removeprefix("migrate_v")),
|
||||||
db_version = data_dict.get("v", 0)
|
)
|
||||||
migrated = False
|
|
||||||
|
|
||||||
if db_version == 0:
|
DBVER = len(migrations) # Used by bootstrap and migrate:sql to set initial version
|
||||||
# Migration v0 -> v1: Remove created_at from orgs (field removed from schema)
|
|
||||||
if "orgs" in data_dict:
|
|
||||||
for org_data in data_dict["orgs"].values():
|
|
||||||
org_data.pop("created_at", None)
|
|
||||||
data_dict["v"] = 1
|
|
||||||
migrated = True
|
|
||||||
_logger.info("Applied schema migration: v0 -> v1 (removed org.created_at)")
|
|
||||||
|
|
||||||
return migrated
|
|
||||||
|
async def apply_all_migrations(
|
||||||
|
data_dict: dict,
|
||||||
|
current_version: int,
|
||||||
|
persist: Callable[[str, int, dict], Awaitable[None]],
|
||||||
|
) -> None:
|
||||||
|
while current_version < DBVER:
|
||||||
|
migrations[current_version](data_dict)
|
||||||
|
current_version += 1
|
||||||
|
await persist(f"migrate:v{current_version}", current_version, data_dict)
|
||||||
|
|||||||
+99
-192
@@ -2,7 +2,7 @@
|
|||||||
Database for WebAuthn passkey authentication.
|
Database for WebAuthn passkey authentication.
|
||||||
|
|
||||||
Read operations: Access _db 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.
|
Context lookup: _db.session_ctx() returns full SessionContext with effective permissions.
|
||||||
Write operations: Functions that validate and commit, or raise ValueError.
|
Write operations: Functions that validate and commit, or raise ValueError.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -10,7 +10,7 @@ import hashlib
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import secrets
|
import secrets
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import uuid7
|
import uuid7
|
||||||
@@ -31,7 +31,6 @@ from paskia.db.structs import (
|
|||||||
SessionContext,
|
SessionContext,
|
||||||
User,
|
User,
|
||||||
)
|
)
|
||||||
from paskia.util.hostutil import normalize_host
|
|
||||||
from paskia.util.passphrase import generate as generate_passphrase
|
from paskia.util.passphrase import generate as generate_passphrase
|
||||||
from paskia.util.passphrase import is_well_formed as _is_passphrase
|
from paskia.util.passphrase import is_well_formed as _is_passphrase
|
||||||
|
|
||||||
@@ -69,21 +68,18 @@ def get_user_organization(user_uuid: UUID) -> tuple[Org, str]:
|
|||||||
Raises ValueError if user not found.
|
Raises ValueError if user not found.
|
||||||
|
|
||||||
Call sites:
|
Call sites:
|
||||||
- Get user's organization when updating user role (admin.py:493)
|
- update_user_role_in_organization: org only
|
||||||
- Get user's organization for user credential listing (admin.py:530)
|
- admin_create_user_registration_link: org only
|
||||||
- Get user's organization for user details API (admin.py:579)
|
- admin_get_user_detail: org and role
|
||||||
- Get user's organization for updating user display name (admin.py:721)
|
- admin_update_user_display_name: org only
|
||||||
- Get user's organization for deleting user credential (admin.py:754)
|
- admin_delete_user_credential: org only
|
||||||
- Get user's organization for deleting user session (admin.py:783)
|
- admin_delete_user_session: org only
|
||||||
"""
|
"""
|
||||||
if user_uuid not in _db.users:
|
if user_uuid not in _db.users:
|
||||||
raise ValueError(f"User {user_uuid} not found")
|
raise ValueError(f"User {user_uuid} not found")
|
||||||
role_uuid = _db.users[user_uuid].role
|
user = _db.users[user_uuid]
|
||||||
if role_uuid not in _db.roles:
|
role = user.role
|
||||||
raise ValueError(f"Role {role_uuid} not found")
|
return role.org, role.display_name
|
||||||
role_data = _db.roles[role_uuid]
|
|
||||||
org_uuid = role_data.org
|
|
||||||
return _db.orgs[org_uuid], role_data.display_name
|
|
||||||
|
|
||||||
|
|
||||||
def get_organization_users(org_uuid: UUID) -> list[tuple[User, str]]:
|
def get_organization_users(org_uuid: UUID) -> list[tuple[User, str]]:
|
||||||
@@ -91,10 +87,8 @@ def get_organization_users(org_uuid: UUID) -> list[tuple[User, str]]:
|
|||||||
|
|
||||||
Returns list of (User, role_display_name) tuples.
|
Returns list of (User, role_display_name) tuples.
|
||||||
"""
|
"""
|
||||||
role_map = {
|
org = _db.orgs[org_uuid]
|
||||||
rid: r.display_name for rid, r in _db.roles.items() if r.org == org_uuid
|
return [(u, u.role.display_name) for role in org.roles for u in role.users]
|
||||||
}
|
|
||||||
return [(u, role_map[u.role]) for u in _db.users.values() if u.role in role_map]
|
|
||||||
|
|
||||||
|
|
||||||
def get_user_credential_ids(user_uuid: UUID) -> list[bytes]:
|
def get_user_credential_ids(user_uuid: UUID) -> list[bytes]:
|
||||||
@@ -102,7 +96,8 @@ def get_user_credential_ids(user_uuid: UUID) -> list[bytes]:
|
|||||||
|
|
||||||
Returns empty list if user has no credentials.
|
Returns empty list if user has no credentials.
|
||||||
"""
|
"""
|
||||||
return [c.credential_id for c in _db.credentials.values() if c.user == user_uuid]
|
assert user_uuid
|
||||||
|
return [c.credential_id for c in _db.users[user_uuid].credentials]
|
||||||
|
|
||||||
|
|
||||||
def _reset_key(passphrase: str) -> bytes:
|
def _reset_key(passphrase: str) -> bytes:
|
||||||
@@ -126,92 +121,6 @@ def get_reset_token(passphrase: str) -> ResetToken | None:
|
|||||||
return _db.reset_tokens.get(key)
|
return _db.reset_tokens.get(key)
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------------
|
|
||||||
# Context lookup
|
|
||||||
# -------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
def get_session_context(
|
|
||||||
session_key: str, host: str | None = None
|
|
||||||
) -> SessionContext | None:
|
|
||||||
"""Get full session context with effective permissions.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
session_key: The session key string
|
|
||||||
host: Optional host for binding/validation and domain-scoped permissions
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
SessionContext if valid, None if session not found, expired, or host mismatch
|
|
||||||
|
|
||||||
Call sites:
|
|
||||||
- Example usage in docstring (db/__init__.py:16)
|
|
||||||
- Get session context from auth token (util/permutil.py:43)
|
|
||||||
"""
|
|
||||||
|
|
||||||
if session_key not in _db.sessions:
|
|
||||||
return None
|
|
||||||
|
|
||||||
s = _db.sessions[session_key]
|
|
||||||
if s.expiry < datetime.now(timezone.utc):
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Validate host matches (sessions are always created with a host)
|
|
||||||
if host is not None and s.host != host:
|
|
||||||
# Session bound to different host
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Validate user exists
|
|
||||||
if s.user not in _db.users:
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Validate role exists
|
|
||||||
role_uuid = _db.users[s.user].role
|
|
||||||
if role_uuid not in _db.roles:
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Validate org exists
|
|
||||||
org_uuid = _db.roles[role_uuid].org
|
|
||||||
if org_uuid not in _db.orgs:
|
|
||||||
return None
|
|
||||||
|
|
||||||
session = _db.sessions[session_key]
|
|
||||||
user = _db.users[s.user]
|
|
||||||
role = _db.roles[role_uuid]
|
|
||||||
org = _db.orgs[org_uuid]
|
|
||||||
|
|
||||||
# Credential must exist (sessions are cascade-deleted when credential is deleted)
|
|
||||||
if s.credential not in _db.credentials:
|
|
||||||
return None
|
|
||||||
credential = _db.credentials[s.credential]
|
|
||||||
|
|
||||||
# Effective permissions: role's permissions that the org can grant
|
|
||||||
# Also filter by domain if host is provided
|
|
||||||
org_perm_uuids = {pid for pid, p in _db.permissions.items() if org_uuid in p.orgs}
|
|
||||||
normalized_host = normalize_host(host)
|
|
||||||
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
|
||||||
|
|
||||||
effective_perms = []
|
|
||||||
for perm_uuid in role.permission_set:
|
|
||||||
if perm_uuid not in org_perm_uuids:
|
|
||||||
continue
|
|
||||||
if perm_uuid not in _db.permissions:
|
|
||||||
continue
|
|
||||||
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.permissions[perm_uuid])
|
|
||||||
|
|
||||||
return SessionContext(
|
|
||||||
session=session,
|
|
||||||
user=user,
|
|
||||||
org=org,
|
|
||||||
role=role,
|
|
||||||
credential=credential,
|
|
||||||
permissions=effective_perms,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
# Write operations (validate, modify, commit or raise ValueError)
|
# Write operations (validate, modify, commit or raise ValueError)
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
@@ -264,9 +173,9 @@ def create_org(org: Org, *, ctx: SessionContext | None = None) -> None:
|
|||||||
if org.uuid in _db.orgs:
|
if org.uuid in _db.orgs:
|
||||||
raise ValueError(f"Organization {org.uuid} already exists")
|
raise ValueError(f"Organization {org.uuid} already exists")
|
||||||
with _db.transaction("admin:create_org", ctx):
|
with _db.transaction("admin:create_org", ctx):
|
||||||
new_org = Org(display_name=org.display_name)
|
new_org = Org.create(display_name=org.display_name)
|
||||||
_db.orgs[org.uuid] = new_org
|
|
||||||
new_org.uuid = org.uuid
|
new_org.uuid = org.uuid
|
||||||
|
_db.orgs[org.uuid] = new_org
|
||||||
# Create Administration role with org admin permission
|
# Create Administration role with org admin permission
|
||||||
|
|
||||||
admin_role_uuid = uuid7.create()
|
admin_role_uuid = uuid7.create()
|
||||||
@@ -278,7 +187,7 @@ def create_org(org: Org, *, ctx: SessionContext | None = None) -> None:
|
|||||||
break
|
break
|
||||||
role_permissions = {org_admin_perm_uuid: True} if org_admin_perm_uuid else {}
|
role_permissions = {org_admin_perm_uuid: True} if org_admin_perm_uuid else {}
|
||||||
admin_role = Role(
|
admin_role = Role(
|
||||||
org=org.uuid,
|
org_uuid=org.uuid,
|
||||||
display_name="Administration",
|
display_name="Administration",
|
||||||
permissions=role_permissions,
|
permissions=role_permissions,
|
||||||
)
|
)
|
||||||
@@ -304,17 +213,15 @@ def delete_org(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
|
|||||||
if uuid not in _db.orgs:
|
if uuid not in _db.orgs:
|
||||||
raise ValueError(f"Organization {uuid} not found")
|
raise ValueError(f"Organization {uuid} not found")
|
||||||
with _db.transaction("admin:delete_org", ctx):
|
with _db.transaction("admin:delete_org", ctx):
|
||||||
|
org = _db.orgs[uuid]
|
||||||
# Remove org from all permissions
|
# Remove org from all permissions
|
||||||
for p in _db.permissions.values():
|
for p in _db.permissions.values():
|
||||||
p.orgs.pop(uuid, None)
|
p.orgs.pop(uuid, None)
|
||||||
# Delete roles in this org
|
# Delete roles in this org and their users
|
||||||
role_uuids = [rid for rid, r in _db.roles.items() if r.org == uuid]
|
for role in org.roles:
|
||||||
for rid in role_uuids:
|
for user in role.users:
|
||||||
del _db.roles[rid]
|
del _db.users[user.uuid]
|
||||||
# Delete users with those roles
|
del _db.roles[role.uuid]
|
||||||
user_uuids = [uid for uid, u in _db.users.items() if u.role in role_uuids]
|
|
||||||
for uid in user_uuids:
|
|
||||||
del _db.users[uid]
|
|
||||||
del _db.orgs[uuid]
|
del _db.orgs[uuid]
|
||||||
|
|
||||||
|
|
||||||
@@ -356,8 +263,8 @@ def create_role(role: Role, *, ctx: SessionContext | None = None) -> None:
|
|||||||
"""Create a new role."""
|
"""Create a new role."""
|
||||||
if role.uuid in _db.roles:
|
if role.uuid in _db.roles:
|
||||||
raise ValueError(f"Role {role.uuid} already exists")
|
raise ValueError(f"Role {role.uuid} already exists")
|
||||||
if role.org not in _db.orgs:
|
if role.org_uuid not in _db.orgs:
|
||||||
raise ValueError(f"Organization {role.org} not found")
|
raise ValueError(f"Organization {role.org_uuid} not found")
|
||||||
with _db.transaction("admin:create_role", ctx):
|
with _db.transaction("admin:create_role", ctx):
|
||||||
_db.roles[role.uuid] = role
|
_db.roles[role.uuid] = role
|
||||||
|
|
||||||
@@ -408,7 +315,8 @@ def delete_role(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
|
|||||||
if uuid not in _db.roles:
|
if uuid not in _db.roles:
|
||||||
raise ValueError(f"Role {uuid} not found")
|
raise ValueError(f"Role {uuid} not found")
|
||||||
# Check no users have this role
|
# Check no users have this role
|
||||||
if any(u.role == uuid for u in _db.users.values()):
|
role = _db.roles[uuid]
|
||||||
|
if role.users:
|
||||||
raise ValueError(f"Cannot delete role {uuid}: users still assigned")
|
raise ValueError(f"Cannot delete role {uuid}: users still assigned")
|
||||||
with _db.transaction("admin:delete_role", ctx):
|
with _db.transaction("admin:delete_role", ctx):
|
||||||
del _db.roles[uuid]
|
del _db.roles[uuid]
|
||||||
@@ -418,8 +326,8 @@ def create_user(new_user: User, *, ctx: SessionContext | None = None) -> None:
|
|||||||
"""Create a new user."""
|
"""Create a new user."""
|
||||||
if new_user.uuid in _db.users:
|
if new_user.uuid in _db.users:
|
||||||
raise ValueError(f"User {new_user.uuid} already exists")
|
raise ValueError(f"User {new_user.uuid} already exists")
|
||||||
if new_user.role not in _db.roles:
|
if new_user.role_uuid not in _db.roles:
|
||||||
raise ValueError(f"Role {new_user.role} not found")
|
raise ValueError(f"Role {new_user.role_uuid} not found")
|
||||||
with _db.transaction("admin:create_user", ctx):
|
with _db.transaction("admin:create_user", ctx):
|
||||||
_db.users[new_user.uuid] = new_user
|
_db.users[new_user.uuid] = new_user
|
||||||
|
|
||||||
@@ -456,7 +364,7 @@ def update_user_role(
|
|||||||
if role_uuid not in _db.roles:
|
if role_uuid not in _db.roles:
|
||||||
raise ValueError(f"Role {role_uuid} not found")
|
raise ValueError(f"Role {role_uuid} not found")
|
||||||
with _db.transaction("admin:update_user_role", ctx):
|
with _db.transaction("admin:update_user_role", ctx):
|
||||||
_db.users[uuid].role = role_uuid
|
_db.users[uuid].role_uuid = role_uuid
|
||||||
|
|
||||||
|
|
||||||
def update_user_role_in_organization(
|
def update_user_role_in_organization(
|
||||||
@@ -468,39 +376,35 @@ def update_user_role_in_organization(
|
|||||||
"""Update user's role by role name within their current organization."""
|
"""Update user's role by role name within their current organization."""
|
||||||
if user_uuid not in _db.users:
|
if user_uuid not in _db.users:
|
||||||
raise ValueError(f"User {user_uuid} not found")
|
raise ValueError(f"User {user_uuid} not found")
|
||||||
current_role_uuid = _db.users[user_uuid].role
|
user = _db.users[user_uuid]
|
||||||
if current_role_uuid not in _db.roles:
|
org = user.org
|
||||||
raise ValueError("Current role not found")
|
|
||||||
org_uuid = _db.roles[current_role_uuid].org
|
|
||||||
# Find role by name in the same org
|
# Find role by name in the same org
|
||||||
new_role_uuid = None
|
new_role_uuid = None
|
||||||
for rid, r in _db.roles.items():
|
for r in org.roles:
|
||||||
if r.org == org_uuid and r.display_name == role_name:
|
if r.display_name == role_name:
|
||||||
new_role_uuid = rid
|
new_role_uuid = r.uuid
|
||||||
break
|
break
|
||||||
if new_role_uuid is None:
|
if new_role_uuid is None:
|
||||||
raise ValueError(f"Role '{role_name}' not found in organization")
|
raise ValueError(f"Role '{role_name}' not found in organization")
|
||||||
with _db.transaction("admin:update_user_role", ctx):
|
with _db.transaction("admin:update_user_role", ctx):
|
||||||
_db.users[user_uuid].role = new_role_uuid
|
_db.users[user_uuid].role_uuid = new_role_uuid
|
||||||
|
|
||||||
|
|
||||||
def delete_user(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
|
def delete_user(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
|
||||||
"""Delete user and their credentials/sessions."""
|
"""Delete user and their credentials/sessions."""
|
||||||
if uuid not in _db.users:
|
if uuid not in _db.users:
|
||||||
raise ValueError(f"User {uuid} not found")
|
raise ValueError(f"User {uuid} not found")
|
||||||
|
user = _db.users[uuid]
|
||||||
with _db.transaction("admin:delete_user", ctx):
|
with _db.transaction("admin:delete_user", ctx):
|
||||||
# Delete credentials
|
# Delete credentials
|
||||||
cred_uuids = [cid for cid, c in _db.credentials.items() if c.user == uuid]
|
for cred in user.credentials:
|
||||||
for cid in cred_uuids:
|
del _db.credentials[cred.uuid]
|
||||||
del _db.credentials[cid]
|
|
||||||
# Delete sessions
|
# Delete sessions
|
||||||
sess_keys = [k for k, s in _db.sessions.items() if s.user == uuid]
|
for sess in user.sessions:
|
||||||
for k in sess_keys:
|
del _db.sessions[sess.key]
|
||||||
del _db.sessions[k]
|
|
||||||
# Delete reset tokens
|
# Delete reset tokens
|
||||||
token_keys = [k for k, t in _db.reset_tokens.items() if t.user == uuid]
|
for token in user.reset_tokens:
|
||||||
for k in token_keys:
|
del _db.reset_tokens[token.key]
|
||||||
del _db.reset_tokens[k]
|
|
||||||
del _db.users[uuid]
|
del _db.users[uuid]
|
||||||
|
|
||||||
|
|
||||||
@@ -508,8 +412,8 @@ def create_credential(cred: Credential, *, ctx: SessionContext | None = None) ->
|
|||||||
"""Create a new credential."""
|
"""Create a new credential."""
|
||||||
if cred.uuid in _db.credentials:
|
if cred.uuid in _db.credentials:
|
||||||
raise ValueError(f"Credential {cred.uuid} already exists")
|
raise ValueError(f"Credential {cred.uuid} already exists")
|
||||||
if cred.user not in _db.users:
|
if cred.user_uuid not in _db.users:
|
||||||
raise ValueError(f"User {cred.user} not found")
|
raise ValueError(f"User {cred.user_uuid} not found")
|
||||||
with _db.transaction("create_credential", ctx):
|
with _db.transaction("create_credential", ctx):
|
||||||
_db.credentials[cred.uuid] = cred
|
_db.credentials[cred.uuid] = cred
|
||||||
|
|
||||||
@@ -542,20 +446,19 @@ def delete_credential(
|
|||||||
"""
|
"""
|
||||||
if uuid not in _db.credentials:
|
if uuid not in _db.credentials:
|
||||||
raise ValueError(f"Credential {uuid} not found")
|
raise ValueError(f"Credential {uuid} not found")
|
||||||
|
cred = _db.credentials[uuid]
|
||||||
if user_uuid is not None:
|
if user_uuid is not None:
|
||||||
cred_user = _db.credentials[uuid].user
|
if cred.user_uuid != user_uuid:
|
||||||
if cred_user != user_uuid:
|
|
||||||
raise ValueError(f"Credential {uuid} does not belong to user {user_uuid}")
|
raise ValueError(f"Credential {uuid} does not belong to user {user_uuid}")
|
||||||
with _db.transaction("delete_credential", ctx):
|
with _db.transaction("delete_credential", ctx):
|
||||||
# Delete all sessions using this credential
|
# Delete all sessions using this credential
|
||||||
keys = [k for k, s in _db.sessions.items() if s.credential == uuid]
|
for sess in cred.sessions:
|
||||||
for k in keys:
|
print(sess, repr(sess.key))
|
||||||
del _db.sessions[k]
|
del _db.sessions[sess.key]
|
||||||
del _db.credentials[uuid]
|
del _db.credentials[uuid]
|
||||||
|
|
||||||
|
|
||||||
def create_session(
|
def create_session(
|
||||||
key: str,
|
|
||||||
user_uuid: UUID,
|
user_uuid: UUID,
|
||||||
credential_uuid: UUID,
|
credential_uuid: UUID,
|
||||||
host: str,
|
host: str,
|
||||||
@@ -564,23 +467,25 @@ def create_session(
|
|||||||
expiry: datetime,
|
expiry: datetime,
|
||||||
*,
|
*,
|
||||||
ctx: SessionContext | None = None,
|
ctx: SessionContext | None = None,
|
||||||
) -> None:
|
) -> str:
|
||||||
"""Create a new session."""
|
"""Create a new session. Returns the session key."""
|
||||||
if key in _db.sessions:
|
|
||||||
raise ValueError("Session already exists")
|
|
||||||
if user_uuid not in _db.users:
|
if user_uuid not in _db.users:
|
||||||
raise ValueError(f"User {user_uuid} not found")
|
raise ValueError(f"User {user_uuid} not found")
|
||||||
if credential_uuid not in _db.credentials:
|
if credential_uuid not in _db.credentials:
|
||||||
raise ValueError(f"Credential {credential_uuid} not found")
|
raise ValueError(f"Credential {credential_uuid} not found")
|
||||||
|
session = Session.create(
|
||||||
|
user=user_uuid,
|
||||||
|
credential=credential_uuid,
|
||||||
|
host=host,
|
||||||
|
ip=ip,
|
||||||
|
user_agent=user_agent,
|
||||||
|
expiry=expiry,
|
||||||
|
)
|
||||||
|
if session.key in _db.sessions:
|
||||||
|
raise ValueError("Session already exists")
|
||||||
with _db.transaction("create_session", ctx):
|
with _db.transaction("create_session", ctx):
|
||||||
_db.sessions[key] = Session(
|
_db.sessions[session.key] = session
|
||||||
user=user_uuid,
|
return session.key
|
||||||
credential=credential_uuid,
|
|
||||||
host=host,
|
|
||||||
ip=ip,
|
|
||||||
user_agent=user_agent,
|
|
||||||
expiry=expiry,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def update_session(
|
def update_session(
|
||||||
@@ -634,10 +539,12 @@ def delete_sessions_for_user(
|
|||||||
For user logout-all, pass ctx of the user's session.
|
For user logout-all, pass ctx of the user's session.
|
||||||
For admin bulk termination, pass admin's ctx.
|
For admin bulk termination, pass admin's ctx.
|
||||||
"""
|
"""
|
||||||
|
user = _db.users.get(user_uuid)
|
||||||
|
if not user:
|
||||||
|
return
|
||||||
with _db.transaction("admin:delete_sessions_for_user", ctx):
|
with _db.transaction("admin:delete_sessions_for_user", ctx):
|
||||||
keys = [k for k, s in _db.sessions.items() if s.user == user_uuid]
|
for sess in user.sessions:
|
||||||
for k in keys:
|
del _db.sessions[sess.key]
|
||||||
del _db.sessions[k]
|
|
||||||
|
|
||||||
|
|
||||||
def create_reset_token(
|
def create_reset_token(
|
||||||
@@ -662,7 +569,7 @@ def create_reset_token(
|
|||||||
raise ValueError(f"User {user_uuid} not found")
|
raise ValueError(f"User {user_uuid} not found")
|
||||||
with _db.transaction("create_reset_token", ctx):
|
with _db.transaction("create_reset_token", ctx):
|
||||||
_db.reset_tokens[key] = ResetToken(
|
_db.reset_tokens[key] = ResetToken(
|
||||||
user=user_uuid, expiry=expiry, token_type=token_type
|
user_uuid=user_uuid, expiry=expiry, token_type=token_type
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -681,7 +588,7 @@ def delete_reset_token(key: bytes, *, ctx: SessionContext | None = None) -> None
|
|||||||
|
|
||||||
def cleanup_expired() -> int:
|
def cleanup_expired() -> int:
|
||||||
"""Remove expired sessions and reset tokens. Returns count removed."""
|
"""Remove expired sessions and reset tokens. Returns count removed."""
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
count = 0
|
count = 0
|
||||||
with _db.transaction("expiry"):
|
with _db.transaction("expiry"):
|
||||||
expired_sessions = [k for k, s in _db.sessions.items() if s.expiry < now]
|
expired_sessions = [k for k, s in _db.sessions.items() if s.expiry < now]
|
||||||
@@ -726,13 +633,20 @@ def login(
|
|||||||
"""
|
"""
|
||||||
if isinstance(user_uuid, str):
|
if isinstance(user_uuid, str):
|
||||||
user_uuid = UUID(user_uuid)
|
user_uuid = UUID(user_uuid)
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if user_uuid not in _db.users:
|
if user_uuid not in _db.users:
|
||||||
raise ValueError(f"User {user_uuid} not found")
|
raise ValueError(f"User {user_uuid} not found")
|
||||||
if credential_uuid not in _db.credentials:
|
if credential_uuid not in _db.credentials:
|
||||||
raise ValueError(f"Credential {credential_uuid} not found")
|
raise ValueError(f"Credential {credential_uuid} not found")
|
||||||
|
|
||||||
session_key = _create_token()
|
session = Session.create(
|
||||||
|
user=user_uuid,
|
||||||
|
credential=credential_uuid,
|
||||||
|
host=host,
|
||||||
|
ip=ip,
|
||||||
|
user_agent=user_agent,
|
||||||
|
expiry=expiry,
|
||||||
|
)
|
||||||
user_str = str(user_uuid)
|
user_str = str(user_uuid)
|
||||||
with _db.transaction("login", user=user_str):
|
with _db.transaction("login", user=user_str):
|
||||||
# Update user
|
# Update user
|
||||||
@@ -742,15 +656,8 @@ def login(
|
|||||||
_db.credentials[credential_uuid].sign_count = sign_count
|
_db.credentials[credential_uuid].sign_count = sign_count
|
||||||
_db.credentials[credential_uuid].last_used = now
|
_db.credentials[credential_uuid].last_used = now
|
||||||
# Create session
|
# Create session
|
||||||
_db.sessions[session_key] = Session(
|
_db.sessions[session.key] = session
|
||||||
user=user_uuid,
|
return session.key
|
||||||
credential=credential_uuid,
|
|
||||||
host=host,
|
|
||||||
ip=ip,
|
|
||||||
user_agent=user_agent,
|
|
||||||
expiry=expiry,
|
|
||||||
)
|
|
||||||
return session_key
|
|
||||||
|
|
||||||
|
|
||||||
def create_credential_session(
|
def create_credential_session(
|
||||||
@@ -773,13 +680,20 @@ def create_credential_session(
|
|||||||
Returns the generated session token.
|
Returns the generated session token.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
expiry = now + SESSION_LIFETIME
|
expiry = now + SESSION_LIFETIME
|
||||||
session_key = _create_token()
|
|
||||||
|
|
||||||
if user_uuid not in _db.users:
|
if user_uuid not in _db.users:
|
||||||
raise ValueError(f"User {user_uuid} not found")
|
raise ValueError(f"User {user_uuid} not found")
|
||||||
|
|
||||||
|
session = Session.create(
|
||||||
|
user=user_uuid,
|
||||||
|
credential=credential.uuid,
|
||||||
|
host=host,
|
||||||
|
ip=ip,
|
||||||
|
user_agent=user_agent,
|
||||||
|
expiry=expiry,
|
||||||
|
)
|
||||||
user_str = str(user_uuid)
|
user_str = str(user_uuid)
|
||||||
with _db.transaction("create_credential_session", user=user_str):
|
with _db.transaction("create_credential_session", user=user_str):
|
||||||
# Update display name if provided
|
# Update display name if provided
|
||||||
@@ -790,20 +704,13 @@ def create_credential_session(
|
|||||||
_db.credentials[credential.uuid] = credential
|
_db.credentials[credential.uuid] = credential
|
||||||
|
|
||||||
# Create session
|
# Create session
|
||||||
_db.sessions[session_key] = Session(
|
_db.sessions[session.key] = session
|
||||||
user=user_uuid,
|
|
||||||
credential=credential.uuid,
|
|
||||||
host=host,
|
|
||||||
ip=ip,
|
|
||||||
user_agent=user_agent,
|
|
||||||
expiry=expiry,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Delete reset token if provided
|
# Delete reset token if provided
|
||||||
if reset_key:
|
if reset_key:
|
||||||
if reset_key in _db.reset_tokens:
|
if reset_key in _db.reset_tokens:
|
||||||
del _db.reset_tokens[reset_key]
|
del _db.reset_tokens[reset_key]
|
||||||
return session_key
|
return session.key
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
@@ -862,7 +769,7 @@ def bootstrap(
|
|||||||
reset_expiry = reset_expires()
|
reset_expiry = reset_expires()
|
||||||
reset_key = _reset_key(reset_passphrase)
|
reset_key = _reset_key(reset_passphrase)
|
||||||
|
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
|
|
||||||
with _db.transaction("bootstrap"):
|
with _db.transaction("bootstrap"):
|
||||||
# Create auth:admin permission
|
# Create auth:admin permission
|
||||||
@@ -884,13 +791,13 @@ def bootstrap(
|
|||||||
_db.permissions[perm_org_admin_uuid] = perm_org_admin
|
_db.permissions[perm_org_admin_uuid] = perm_org_admin
|
||||||
|
|
||||||
# Create organization
|
# Create organization
|
||||||
new_org = Org(display_name=org_name)
|
new_org = Org.create(display_name=org_name)
|
||||||
new_org.uuid = org_uuid
|
new_org.uuid = org_uuid
|
||||||
_db.orgs[org_uuid] = new_org
|
_db.orgs[org_uuid] = new_org
|
||||||
|
|
||||||
# Create Administration role with both permissions
|
# Create Administration role with both permissions
|
||||||
admin_role = Role(
|
admin_role = Role(
|
||||||
org=org_uuid,
|
org_uuid=org_uuid,
|
||||||
display_name="Administration",
|
display_name="Administration",
|
||||||
permissions={perm_admin_uuid: True, perm_org_admin_uuid: True},
|
permissions={perm_admin_uuid: True, perm_org_admin_uuid: True},
|
||||||
)
|
)
|
||||||
@@ -900,7 +807,7 @@ def bootstrap(
|
|||||||
# Create admin user
|
# Create admin user
|
||||||
admin_user = User(
|
admin_user = User(
|
||||||
display_name=admin_name,
|
display_name=admin_name,
|
||||||
role=role_uuid,
|
role_uuid=role_uuid,
|
||||||
created_at=now,
|
created_at=now,
|
||||||
last_seen=None,
|
last_seen=None,
|
||||||
visits=0,
|
visits=0,
|
||||||
@@ -910,7 +817,7 @@ def bootstrap(
|
|||||||
|
|
||||||
# Create reset token
|
# Create reset token
|
||||||
_db.reset_tokens[reset_key] = ResetToken(
|
_db.reset_tokens[reset_key] = ResetToken(
|
||||||
user=user_uuid,
|
user_uuid=user_uuid,
|
||||||
expiry=reset_expiry,
|
expiry=reset_expiry,
|
||||||
token_type="admin bootstrap",
|
token_type="admin bootstrap",
|
||||||
)
|
)
|
||||||
|
|||||||
+236
-46
@@ -1,9 +1,15 @@
|
|||||||
from datetime import datetime, timezone
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import secrets
|
||||||
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
import uuid7
|
import uuid7
|
||||||
|
|
||||||
|
from paskia import db
|
||||||
|
from paskia.util.hostutil import normalize_host
|
||||||
|
|
||||||
# Sentinel for uuid fields before they are set by create() or DB post init
|
# Sentinel for uuid fields before they are set by create() or DB post init
|
||||||
_UUID_UNSET = UUID(int=0)
|
_UUID_UNSET = UUID(int=0)
|
||||||
|
|
||||||
@@ -22,20 +28,30 @@ class Permission(msgspec.Struct, dict=True, omit_defaults=True):
|
|||||||
orgs: dict[UUID, bool] = {} # org_uuid -> True (which orgs can grant this)
|
orgs: dict[UUID, bool] = {} # org_uuid -> True (which orgs can grant this)
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.uuid: UUID = _UUID_UNSET # Convenience field, not serialized
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def org_set(self) -> set[UUID]:
|
def org_set(self) -> set[UUID]:
|
||||||
"""Get orgs that can grant this permission as a set."""
|
"""Get orgs that can grant this permission as a set."""
|
||||||
return set(self.orgs.keys())
|
return set(self.orgs.keys())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def orgs_list(self) -> list[Org]:
|
||||||
|
"""Get list of Org objects that can grant this permission."""
|
||||||
|
return [
|
||||||
|
db.data().orgs[org_uuid]
|
||||||
|
for org_uuid in self.orgs.keys()
|
||||||
|
if org_uuid in db.data().orgs
|
||||||
|
]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
cls,
|
cls,
|
||||||
scope: str,
|
scope: str,
|
||||||
display_name: str,
|
display_name: str,
|
||||||
domain: str | None = None,
|
domain: str | None = None,
|
||||||
) -> "Permission":
|
) -> Permission:
|
||||||
"""Create a new Permission with auto-generated uuid7."""
|
"""Create a new Permission with auto-generated uuid7."""
|
||||||
perm = cls(
|
perm = cls(
|
||||||
scope=scope,
|
scope=scope,
|
||||||
@@ -46,36 +62,84 @@ class Permission(msgspec.Struct, dict=True, omit_defaults=True):
|
|||||||
return perm
|
return perm
|
||||||
|
|
||||||
|
|
||||||
|
class Org(msgspec.Struct, dict=True):
|
||||||
|
"""Organization data structure."""
|
||||||
|
|
||||||
|
display_name: str
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def roles(self) -> list[Role]:
|
||||||
|
"""Get all roles that belong to this organization."""
|
||||||
|
return [r for r in db.data().roles.values() if r.org_uuid == self.uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def permissions(self) -> list[Permission]:
|
||||||
|
"""Get all permissions that this organization can grant."""
|
||||||
|
return [p for p in db.data().permissions.values() if self.uuid in p.orgs]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, display_name: str) -> Org:
|
||||||
|
"""Create a new Org with auto-generated uuid7."""
|
||||||
|
org = cls(display_name=display_name)
|
||||||
|
org.uuid = uuid7.create()
|
||||||
|
return org
|
||||||
|
|
||||||
|
|
||||||
class Role(msgspec.Struct, dict=True, omit_defaults=True):
|
class Role(msgspec.Struct, dict=True, omit_defaults=True):
|
||||||
"""Role data structure.
|
"""Role data structure.
|
||||||
|
|
||||||
Mutable fields: display_name, permissions
|
Mutable fields: display_name, permissions
|
||||||
Immutable fields: org (set at creation, never modified)
|
Immutable fields: org_uuid (set at creation, never modified)
|
||||||
uuid is generated at creation.
|
uuid is generated at creation.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
org: UUID
|
org_uuid: UUID = msgspec.field(name="org")
|
||||||
display_name: str
|
display_name: str
|
||||||
permissions: dict[UUID, bool] = {} # permission_uuid -> True
|
permissions: dict[UUID, bool] = {} # permission_uuid -> True
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.uuid: UUID = _UUID_UNSET # Convenience field, not serialized
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def permission_set(self) -> set[UUID]:
|
def permission_set(self) -> set[UUID]:
|
||||||
"""Get permissions as a set of UUIDs."""
|
"""Get permissions as a set of UUIDs."""
|
||||||
return set(self.permissions.keys())
|
return set(self.permissions.keys())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def permissions_list(self) -> list[Permission]:
|
||||||
|
"""Get list of Permission objects for this role."""
|
||||||
|
return [
|
||||||
|
db.data().permissions[perm_uuid]
|
||||||
|
for perm_uuid in self.permissions.keys()
|
||||||
|
if perm_uuid in db.data().permissions
|
||||||
|
]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def org(self) -> Org:
|
||||||
|
"""Get the organization object this role belongs to."""
|
||||||
|
return db.data().orgs[self.org_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def users(self) -> list[User]:
|
||||||
|
"""Get all users that have this role."""
|
||||||
|
return [u for u in db.data().users.values() if u.role_uuid == self.uuid]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
cls,
|
cls,
|
||||||
org: UUID,
|
org: UUID | Org,
|
||||||
display_name: str,
|
display_name: str,
|
||||||
permissions: set[UUID] | None = None,
|
permissions: set[UUID] | None = None,
|
||||||
) -> "Role":
|
) -> Role:
|
||||||
"""Create a new Role with auto-generated uuid7."""
|
"""Create a new Role with auto-generated uuid7."""
|
||||||
|
org_uuid = org if isinstance(org, UUID) else org.uuid
|
||||||
role = cls(
|
role = cls(
|
||||||
org=org,
|
org_uuid=org_uuid,
|
||||||
display_name=display_name,
|
display_name=display_name,
|
||||||
permissions={p: True for p in (permissions or set())},
|
permissions={p: True for p in (permissions or set())},
|
||||||
)
|
)
|
||||||
@@ -83,52 +147,62 @@ class Role(msgspec.Struct, dict=True, omit_defaults=True):
|
|||||||
return role
|
return role
|
||||||
|
|
||||||
|
|
||||||
class Org(msgspec.Struct, dict=True):
|
|
||||||
"""Organization data structure."""
|
|
||||||
|
|
||||||
display_name: str
|
|
||||||
|
|
||||||
def __post_init__(self):
|
|
||||||
self.uuid: UUID = _UUID_UNSET # Convenience field, not serialized
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, display_name: str) -> "Org":
|
|
||||||
"""Create a new Org with auto-generated uuid7."""
|
|
||||||
org = cls(display_name=display_name)
|
|
||||||
org.uuid = uuid7.create()
|
|
||||||
return org
|
|
||||||
|
|
||||||
|
|
||||||
class User(msgspec.Struct, dict=True):
|
class User(msgspec.Struct, dict=True):
|
||||||
"""User data structure.
|
"""User data structure.
|
||||||
|
|
||||||
Mutable fields: display_name, role, last_seen, visits
|
Mutable fields: display_name, role_uuid, last_seen, visits
|
||||||
Immutable fields: created_at (set at creation, never modified)
|
Immutable fields: created_at (set at creation, never modified)
|
||||||
uuid is derived from created_at using uuid7.
|
uuid is derived from created_at using uuid7.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
display_name: str
|
display_name: str
|
||||||
role: UUID
|
role_uuid: UUID = msgspec.field(name="role")
|
||||||
created_at: datetime
|
created_at: datetime
|
||||||
last_seen: datetime | None = None
|
last_seen: datetime | None = None
|
||||||
visits: int = 0
|
visits: int = 0
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.uuid: UUID = _UUID_UNSET # Convenience field, not serialized
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role(self) -> Role:
|
||||||
|
"""Get the role object this user has."""
|
||||||
|
return db.data().roles[self.role_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def org(self) -> Org:
|
||||||
|
"""Get the organization this user belongs to (via role)."""
|
||||||
|
return self.role.org
|
||||||
|
|
||||||
|
@property
|
||||||
|
def credentials(self) -> list[Credential]:
|
||||||
|
"""Get all credentials for this user."""
|
||||||
|
return [c for c in db.data().credentials.values() if c.user_uuid == self.uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sessions(self) -> list[Session]:
|
||||||
|
"""Get all sessions for this user."""
|
||||||
|
return [s for s in db.data().sessions.values() if s.user_uuid == self.uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def reset_tokens(self) -> list[ResetToken]:
|
||||||
|
"""Get all reset tokens for this user."""
|
||||||
|
return [t for t in db.data().reset_tokens.values() if t.user_uuid == self.uuid]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
cls,
|
cls,
|
||||||
display_name: str,
|
display_name: str,
|
||||||
role: UUID,
|
role: UUID | Role,
|
||||||
created_at: datetime | None = None,
|
created_at: datetime | None = None,
|
||||||
) -> "User":
|
) -> User:
|
||||||
"""Create a new User with auto-generated uuid7."""
|
"""Create a new User with auto-generated uuid7."""
|
||||||
|
role_uuid = role if isinstance(role, UUID) else role.uuid
|
||||||
user = cls(
|
user = cls(
|
||||||
display_name=display_name,
|
display_name=display_name,
|
||||||
role=role,
|
role_uuid=role_uuid,
|
||||||
created_at=created_at or datetime.now(timezone.utc),
|
created_at=created_at or datetime.now(UTC),
|
||||||
)
|
)
|
||||||
user.uuid = uuid7.create(user.created_at)
|
user.uuid = uuid7.create(user.created_at)
|
||||||
return user
|
return user
|
||||||
@@ -143,7 +217,7 @@ class Credential(msgspec.Struct, dict=True):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
credential_id: bytes # Long binary ID from the authenticator
|
credential_id: bytes # Long binary ID from the authenticator
|
||||||
user: UUID
|
user_uuid: UUID = msgspec.field(name="user")
|
||||||
aaguid: UUID
|
aaguid: UUID
|
||||||
public_key: bytes
|
public_key: bytes
|
||||||
sign_count: int
|
sign_count: int
|
||||||
@@ -152,23 +226,37 @@ class Credential(msgspec.Struct, dict=True):
|
|||||||
last_verified: datetime | None = None
|
last_verified: datetime | None = None
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.uuid: UUID = _UUID_UNSET # Convenience field, not serialized
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def user(self) -> User:
|
||||||
|
"""Get the User object for this credential."""
|
||||||
|
return db.data().users[self.user_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sessions(self) -> list[Session]:
|
||||||
|
"""Get all sessions using this credential."""
|
||||||
|
return [
|
||||||
|
s for s in db.data().sessions.values() if s.credential_uuid == self.uuid
|
||||||
|
]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
cls,
|
cls,
|
||||||
credential_id: bytes,
|
credential_id: bytes,
|
||||||
user: UUID,
|
user: UUID | User,
|
||||||
aaguid: UUID,
|
aaguid: UUID,
|
||||||
public_key: bytes,
|
public_key: bytes,
|
||||||
sign_count: int,
|
sign_count: int,
|
||||||
created_at: datetime | None = None,
|
created_at: datetime | None = None,
|
||||||
) -> "Credential":
|
) -> Credential:
|
||||||
"""Create a new Credential with auto-generated uuid7."""
|
"""Create a new Credential with auto-generated uuid7."""
|
||||||
now = created_at or datetime.now(timezone.utc)
|
user_uuid = user if isinstance(user, UUID) else user.uuid
|
||||||
|
now = created_at or datetime.now(UTC)
|
||||||
cred = cls(
|
cred = cls(
|
||||||
credential_id=credential_id,
|
credential_id=credential_id,
|
||||||
user=user,
|
user_uuid=user_uuid,
|
||||||
aaguid=aaguid,
|
aaguid=aaguid,
|
||||||
public_key=public_key,
|
public_key=public_key,
|
||||||
sign_count=sign_count,
|
sign_count=sign_count,
|
||||||
@@ -184,19 +272,30 @@ class Session(msgspec.Struct, dict=True):
|
|||||||
"""Session data structure.
|
"""Session data structure.
|
||||||
|
|
||||||
Mutable fields: expiry (updated on session refresh)
|
Mutable fields: expiry (updated on session refresh)
|
||||||
Immutable fields: user, credential, host, ip, user_agent
|
Immutable fields: user_uuid, credential_uuid, host, ip, user_agent
|
||||||
key is stored in the dict key, not in the struct.
|
key is stored in the dict key, not in the struct.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
user: UUID
|
user_uuid: UUID = msgspec.field(name="user")
|
||||||
credential: UUID
|
credential_uuid: UUID = msgspec.field(name="credential")
|
||||||
host: str
|
host: str
|
||||||
ip: str
|
ip: str
|
||||||
user_agent: str
|
user_agent: str
|
||||||
expiry: datetime
|
expiry: datetime
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.key: str = "" # Convenience field, not serialized
|
if not hasattr(self, "key"):
|
||||||
|
self.key: str = ""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def user(self) -> User:
|
||||||
|
"""Get the User object for this session."""
|
||||||
|
return db.data().users[self.user_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def credential(self) -> Credential:
|
||||||
|
"""Get the Credential object for this session."""
|
||||||
|
return db.data().credentials[self.credential_uuid]
|
||||||
|
|
||||||
def metadata(self) -> dict:
|
def metadata(self) -> dict:
|
||||||
"""Return session metadata for backwards compatibility."""
|
"""Return session metadata for backwards compatibility."""
|
||||||
@@ -206,6 +305,32 @@ class Session(msgspec.Struct, dict=True):
|
|||||||
"expiry": self.expiry.isoformat(),
|
"expiry": self.expiry.isoformat(),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
user: UUID | User,
|
||||||
|
credential: UUID | Credential,
|
||||||
|
host: str,
|
||||||
|
ip: str,
|
||||||
|
user_agent: str,
|
||||||
|
expiry: datetime,
|
||||||
|
) -> Session:
|
||||||
|
"""Create a new Session with auto-generated key."""
|
||||||
|
user_uuid = user if isinstance(user, UUID) else user.uuid
|
||||||
|
credential_uuid = (
|
||||||
|
credential if isinstance(credential, UUID) else credential.uuid
|
||||||
|
)
|
||||||
|
session = cls(
|
||||||
|
user_uuid=user_uuid,
|
||||||
|
credential_uuid=credential_uuid,
|
||||||
|
host=host,
|
||||||
|
ip=ip,
|
||||||
|
user_agent=user_agent,
|
||||||
|
expiry=expiry,
|
||||||
|
)
|
||||||
|
session.key = secrets.token_urlsafe(12)
|
||||||
|
return session
|
||||||
|
|
||||||
|
|
||||||
class ResetToken(msgspec.Struct, dict=True):
|
class ResetToken(msgspec.Struct, dict=True):
|
||||||
"""Reset/device-addition token data structure.
|
"""Reset/device-addition token data structure.
|
||||||
@@ -214,12 +339,18 @@ class ResetToken(msgspec.Struct, dict=True):
|
|||||||
key is stored in the dict key, not in the struct.
|
key is stored in the dict key, not in the struct.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
user: UUID
|
user_uuid: UUID = msgspec.field(name="user")
|
||||||
expiry: datetime
|
expiry: datetime
|
||||||
token_type: str
|
token_type: str
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self.key: bytes = b"" # Convenience field, not serialized
|
if not hasattr(self, "key"):
|
||||||
|
self.key: bytes = b""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def user(self) -> User:
|
||||||
|
"""Get the User object for this reset token."""
|
||||||
|
return db.data().users[self.user_uuid]
|
||||||
|
|
||||||
|
|
||||||
class SessionContext(msgspec.Struct):
|
class SessionContext(msgspec.Struct):
|
||||||
@@ -246,7 +377,6 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
|||||||
credentials: dict[UUID, Credential] = {}
|
credentials: dict[UUID, Credential] = {}
|
||||||
sessions: dict[str, Session] = {}
|
sessions: dict[str, Session] = {}
|
||||||
reset_tokens: dict[bytes, ResetToken] = {}
|
reset_tokens: dict[bytes, ResetToken] = {}
|
||||||
v: int = 0
|
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
# Store reference for persistence (not serialized)
|
# Store reference for persistence (not serialized)
|
||||||
@@ -270,3 +400,63 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
|||||||
def transaction(self, action, ctx=None, *, user=None):
|
def transaction(self, action, ctx=None, *, user=None):
|
||||||
"""Wrap writes in transaction. Delegates to JsonlStore."""
|
"""Wrap writes in transaction. Delegates to JsonlStore."""
|
||||||
return self._store.transaction(action, ctx, user=user)
|
return self._store.transaction(action, ctx, user=user)
|
||||||
|
|
||||||
|
def session_ctx(
|
||||||
|
self, session_key: str, host: str | None = None
|
||||||
|
) -> SessionContext | None:
|
||||||
|
"""Get full session context with effective permissions.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session_key: The session key string
|
||||||
|
host: Optional host for binding/validation and domain-scoped permissions
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
SessionContext if valid, None if session not found, expired, or host mismatch
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
s = self.sessions[session_key]
|
||||||
|
except KeyError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Validate host matches (sessions are always created with a host)
|
||||||
|
if s.host != host:
|
||||||
|
# Session bound to different host
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
user = s.user
|
||||||
|
role = user.role
|
||||||
|
org = role.org
|
||||||
|
credential = s.credential
|
||||||
|
except KeyError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Effective permissions: role's permissions that the org can grant
|
||||||
|
# Also filter by domain if host is provided
|
||||||
|
org_perm_uuids = {p.uuid for p in org.permissions}
|
||||||
|
normalized_host = normalize_host(host)
|
||||||
|
host_without_port = (
|
||||||
|
normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
||||||
|
)
|
||||||
|
|
||||||
|
effective_perms = []
|
||||||
|
for perm_uuid in role.permission_set:
|
||||||
|
if perm_uuid not in org_perm_uuids:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
p = self.permissions[perm_uuid]
|
||||||
|
except KeyError:
|
||||||
|
continue
|
||||||
|
# Check domain restriction
|
||||||
|
if p.domain is not None and p.domain != host_without_port:
|
||||||
|
continue
|
||||||
|
effective_perms.append(p)
|
||||||
|
|
||||||
|
return SessionContext(
|
||||||
|
session=s,
|
||||||
|
user=user,
|
||||||
|
org=org,
|
||||||
|
role=role,
|
||||||
|
credential=credential,
|
||||||
|
permissions=effective_perms,
|
||||||
|
)
|
||||||
|
|||||||
@@ -189,6 +189,7 @@ def main():
|
|||||||
|
|
||||||
run_kwargs: dict = {
|
run_kwargs: dict = {
|
||||||
"log_level": "info",
|
"log_level": "info",
|
||||||
|
"access_log": False, # We use custom AccessLogMiddleware instead
|
||||||
}
|
}
|
||||||
|
|
||||||
if devmode:
|
if devmode:
|
||||||
|
|||||||
+72
-195
@@ -1,5 +1,4 @@
|
|||||||
import logging
|
import logging
|
||||||
from datetime import timezone
|
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import Body, FastAPI, HTTPException, Query, Request, Response
|
from fastapi import Body, FastAPI, HTTPException, Query, Request, Response
|
||||||
@@ -13,6 +12,7 @@ from paskia.db import Permission as PermDC
|
|||||||
from paskia.db import Role as RoleDC
|
from paskia.db import Role as RoleDC
|
||||||
from paskia.db import User as UserDC
|
from paskia.db import User as UserDC
|
||||||
from paskia.fastapi import authz
|
from paskia.fastapi import authz
|
||||||
|
from paskia.fastapi.response import MsgspecResponse
|
||||||
from paskia.fastapi.session import AUTH_COOKIE
|
from paskia.fastapi.session import AUTH_COOKIE
|
||||||
from paskia.globals import passkey
|
from paskia.globals import passkey
|
||||||
from paskia.util import (
|
from paskia.util import (
|
||||||
@@ -20,46 +20,26 @@ from paskia.util import (
|
|||||||
passphrase,
|
passphrase,
|
||||||
permutil,
|
permutil,
|
||||||
querysafe,
|
querysafe,
|
||||||
useragent,
|
|
||||||
vitedev,
|
vitedev,
|
||||||
)
|
)
|
||||||
|
from paskia.util.apistructs import ApiPermission, ApiSession, format_datetime
|
||||||
from paskia.util.hostutil import normalize_host
|
from paskia.util.hostutil import normalize_host
|
||||||
|
|
||||||
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
|
||||||
|
|
||||||
def is_global_admin(ctx) -> bool:
|
def master_admin(ctx) -> bool:
|
||||||
"""Check if user has global admin permission."""
|
return any(p.scope == "auth:admin" for p in ctx.permissions)
|
||||||
effective_scopes = (
|
|
||||||
{p.scope for p in (ctx.permissions or [])}
|
|
||||||
if ctx.permissions
|
def org_admin(ctx, org_uuid: UUID) -> bool:
|
||||||
else set(ctx.role.permissions or [])
|
return ctx.org.uuid == org_uuid and any(
|
||||||
|
p.scope == "auth:org:admin" for p in ctx.permissions
|
||||||
)
|
)
|
||||||
return "auth:admin" in effective_scopes
|
|
||||||
|
|
||||||
|
|
||||||
def is_org_admin(ctx, org_uuid: UUID | None = None) -> bool:
|
|
||||||
"""Check if user has org admin permission.
|
|
||||||
|
|
||||||
If org_uuid is provided, checks if user is admin of that specific org.
|
|
||||||
If org_uuid is None, checks if user is admin of their own org.
|
|
||||||
"""
|
|
||||||
effective_scopes = (
|
|
||||||
{p.scope for p in (ctx.permissions or [])}
|
|
||||||
if ctx.permissions
|
|
||||||
else set(ctx.role.permissions or [])
|
|
||||||
)
|
|
||||||
if "auth:org:admin" not in effective_scopes:
|
|
||||||
return False
|
|
||||||
if org_uuid is None:
|
|
||||||
return True
|
|
||||||
# User must belong to the target org (via their role)
|
|
||||||
return ctx.org.uuid == org_uuid
|
|
||||||
|
|
||||||
|
|
||||||
def can_manage_org(ctx, org_uuid: UUID) -> bool:
|
def can_manage_org(ctx, org_uuid: UUID) -> bool:
|
||||||
"""Check if user can manage the specified organization."""
|
return master_admin(ctx) or org_admin(ctx, org_uuid)
|
||||||
return is_global_admin(ctx) or is_org_admin(ctx, org_uuid)
|
|
||||||
|
|
||||||
|
|
||||||
@app.exception_handler(ValueError)
|
@app.exception_handler(ValueError)
|
||||||
@@ -99,42 +79,38 @@ async def admin_list_orgs(request: Request, auth=AUTH_COOKIE):
|
|||||||
host=request.headers.get("host"),
|
host=request.headers.get("host"),
|
||||||
)
|
)
|
||||||
orgs = list(db.data().orgs.values())
|
orgs = list(db.data().orgs.values())
|
||||||
if not is_global_admin(ctx):
|
if not master_admin(ctx):
|
||||||
# Org admins can only see their own organization
|
# Org admins can only see their own organization
|
||||||
orgs = [o for o in orgs if o.uuid == ctx.org.uuid]
|
orgs = [o for o in orgs if o.uuid == ctx.org.uuid]
|
||||||
|
|
||||||
def role_to_dict(r):
|
def org_to_dict(o):
|
||||||
return {
|
|
||||||
"uuid": str(r.uuid),
|
|
||||||
"org": str(r.org),
|
|
||||||
"display_name": r.display_name,
|
|
||||||
"permissions": list(r.permissions.keys()),
|
|
||||||
}
|
|
||||||
|
|
||||||
async def org_to_dict(o):
|
|
||||||
users = db.get_organization_users(o.uuid)
|
users = db.get_organization_users(o.uuid)
|
||||||
return {
|
return {
|
||||||
"uuid": str(o.uuid),
|
"uuid": o.uuid,
|
||||||
"display_name": o.display_name,
|
"display_name": o.display_name,
|
||||||
"permissions": {
|
"permissions": {p.uuid for p in o.permissions},
|
||||||
pid for pid, p in db.data().permissions.items() if o.uuid in p.orgs
|
|
||||||
},
|
|
||||||
"roles": [
|
"roles": [
|
||||||
role_to_dict(r) for r in db.data().roles.values() if r.org == o.uuid
|
{
|
||||||
|
"uuid": r.uuid,
|
||||||
|
"org": r.org_uuid,
|
||||||
|
"display_name": r.display_name,
|
||||||
|
"permissions": list(r.permissions.keys()),
|
||||||
|
}
|
||||||
|
for r in o.roles
|
||||||
],
|
],
|
||||||
"users": [
|
"users": [
|
||||||
{
|
{
|
||||||
"uuid": str(u.uuid),
|
"uuid": u.uuid,
|
||||||
"display_name": u.display_name,
|
"display_name": u.display_name,
|
||||||
"role": role_name,
|
"role": role_name,
|
||||||
"visits": u.visits,
|
"visits": u.visits,
|
||||||
"last_seen": u.last_seen.isoformat() if u.last_seen else None,
|
"last_seen": u.last_seen,
|
||||||
}
|
}
|
||||||
for (u, role_name) in users
|
for (u, role_name) in users
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
return [await org_to_dict(o) for o in orgs]
|
return MsgspecResponse([org_to_dict(o) for o in orgs])
|
||||||
|
|
||||||
|
|
||||||
@app.post("/orgs")
|
@app.post("/orgs")
|
||||||
@@ -283,7 +259,8 @@ async def admin_create_role(
|
|||||||
perms = payload.get("permissions") or []
|
perms = payload.get("permissions") or []
|
||||||
if org_uuid not in db.data().orgs:
|
if org_uuid not in db.data().orgs:
|
||||||
raise HTTPException(status_code=404, detail="Organization not found")
|
raise HTTPException(status_code=404, detail="Organization not found")
|
||||||
grantable = {pid for pid, p in db.data().permissions.items() if org_uuid in p.orgs}
|
org = db.data().orgs[org_uuid]
|
||||||
|
grantable = {p.uuid for p in org.permissions}
|
||||||
|
|
||||||
# Normalize permission IDs to UUIDs
|
# Normalize permission IDs to UUIDs
|
||||||
permission_uuids: set[UUID] = set()
|
permission_uuids: set[UUID] = set()
|
||||||
@@ -324,7 +301,7 @@ async def admin_update_role_name(
|
|||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
role = db.data().roles.get(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if not role or role.org != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
display_name = payload.get("display_name")
|
display_name = payload.get("display_name")
|
||||||
@@ -356,7 +333,7 @@ async def admin_add_role_permission(
|
|||||||
)
|
)
|
||||||
|
|
||||||
role = db.data().roles.get(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if not role or role.org != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
# Verify permission exists and org can grant it
|
# Verify permission exists and org can grant it
|
||||||
@@ -391,7 +368,7 @@ async def admin_remove_role_permission(
|
|||||||
)
|
)
|
||||||
|
|
||||||
role = db.data().roles.get(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if not role or role.org != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
# Sanity check: prevent admin from removing their own access
|
# Sanity check: prevent admin from removing their own access
|
||||||
@@ -432,7 +409,7 @@ async def admin_delete_role(
|
|||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
role = db.data().roles.get(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if not role or role.org != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
# Sanity check: prevent admin from deleting their own role
|
# Sanity check: prevent admin from deleting their own role
|
||||||
@@ -468,8 +445,11 @@ async def admin_create_user(
|
|||||||
if not display_name or not role_name:
|
if not display_name or not role_name:
|
||||||
raise ValueError("display_name and role are required")
|
raise ValueError("display_name and role are required")
|
||||||
|
|
||||||
roles = [r for r in db.data().roles.values() if r.org == org_uuid]
|
org = db.data().orgs[org_uuid]
|
||||||
role_obj = next((r for r in roles if r.display_name == role_name), None)
|
role_obj = next(
|
||||||
|
(r for r in org.roles if r.display_name == role_name),
|
||||||
|
None,
|
||||||
|
)
|
||||||
if not role_obj:
|
if not role_obj:
|
||||||
raise ValueError("Role not found in organization")
|
raise ValueError("Role not found in organization")
|
||||||
user = UserDC.create(
|
user = UserDC.create(
|
||||||
@@ -507,7 +487,7 @@ async def admin_update_user_role(
|
|||||||
raise ValueError("User not found")
|
raise ValueError("User not found")
|
||||||
if user_org.uuid != org_uuid:
|
if user_org.uuid != org_uuid:
|
||||||
raise ValueError("User does not belong to this organization")
|
raise ValueError("User does not belong to this organization")
|
||||||
roles = [r for r in db.data().roles.values() if r.org == org_uuid]
|
roles = user_org.roles
|
||||||
if not any(r.display_name == new_role for r in roles):
|
if not any(r.display_name == new_role for r in roles):
|
||||||
raise ValueError("Role not found in organization")
|
raise ValueError("Role not found in organization")
|
||||||
|
|
||||||
@@ -572,11 +552,7 @@ async def admin_create_user_registration_link(
|
|||||||
url = hostutil.reset_link_url(token)
|
url = hostutil.reset_link_url(token)
|
||||||
return {
|
return {
|
||||||
"url": url,
|
"url": url,
|
||||||
"expires": (
|
"expires": format_datetime(expiry),
|
||||||
expiry.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if expiry.tzinfo
|
|
||||||
else expiry.replace(tzinfo=timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -604,120 +580,39 @@ async def admin_get_user_detail(
|
|||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
user = db.data().users.get(user_uuid)
|
user = db.data().users.get(user_uuid)
|
||||||
user_creds = [c for c in db.data().credentials.values() if c.user == user_uuid]
|
normalized_host = hostutil.normalize_host(request.headers.get("host"))
|
||||||
creds: list[dict] = []
|
|
||||||
aaguids: set[str] = set()
|
return MsgspecResponse(
|
||||||
for c in user_creds:
|
{
|
||||||
aaguid_str = str(c.aaguid)
|
"display_name": user.display_name,
|
||||||
aaguids.add(aaguid_str)
|
"org": {"display_name": user_org.display_name},
|
||||||
creds.append(
|
"role": role_name,
|
||||||
{
|
"visits": user.visits,
|
||||||
"credential": str(c.uuid),
|
"created_at": user.created_at,
|
||||||
"aaguid": aaguid_str,
|
"last_seen": user.last_seen,
|
||||||
"created_at": (
|
"credentials": [
|
||||||
c.created_at.astimezone(timezone.utc)
|
{
|
||||||
.isoformat()
|
"credential": c.uuid,
|
||||||
.replace("+00:00", "Z")
|
"aaguid": c.aaguid,
|
||||||
if c.created_at.tzinfo
|
"created_at": c.created_at,
|
||||||
else c.created_at.replace(tzinfo=timezone.utc)
|
"last_used": c.last_used,
|
||||||
.isoformat()
|
"last_verified": c.last_verified,
|
||||||
.replace("+00:00", "Z")
|
"sign_count": c.sign_count,
|
||||||
),
|
}
|
||||||
"last_used": (
|
for c in user.credentials
|
||||||
c.last_used.astimezone(timezone.utc)
|
],
|
||||||
.isoformat()
|
"aaguid_info": aaguid_mod.filter(c.aaguid for c in user.credentials),
|
||||||
.replace("+00:00", "Z")
|
"sessions": [
|
||||||
if c.last_used and c.last_used.tzinfo
|
ApiSession.from_db(
|
||||||
else (
|
s,
|
||||||
c.last_used.replace(tzinfo=timezone.utc)
|
current_key=auth,
|
||||||
.isoformat()
|
normalized_host=normalized_host,
|
||||||
.replace("+00:00", "Z")
|
expires_delta=EXPIRES,
|
||||||
if c.last_used
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
),
|
|
||||||
"last_verified": (
|
|
||||||
c.last_verified.astimezone(timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if c.last_verified and c.last_verified.tzinfo
|
|
||||||
else (
|
|
||||||
c.last_verified.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if c.last_verified
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
if c.last_verified
|
for s in user.sessions
|
||||||
else None,
|
],
|
||||||
"sign_count": c.sign_count,
|
}
|
||||||
}
|
)
|
||||||
)
|
|
||||||
|
|
||||||
aaguid_info = aaguid_mod.filter(aaguids)
|
|
||||||
|
|
||||||
# Get sessions for the user
|
|
||||||
normalized_request_host = hostutil.normalize_host(request.headers.get("host"))
|
|
||||||
session_records = [s for s in db.data().sessions.values() if s.user == user_uuid]
|
|
||||||
current_session_key = auth
|
|
||||||
sessions_payload: list[dict] = []
|
|
||||||
for entry in session_records:
|
|
||||||
renewed = entry.expiry - EXPIRES
|
|
||||||
sessions_payload.append(
|
|
||||||
{
|
|
||||||
"id": entry.key,
|
|
||||||
"credential": str(entry.credential),
|
|
||||||
"host": entry.host,
|
|
||||||
"ip": entry.ip,
|
|
||||||
"user_agent": useragent.compact_user_agent(entry.user_agent),
|
|
||||||
"last_renewed": (
|
|
||||||
renewed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if renewed.tzinfo
|
|
||||||
else renewed.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
),
|
|
||||||
"is_current": entry.key == current_session_key,
|
|
||||||
"is_current_host": bool(
|
|
||||||
normalized_request_host
|
|
||||||
and entry.host
|
|
||||||
and entry.host == normalized_request_host
|
|
||||||
),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"display_name": user.display_name,
|
|
||||||
"org": {"display_name": user_org.display_name},
|
|
||||||
"role": role_name,
|
|
||||||
"visits": user.visits,
|
|
||||||
"created_at": (
|
|
||||||
user.created_at.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if user.created_at and user.created_at.tzinfo
|
|
||||||
else (
|
|
||||||
user.created_at.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if user.created_at
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
),
|
|
||||||
"last_seen": (
|
|
||||||
user.last_seen.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if user.last_seen and user.last_seen.tzinfo
|
|
||||||
else (
|
|
||||||
user.last_seen.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if user.last_seen
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
),
|
|
||||||
"credentials": creds,
|
|
||||||
"aaguid_info": aaguid_info,
|
|
||||||
"sessions": sessions_payload,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.patch("/orgs/{org_uuid}/users/{user_uuid}/display-name")
|
@app.patch("/orgs/{org_uuid}/users/{user_uuid}/display-name")
|
||||||
@@ -808,7 +703,7 @@ async def admin_delete_user_session(
|
|||||||
)
|
)
|
||||||
|
|
||||||
target_session = db.data().sessions.get(session_id)
|
target_session = db.data().sessions.get(session_id)
|
||||||
if not target_session or target_session.user != user_uuid:
|
if not target_session or target_session.user_uuid != user_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Session not found")
|
raise HTTPException(status_code=404, detail="Session not found")
|
||||||
|
|
||||||
db.delete_session(session_id, ctx=ctx)
|
db.delete_session(session_id, ctx=ctx)
|
||||||
@@ -821,14 +716,6 @@ async def admin_delete_user_session(
|
|||||||
# -------------------- Permissions (global) --------------------
|
# -------------------- Permissions (global) --------------------
|
||||||
|
|
||||||
|
|
||||||
def _perm_to_dict(p):
|
|
||||||
"""Convert Permission to dict, omitting domain if None."""
|
|
||||||
d = {"uuid": str(p.uuid), "scope": p.scope, "display_name": p.display_name}
|
|
||||||
if p.domain is not None:
|
|
||||||
d["domain"] = p.domain
|
|
||||||
return d
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_permission_domain(domain: str | None) -> None:
|
def _validate_permission_domain(domain: str | None) -> None:
|
||||||
"""Validate that domain is rp_id or a subdomain of it."""
|
"""Validate that domain is rp_id or a subdomain of it."""
|
||||||
if domain is None:
|
if domain is None:
|
||||||
@@ -919,18 +806,8 @@ async def admin_list_permissions(request: Request, auth=AUTH_COOKIE):
|
|||||||
match=permutil.has_any,
|
match=permutil.has_any,
|
||||||
host=request.headers.get("host"),
|
host=request.headers.get("host"),
|
||||||
)
|
)
|
||||||
perms = list(db.data().permissions.values())
|
perms = db.data().permissions.values() if master_admin(ctx) else ctx.org.permissions
|
||||||
|
return MsgspecResponse([ApiPermission.from_db(p) for p in perms])
|
||||||
# Global admins see all permissions
|
|
||||||
if is_global_admin(ctx):
|
|
||||||
return [_perm_to_dict(p) for p in perms]
|
|
||||||
|
|
||||||
# Org admins only see permissions their org can grant (by UUID)
|
|
||||||
grantable = {
|
|
||||||
pid for pid, p in db.data().permissions.items() if ctx.org.uuid in p.orgs
|
|
||||||
}
|
|
||||||
filtered_perms = [p for p in perms if p.uuid in grantable]
|
|
||||||
return [_perm_to_dict(p) for p in filtered_perms]
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/permissions")
|
@app.post("/permissions")
|
||||||
|
|||||||
+49
-55
@@ -1,6 +1,6 @@
|
|||||||
import logging
|
import logging
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import UTC, datetime, timedelta
|
||||||
|
|
||||||
from fastapi import (
|
from fastapi import (
|
||||||
Depends,
|
Depends,
|
||||||
@@ -14,12 +14,9 @@ from fastapi.responses import JSONResponse
|
|||||||
from fastapi.security import HTTPBearer
|
from fastapi.security import HTTPBearer
|
||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.authsession import (
|
from paskia.authsession import EXPIRES, expires, get_reset
|
||||||
EXPIRES,
|
|
||||||
get_reset,
|
|
||||||
refresh_session_token,
|
|
||||||
)
|
|
||||||
from paskia.fastapi import authz, session, user
|
from paskia.fastapi import authz, session, user
|
||||||
|
from paskia.fastapi.response import MsgspecResponse
|
||||||
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME
|
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME
|
||||||
from paskia.globals import passkey as global_passkey
|
from paskia.globals import passkey as global_passkey
|
||||||
from paskia.util import hostutil, htmlutil, passphrase, userinfo, vitedev
|
from paskia.util import hostutil, htmlutil, passphrase, userinfo, vitedev
|
||||||
@@ -90,44 +87,23 @@ async def validate_token(
|
|||||||
raise
|
raise
|
||||||
renewed = False
|
renewed = False
|
||||||
if auth:
|
if auth:
|
||||||
consumed = EXPIRES - (ctx.session.expiry - datetime.now(timezone.utc))
|
consumed = EXPIRES - (ctx.session.expiry - datetime.now(UTC))
|
||||||
if not timedelta(0) < consumed < _REFRESH_INTERVAL:
|
if not timedelta(0) < consumed < _REFRESH_INTERVAL:
|
||||||
try:
|
db.update_session(
|
||||||
refresh_session_token(
|
auth,
|
||||||
auth,
|
ip=request.client.host if request.client else "",
|
||||||
ip=request.client.host if request.client else "",
|
user_agent=request.headers.get("user-agent") or "",
|
||||||
user_agent=request.headers.get("user-agent") or "",
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
session.set_session_cookie(response, auth)
|
session.set_session_cookie(response, auth)
|
||||||
renewed = True
|
renewed = True
|
||||||
except ValueError:
|
return MsgspecResponse(
|
||||||
# Session disappeared, e.g. due to concurrent logout; global handler will clear
|
{
|
||||||
raise authz.AuthException(
|
"valid": True,
|
||||||
status_code=401, detail="Session expired", mode="login"
|
"renewed": renewed,
|
||||||
)
|
"ctx": userinfo.build_session_context(ctx),
|
||||||
return {
|
}
|
||||||
"valid": True,
|
)
|
||||||
"renewed": renewed,
|
|
||||||
"ctx": userinfo.format_session_context(ctx),
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.get("/token-info")
|
|
||||||
async def token_info(credentials=Depends(bearer_auth)):
|
|
||||||
"""Get reset/device-add token info. Pass token via Bearer header."""
|
|
||||||
token = credentials.credentials
|
|
||||||
if not passphrase.is_well_formed(token):
|
|
||||||
raise HTTPException(400, "Invalid token format")
|
|
||||||
try:
|
|
||||||
reset_token = get_reset(token)
|
|
||||||
except ValueError as e:
|
|
||||||
raise HTTPException(401, str(e))
|
|
||||||
|
|
||||||
u = db.data().users.get(reset_token.user)
|
|
||||||
return {
|
|
||||||
"token_type": reset_token.token_type,
|
|
||||||
"display_name": u.display_name,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.get("/forward")
|
@app.get("/forward")
|
||||||
@@ -170,11 +146,9 @@ async def forward_authentication(
|
|||||||
"Remote-Role": str(ctx.role.uuid),
|
"Remote-Role": str(ctx.role.uuid),
|
||||||
"Remote-Role-Name": ctx.role.display_name,
|
"Remote-Role-Name": ctx.role.display_name,
|
||||||
"Remote-Session-Expires": (
|
"Remote-Session-Expires": (
|
||||||
ctx.session.expiry.astimezone(timezone.utc)
|
ctx.session.expiry.astimezone(UTC).isoformat().replace("+00:00", "Z")
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if ctx.session.expiry.tzinfo
|
if ctx.session.expiry.tzinfo
|
||||||
else ctx.session.expiry.replace(tzinfo=timezone.utc)
|
else ctx.session.expiry.replace(tzinfo=UTC)
|
||||||
.isoformat()
|
.isoformat()
|
||||||
.replace("+00:00", "Z")
|
.replace("+00:00", "Z")
|
||||||
),
|
),
|
||||||
@@ -233,24 +207,44 @@ async def api_user_info(
|
|||||||
detail="Authentication required",
|
detail="Authentication required",
|
||||||
mode="login",
|
mode="login",
|
||||||
)
|
)
|
||||||
ctx = db.get_session_context(auth, request.headers.get("host"))
|
ctx = db.data().session_ctx(auth, request.headers.get("host"))
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise HTTPException(401, "Session expired")
|
raise HTTPException(401, "Session expired")
|
||||||
|
|
||||||
return await userinfo.format_user_info(
|
return MsgspecResponse(
|
||||||
user_uuid=ctx.user.uuid,
|
await userinfo.build_user_info(
|
||||||
auth=auth,
|
user_uuid=ctx.user.uuid,
|
||||||
session_record=ctx.session,
|
auth=auth,
|
||||||
request_host=request.headers.get("host"),
|
session_record=ctx.session,
|
||||||
|
request_host=request.headers.get("host"),
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/token-info")
|
||||||
|
async def token_info(credentials=Depends(bearer_auth)):
|
||||||
|
"""Get reset/device-add token info. Pass token via Bearer header."""
|
||||||
|
token = credentials.credentials
|
||||||
|
if not passphrase.is_well_formed(token):
|
||||||
|
raise HTTPException(400, "Invalid token format")
|
||||||
|
try:
|
||||||
|
reset_token = get_reset(token)
|
||||||
|
except ValueError as e:
|
||||||
|
raise HTTPException(401, str(e))
|
||||||
|
|
||||||
|
u = reset_token.user
|
||||||
|
return {
|
||||||
|
"token_type": reset_token.token_type,
|
||||||
|
"display_name": u.display_name,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@app.post("/logout")
|
@app.post("/logout")
|
||||||
async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
||||||
if not auth:
|
if not auth:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.get_session_context(auth, host)
|
ctx = db.data().session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
with suppress(Exception):
|
with suppress(Exception):
|
||||||
@@ -263,7 +257,7 @@ async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
|||||||
async def api_set_session(
|
async def api_set_session(
|
||||||
request: Request, response: Response, auth=Depends(bearer_auth)
|
request: Request, response: Response, auth=Depends(bearer_auth)
|
||||||
):
|
):
|
||||||
ctx = db.get_session_context(auth.credentials, request.headers.get("host"))
|
ctx = db.data().session_ctx(auth.credentials, request.headers.get("host"))
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise HTTPException(401, "Session expired")
|
raise HTTPException(401, "Session expired")
|
||||||
session.set_session_cookie(response, auth.credentials)
|
session.set_session_cookie(response, auth.credentials)
|
||||||
|
|||||||
@@ -0,0 +1,218 @@
|
|||||||
|
"""Custom access logging middleware for FastAPI/Uvicorn."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from ipaddress import IPv6Address
|
||||||
|
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
from starlette.requests import Request
|
||||||
|
from starlette.responses import Response
|
||||||
|
|
||||||
|
logger = logging.getLogger("paskia.access")
|
||||||
|
|
||||||
|
_RESET = "\033[0m"
|
||||||
|
_STATUS_INFO = "\033[32m" # 1xx (green)
|
||||||
|
_STATUS_OK = "\033[92m" # 2xx (bright green)
|
||||||
|
_STATUS_REDIRECT = "\033[32m" # 3xx (green)
|
||||||
|
_STATUS_CLIENT_ERR = "\033[0;31m" # 4xx (red)
|
||||||
|
_STATUS_SERVER_ERR = "\033[1;31m" # 5xx (bright red)
|
||||||
|
_METHOD_READ = "\033[0;34m" # GET, HEAD, OPTIONS (blue)
|
||||||
|
_METHOD_WRITE = "\033[1;34m" # POST, PUT, DELETE, PATCH (bright blue)
|
||||||
|
_HOST = "\033[1;30m" # hostname (dark grey)
|
||||||
|
_PATH = "\033[0m" # path (default)
|
||||||
|
_TIMING = "\033[2m" # timing (dim)
|
||||||
|
_WS_OPEN = "\033[1;33m" # WebSocket connect (bright yellow)
|
||||||
|
_WS_CLOSE = "\033[0;33m" # WebSocket disconnect (yellow)
|
||||||
|
_WS_STATUS = "\033[1;30m" # WebSocket close status (dark grey)
|
||||||
|
|
||||||
|
|
||||||
|
def format_ipv6_network(ip: str) -> str:
|
||||||
|
"""Format IPv6 address to show only network part (first 64 bits)."""
|
||||||
|
try:
|
||||||
|
addr = IPv6Address(ip)
|
||||||
|
# Get the integer representation and mask to first 64 bits
|
||||||
|
network_int = int(addr) >> 64
|
||||||
|
# Format as IPv6 with trailing ::
|
||||||
|
# Split into 4 groups of 16 bits
|
||||||
|
groups = []
|
||||||
|
for _ in range(4):
|
||||||
|
groups.insert(0, format(network_int & 0xFFFF, "x"))
|
||||||
|
network_int >>= 16
|
||||||
|
# Compress consecutive zero groups
|
||||||
|
result = ":".join(groups) + "::"
|
||||||
|
# Simplify leading zeros in groups and compress
|
||||||
|
return str(IPv6Address(result + "0"))
|
||||||
|
except Exception:
|
||||||
|
return ip
|
||||||
|
|
||||||
|
|
||||||
|
def format_client_ip(ip: str) -> str:
|
||||||
|
"""Format client IP, compressing IPv6 to network part only."""
|
||||||
|
if not ip or ip == "-":
|
||||||
|
return "-"
|
||||||
|
if ":" in ip:
|
||||||
|
return format_ipv6_network(ip)
|
||||||
|
return ip
|
||||||
|
|
||||||
|
|
||||||
|
def status_color(status: int) -> str:
|
||||||
|
"""Return color code based on HTTP status."""
|
||||||
|
if status < 200:
|
||||||
|
return _STATUS_INFO
|
||||||
|
if status < 300:
|
||||||
|
return _STATUS_OK
|
||||||
|
if status < 400:
|
||||||
|
return _STATUS_REDIRECT
|
||||||
|
if status < 500:
|
||||||
|
return _STATUS_CLIENT_ERR
|
||||||
|
return _STATUS_SERVER_ERR
|
||||||
|
|
||||||
|
|
||||||
|
def method_color(method: str) -> str:
|
||||||
|
"""Return color code based on HTTP method."""
|
||||||
|
if method in ("GET", "HEAD", "OPTIONS"):
|
||||||
|
return _METHOD_READ
|
||||||
|
return _METHOD_WRITE
|
||||||
|
|
||||||
|
|
||||||
|
def format_access_log(
|
||||||
|
client: str, status: int, method: str, host: str, path: str, duration_ms: float
|
||||||
|
) -> str:
|
||||||
|
"""Format access log line with colors and aligned fields."""
|
||||||
|
use_color = sys.stderr.isatty()
|
||||||
|
|
||||||
|
# Format components with fixed widths for alignment
|
||||||
|
ip = format_client_ip(client).ljust(15) # IPv4 max 15 chars
|
||||||
|
timing = f"{duration_ms:.0f}ms"
|
||||||
|
method_padded = method.ljust(7) # Longest method is OPTIONS (7)
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
status_str = f"{status_color(status)}{status}{_RESET}"
|
||||||
|
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||||
|
method_str = f"{method_color(method)}{method_padded}{_RESET}"
|
||||||
|
host_str = f"{_HOST}{host}{_RESET}"
|
||||||
|
path_str = f"{_PATH}{path}{_RESET}"
|
||||||
|
else:
|
||||||
|
status_str = str(status)
|
||||||
|
timing_str = timing
|
||||||
|
method_str = method_padded
|
||||||
|
host_str = host
|
||||||
|
path_str = path
|
||||||
|
|
||||||
|
# Format: "IP STATUS METHOD host path TIMING"
|
||||||
|
return f"{ip} {status_str} {method_str} {host_str}{path_str} {timing_str}"
|
||||||
|
|
||||||
|
|
||||||
|
# WebSocket connection counter (mod 100)
|
||||||
|
_ws_counter = 0
|
||||||
|
|
||||||
|
|
||||||
|
def _next_ws_id() -> int:
|
||||||
|
"""Get next WebSocket connection ID (0-99)."""
|
||||||
|
global _ws_counter
|
||||||
|
ws_id = _ws_counter
|
||||||
|
_ws_counter = (_ws_counter + 1) % 100
|
||||||
|
return ws_id
|
||||||
|
|
||||||
|
|
||||||
|
def log_ws_open(client: str, host: str, path: str) -> int:
|
||||||
|
"""Log WebSocket connection open. Returns connection ID for use in close."""
|
||||||
|
use_color = sys.stderr.isatty()
|
||||||
|
ws_id = _next_ws_id()
|
||||||
|
|
||||||
|
ip = format_client_ip(client).ljust(15)
|
||||||
|
id_str = f"{ws_id:02d}".ljust(7) # Align with method field (7 chars)
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
# 🔌 aligned with status (takes ~2 char width), ID aligned with method
|
||||||
|
prefix = f"🔌 {_WS_OPEN}{id_str}{_RESET}"
|
||||||
|
host_str = f"{_HOST}{host}{_RESET}"
|
||||||
|
path_str = f"{_PATH}{path}{_RESET}"
|
||||||
|
else:
|
||||||
|
prefix = f"WS+ {id_str}"
|
||||||
|
host_str = host
|
||||||
|
path_str = path
|
||||||
|
|
||||||
|
logger.info(f"{ip} {prefix} {host_str}{path_str}")
|
||||||
|
return ws_id
|
||||||
|
|
||||||
|
|
||||||
|
# WebSocket close codes to human-readable status
|
||||||
|
WS_CLOSE_CODES = {
|
||||||
|
1000: "ok",
|
||||||
|
1001: "going away",
|
||||||
|
1002: "protocol error",
|
||||||
|
1003: "unsupported",
|
||||||
|
1005: "no status",
|
||||||
|
1006: "abnormal",
|
||||||
|
1007: "invalid data",
|
||||||
|
1008: "policy violation",
|
||||||
|
1009: "too large",
|
||||||
|
1010: "extension required",
|
||||||
|
1011: "server error",
|
||||||
|
1012: "restarting",
|
||||||
|
1013: "try again",
|
||||||
|
1014: "bad gateway",
|
||||||
|
1015: "tls error",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def log_ws_close(
|
||||||
|
client: str, ws_id: int, close_code: int | None, duration_ms: float
|
||||||
|
) -> None:
|
||||||
|
"""Log WebSocket connection close with duration and status."""
|
||||||
|
use_color = sys.stderr.isatty()
|
||||||
|
|
||||||
|
ip = format_client_ip(client).ljust(15)
|
||||||
|
id_str = f"{ws_id:02d}".ljust(7) # Align with method field (7 chars)
|
||||||
|
timing = f"{duration_ms:.0f}ms"
|
||||||
|
|
||||||
|
# Convert close code to status text
|
||||||
|
if close_code is None:
|
||||||
|
status = "closed"
|
||||||
|
else:
|
||||||
|
status = WS_CLOSE_CODES.get(close_code, f"code {close_code}")
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
# 🔌 aligned with status, ID aligned with method
|
||||||
|
prefix = f"🔌 {_WS_CLOSE}{id_str}{_RESET}"
|
||||||
|
status_str = f"{_WS_STATUS}{status}{_RESET}"
|
||||||
|
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||||
|
else:
|
||||||
|
prefix = f"WS- {id_str}"
|
||||||
|
status_str = status
|
||||||
|
timing_str = timing
|
||||||
|
|
||||||
|
logger.info(f"{ip} {prefix} {status_str} {timing_str}")
|
||||||
|
|
||||||
|
|
||||||
|
class AccessLogMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Middleware that logs HTTP requests with custom format."""
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next) -> Response:
|
||||||
|
start = time.perf_counter()
|
||||||
|
response = await call_next(request)
|
||||||
|
duration_ms = (time.perf_counter() - start) * 1000
|
||||||
|
|
||||||
|
client = request.client.host if request.client else "-"
|
||||||
|
host = request.headers.get("host", "-")
|
||||||
|
method = request.method
|
||||||
|
path = request.url.path
|
||||||
|
if request.url.query:
|
||||||
|
path = f"{path}?{request.url.query}"
|
||||||
|
status = response.status_code
|
||||||
|
|
||||||
|
line = format_access_log(client, status, method, host, path, duration_ms)
|
||||||
|
logger.info(line)
|
||||||
|
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
def configure_access_logging():
|
||||||
|
"""Configure the access logger to output to stderr."""
|
||||||
|
handler = logging.StreamHandler(sys.stderr)
|
||||||
|
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||||
|
logger.addHandler(handler)
|
||||||
|
logger.setLevel(logging.INFO)
|
||||||
|
logger.propagate = False
|
||||||
@@ -10,10 +10,16 @@ from fastapi_vue import Frontend
|
|||||||
|
|
||||||
from paskia import globals
|
from paskia import globals
|
||||||
from paskia.db import start_background, stop_background
|
from paskia.db import start_background, stop_background
|
||||||
|
from paskia.db.logging import configure_db_logging
|
||||||
from paskia.fastapi import admin, api, auth_host, ws
|
from paskia.fastapi import admin, api, auth_host, ws
|
||||||
|
from paskia.fastapi.logging import AccessLogMiddleware, configure_access_logging
|
||||||
from paskia.fastapi.session import AUTH_COOKIE
|
from paskia.fastapi.session import AUTH_COOKIE
|
||||||
from paskia.util import hostutil, passphrase, vitedev
|
from paskia.util import hostutil, passphrase, vitedev
|
||||||
|
|
||||||
|
# Configure custom logging
|
||||||
|
configure_access_logging()
|
||||||
|
configure_db_logging()
|
||||||
|
|
||||||
# Vue Frontend static files
|
# Vue Frontend static files
|
||||||
frontend = Frontend(
|
frontend = Frontend(
|
||||||
Path(__file__).parent.parent / "frontend-build",
|
Path(__file__).parent.parent / "frontend-build",
|
||||||
@@ -48,10 +54,11 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
|||||||
# Re-raise to fail fast
|
# Re-raise to fail fast
|
||||||
raise
|
raise
|
||||||
|
|
||||||
# Restore info level logging after startup (suppressed during uvicorn init in dev mode)
|
# Restore uvicorn info logging (suppressed during startup in dev mode)
|
||||||
|
# Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages
|
||||||
if frontend.devmode:
|
if frontend.devmode:
|
||||||
logging.getLogger("uvicorn").setLevel(logging.INFO)
|
logging.getLogger("uvicorn").setLevel(logging.INFO)
|
||||||
logging.getLogger("uvicorn.access").setLevel(logging.INFO)
|
logging.getLogger("uvicorn.error").setLevel(logging.WARNING)
|
||||||
|
|
||||||
await frontend.load()
|
await frontend.load()
|
||||||
await start_background()
|
await start_background()
|
||||||
@@ -67,6 +74,9 @@ app = FastAPI(
|
|||||||
openapi_url=None,
|
openapi_url=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Custom access logging (uvicorn's access_log is disabled)
|
||||||
|
app.add_middleware(AccessLogMiddleware)
|
||||||
|
|
||||||
# Apply redirections to auth-host if configured (deny access to restricted endpoints, remove /auth/)
|
# Apply redirections to auth-host if configured (deny access to restricted endpoints, remove /auth/)
|
||||||
app.middleware("http")(auth_host.redirect_middleware)
|
app.middleware("http")(auth_host.redirect_middleware)
|
||||||
|
|
||||||
|
|||||||
@@ -324,7 +324,7 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
token_str = passphrase.generate()
|
token_str = passphrase.generate()
|
||||||
expiry = expires()
|
expiry = expires()
|
||||||
db.create_reset_token(
|
db.create_reset_token(
|
||||||
user_uuid=cred.user,
|
user_uuid=cred.user_uuid,
|
||||||
passphrase=token_str,
|
passphrase=token_str,
|
||||||
expiry=expiry,
|
expiry=expiry,
|
||||||
token_type="device addition",
|
token_type="device addition",
|
||||||
@@ -333,7 +333,7 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
# Also create a session so the device is logged in
|
# Also create a session so the device is logged in
|
||||||
normalized_host = hostutil.normalize_host(request.host)
|
normalized_host = hostutil.normalize_host(request.host)
|
||||||
session_token = db.login(
|
session_token = db.login(
|
||||||
user_uuid=cred.user,
|
user_uuid=cred.user_uuid,
|
||||||
credential_uuid=cred.uuid,
|
credential_uuid=cred.uuid,
|
||||||
sign_count=new_sign_count,
|
sign_count=new_sign_count,
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
@@ -346,7 +346,7 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
|
|
||||||
normalized_host = hostutil.normalize_host(request.host)
|
normalized_host = hostutil.normalize_host(request.host)
|
||||||
session_token = db.login(
|
session_token = db.login(
|
||||||
user_uuid=cred.user,
|
user_uuid=cred.user_uuid,
|
||||||
credential_uuid=cred.uuid,
|
credential_uuid=cred.uuid,
|
||||||
sign_count=new_sign_count,
|
sign_count=new_sign_count,
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
@@ -359,7 +359,7 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
completed = await remoteauth.instance.complete_request(
|
completed = await remoteauth.instance.complete_request(
|
||||||
token=request.key,
|
token=request.key,
|
||||||
session_token=session_token,
|
session_token=session_token,
|
||||||
user_uuid=cred.user,
|
user_uuid=cred.user_uuid,
|
||||||
credential_uuid=cred.uuid,
|
credential_uuid=cred.uuid,
|
||||||
reset_token=reset_token,
|
reset_token=reset_token,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -10,8 +10,6 @@ display name. If multiple users match, they are listed and the command
|
|||||||
aborts. A new one-time reset link is always created.
|
aborts. A new one-time reset link is always created.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,22 @@
|
|||||||
|
"""FastAPI response utilities for msgspec.Struct serialization."""
|
||||||
|
|
||||||
|
import msgspec
|
||||||
|
from fastapi import Response
|
||||||
|
|
||||||
|
|
||||||
|
class MsgspecResponse(Response):
|
||||||
|
"""Response that uses msgspec for JSON encoding.
|
||||||
|
|
||||||
|
Use this for returning msgspec.Struct, dict, or list with proper serialization.
|
||||||
|
"""
|
||||||
|
|
||||||
|
media_type = "application/json"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
content: msgspec.Struct | dict | list,
|
||||||
|
status_code: int = 200,
|
||||||
|
headers: dict | None = None,
|
||||||
|
):
|
||||||
|
body = msgspec.json.encode(content)
|
||||||
|
super().__init__(content=body, status_code=status_code, headers=headers)
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
from datetime import timezone
|
from datetime import UTC
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import (
|
from fastapi import (
|
||||||
@@ -43,7 +43,7 @@ async def user_update_display_name(
|
|||||||
status_code=401, detail="Authentication Required", mode="login"
|
status_code=401, detail="Authentication Required", mode="login"
|
||||||
)
|
)
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.get_session_context(auth, host)
|
ctx = db.data().session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
@@ -62,7 +62,7 @@ async def api_logout_all(request: Request, response: Response, auth=AUTH_COOKIE)
|
|||||||
if not auth:
|
if not auth:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.get_session_context(auth, host)
|
ctx = db.data().session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
@@ -84,14 +84,14 @@ async def api_delete_session(
|
|||||||
status_code=401, detail="Authentication Required", mode="login"
|
status_code=401, detail="Authentication Required", mode="login"
|
||||||
)
|
)
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.get_session_context(auth, host)
|
ctx = db.data().session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
)
|
)
|
||||||
|
|
||||||
target_session = db.data().sessions.get(session_id)
|
target_session = db.data().sessions.get(session_id)
|
||||||
if not target_session or target_session.user != ctx.user.uuid:
|
if not target_session or target_session.user_uuid != ctx.user.uuid:
|
||||||
raise HTTPException(status_code=404, detail="Session not found")
|
raise HTTPException(status_code=404, detail="Session not found")
|
||||||
|
|
||||||
db.delete_session(session_id, ctx=ctx)
|
db.delete_session(session_id, ctx=ctx)
|
||||||
@@ -141,8 +141,8 @@ async def api_create_link(
|
|||||||
"message": "Registration link generated successfully",
|
"message": "Registration link generated successfully",
|
||||||
"url": url,
|
"url": url,
|
||||||
"expires": (
|
"expires": (
|
||||||
expiry.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
expiry.astimezone(UTC).isoformat().replace("+00:00", "Z")
|
||||||
if expiry.tzinfo
|
if expiry.tzinfo
|
||||||
else expiry.replace(tzinfo=timezone.utc).isoformat().replace("+00:00", "Z")
|
else expiry.replace(tzinfo=UTC).isoformat().replace("+00:00", "Z")
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -38,11 +38,11 @@ async def websocket_register_add(
|
|||||||
f"The reset link for {passkey.instance.rp_name} is invalid or has expired"
|
f"The reset link for {passkey.instance.rp_name} is invalid or has expired"
|
||||||
)
|
)
|
||||||
s = get_reset(reset)
|
s = get_reset(reset)
|
||||||
user_uuid = s.user
|
user_uuid = s.user_uuid
|
||||||
else:
|
else:
|
||||||
# Require recent authentication for adding a new passkey
|
# Require recent authentication for adding a new passkey
|
||||||
ctx = await authz.verify(auth, perm=[], host=host, max_age="5m")
|
ctx = await authz.verify(auth, perm=[], host=host, max_age="5m")
|
||||||
user_uuid = ctx.session.user
|
user_uuid = ctx.session.user_uuid
|
||||||
s = ctx.session
|
s = ctx.session
|
||||||
|
|
||||||
# Get user information and determine effective user_name for this registration
|
# Get user information and determine effective user_name for this registration
|
||||||
@@ -91,7 +91,7 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
|
|||||||
session_user_uuid = None
|
session_user_uuid = None
|
||||||
credential_ids = None
|
credential_ids = None
|
||||||
if auth:
|
if auth:
|
||||||
ctx = db.get_session_context(auth, host)
|
ctx = db.data().session_ctx(auth, host)
|
||||||
if ctx:
|
if ctx:
|
||||||
session_user_uuid = ctx.user.uuid
|
session_user_uuid = ctx.user.uuid
|
||||||
credential_ids = db.get_user_credential_ids(session_user_uuid) or None
|
credential_ids = db.get_user_credential_ids(session_user_uuid) or None
|
||||||
@@ -99,7 +99,7 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
|
|||||||
cred, new_sign_count = await authenticate_chat(ws, origin, credential_ids)
|
cred, new_sign_count = await authenticate_chat(ws, origin, credential_ids)
|
||||||
|
|
||||||
# If reauth mode, verify the credential belongs to the session's user
|
# If reauth mode, verify the credential belongs to the session's user
|
||||||
if session_user_uuid and cred.user != session_user_uuid:
|
if session_user_uuid and cred.user_uuid != session_user_uuid:
|
||||||
raise ValueError("This passkey belongs to a different account")
|
raise ValueError("This passkey belongs to a different account")
|
||||||
|
|
||||||
# Create session and update user/credential in a single transaction
|
# Create session and update user/credential in a single transaction
|
||||||
@@ -114,7 +114,7 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
|
|||||||
raise ValueError(f"Host must be the same as or a subdomain of {rp_id}")
|
raise ValueError(f"Host must be the same as or a subdomain of {rp_id}")
|
||||||
|
|
||||||
token = db.login(
|
token = db.login(
|
||||||
user_uuid=cred.user,
|
user_uuid=cred.user_uuid,
|
||||||
credential_uuid=cred.uuid,
|
credential_uuid=cred.uuid,
|
||||||
sign_count=new_sign_count,
|
sign_count=new_sign_count,
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
@@ -125,7 +125,7 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
|
|||||||
|
|
||||||
await ws.send_json(
|
await ws.send_json(
|
||||||
{
|
{
|
||||||
"user": str(cred.user),
|
"user": str(cred.user_uuid),
|
||||||
"session_token": token,
|
"session_token": token,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ Shared WebSocket utilities for FastAPI endpoints.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
|
|
||||||
import base64url
|
import base64url
|
||||||
@@ -10,6 +11,7 @@ from fastapi import WebSocket, WebSocketDisconnect
|
|||||||
from webauthn.helpers.exceptions import InvalidAuthenticationResponse
|
from webauthn.helpers.exceptions import InvalidAuthenticationResponse
|
||||||
|
|
||||||
from paskia.fastapi import authz
|
from paskia.fastapi import authz
|
||||||
|
from paskia.fastapi.logging import log_ws_close, log_ws_open
|
||||||
from paskia.globals import passkey
|
from paskia.globals import passkey
|
||||||
from paskia.util import pow
|
from paskia.util import pow
|
||||||
|
|
||||||
@@ -19,11 +21,19 @@ def websocket_error_handler(func):
|
|||||||
|
|
||||||
@wraps(func)
|
@wraps(func)
|
||||||
async def wrapper(ws: WebSocket, *args, **kwargs):
|
async def wrapper(ws: WebSocket, *args, **kwargs):
|
||||||
|
client = ws.client.host if ws.client else "-"
|
||||||
|
host = ws.headers.get("host", "-")
|
||||||
|
path = ws.url.path
|
||||||
|
|
||||||
|
start = time.perf_counter()
|
||||||
|
ws_id = log_ws_open(client, host, path)
|
||||||
|
close_code = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await ws.accept()
|
await ws.accept()
|
||||||
return await func(ws, *args, **kwargs)
|
return await func(ws, *args, **kwargs)
|
||||||
except WebSocketDisconnect:
|
except WebSocketDisconnect as e:
|
||||||
pass
|
close_code = e.code
|
||||||
except authz.AuthException as e:
|
except authz.AuthException as e:
|
||||||
await ws.send_json(
|
await ws.send_json(
|
||||||
{
|
{
|
||||||
@@ -36,6 +46,9 @@ def websocket_error_handler(func):
|
|||||||
except Exception:
|
except Exception:
|
||||||
logging.exception("Internal Server Error")
|
logging.exception("Internal Server Error")
|
||||||
await ws.send_json({"status": 500, "detail": "Internal Server Error"})
|
await ws.send_json({"status": 500, "detail": "Internal Server Error"})
|
||||||
|
finally:
|
||||||
|
duration_ms = (time.perf_counter() - start) * 1000
|
||||||
|
log_ws_close(client, ws_id, close_code, duration_ms)
|
||||||
|
|
||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ Or via the CLI entry point (if installed):
|
|||||||
import argparse
|
import argparse
|
||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
import base64url
|
||||||
@@ -154,7 +154,7 @@ async def migrate_from_sql(
|
|||||||
if perm_uuid:
|
if perm_uuid:
|
||||||
new_permissions[perm_uuid] = True
|
new_permissions[perm_uuid] = True
|
||||||
new_role = Role(
|
new_role = Role(
|
||||||
org=role.org_uuid,
|
org_uuid=role.org_uuid,
|
||||||
display_name=role.display_name,
|
display_name=role.display_name,
|
||||||
permissions=new_permissions,
|
permissions=new_permissions,
|
||||||
)
|
)
|
||||||
@@ -172,8 +172,8 @@ async def migrate_from_sql(
|
|||||||
user_key: UUID = legacy_user.uuid
|
user_key: UUID = legacy_user.uuid
|
||||||
new_user = User(
|
new_user = User(
|
||||||
display_name=legacy_user.display_name,
|
display_name=legacy_user.display_name,
|
||||||
role=legacy_user.role_uuid,
|
role_uuid=legacy_user.role_uuid,
|
||||||
created_at=legacy_user.created_at or datetime.now(timezone.utc),
|
created_at=legacy_user.created_at or datetime.now(UTC),
|
||||||
last_seen=legacy_user.last_seen,
|
last_seen=legacy_user.last_seen,
|
||||||
visits=legacy_user.visits,
|
visits=legacy_user.visits,
|
||||||
)
|
)
|
||||||
@@ -190,7 +190,7 @@ async def migrate_from_sql(
|
|||||||
cred_key: UUID = legacy_cred.uuid
|
cred_key: UUID = legacy_cred.uuid
|
||||||
new_cred = Credential(
|
new_cred = Credential(
|
||||||
credential_id=legacy_cred.credential_id,
|
credential_id=legacy_cred.credential_id,
|
||||||
user=legacy_cred.user_uuid,
|
user_uuid=legacy_cred.user_uuid,
|
||||||
aaguid=legacy_cred.aaguid,
|
aaguid=legacy_cred.aaguid,
|
||||||
public_key=legacy_cred.public_key,
|
public_key=legacy_cred.public_key,
|
||||||
sign_count=legacy_cred.sign_count,
|
sign_count=legacy_cred.sign_count,
|
||||||
@@ -217,8 +217,8 @@ async def migrate_from_sql(
|
|||||||
# Already in new format or unknown - try to use as-is
|
# Already in new format or unknown - try to use as-is
|
||||||
session_key = base64url.enc(old_key[:12])
|
session_key = base64url.enc(old_key[:12])
|
||||||
db.sessions[session_key] = Session(
|
db.sessions[session_key] = Session(
|
||||||
user=sess.user_uuid,
|
user_uuid=sess.user_uuid,
|
||||||
credential=sess.credential_uuid,
|
credential_uuid=sess.credential_uuid,
|
||||||
host=sess.host,
|
host=sess.host,
|
||||||
ip=sess.ip,
|
ip=sess.ip,
|
||||||
user_agent=sess.user_agent,
|
user_agent=sess.user_agent,
|
||||||
@@ -241,14 +241,14 @@ async def migrate_from_sql(
|
|||||||
# Already in new format or unknown - truncate to 9 bytes
|
# Already in new format or unknown - truncate to 9 bytes
|
||||||
token_key = old_key[:9]
|
token_key = old_key[:9]
|
||||||
db.reset_tokens[token_key] = ResetToken(
|
db.reset_tokens[token_key] = ResetToken(
|
||||||
user=token.user_uuid,
|
user_uuid=token.user_uuid,
|
||||||
expiry=token.expiry,
|
expiry=token.expiry,
|
||||||
token_type=token.token_type,
|
token_type=token.token_type,
|
||||||
)
|
)
|
||||||
print(f" Migrated {len(token_models)} reset tokens")
|
print(f" Migrated {len(token_models)} reset tokens")
|
||||||
|
|
||||||
# Queue and flush all changes using the transaction mechanism
|
# Queue and flush all changes using the transaction mechanism
|
||||||
with db.transaction("migrate"):
|
with db.transaction("migrate:sql"):
|
||||||
pass # All data already added to _data, transaction commits on exit
|
pass # All data already added to _data, transaction commits on exit
|
||||||
|
|
||||||
await store.flush()
|
await store.flush()
|
||||||
|
|||||||
+26
-19
@@ -9,7 +9,7 @@ DO NOT use this module for new code. Use paskia.db instead.
|
|||||||
|
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from sqlalchemy import (
|
from sqlalchemy import (
|
||||||
@@ -25,11 +25,6 @@ from sqlalchemy.dialects.sqlite import BLOB
|
|||||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||||
|
|
||||||
from paskia.db import (
|
|
||||||
Org,
|
|
||||||
Role,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# Legacy User class for SQL schema (uses 'role_uuid' not 'role')
|
# Legacy User class for SQL schema (uses 'role_uuid' not 'role')
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -71,6 +66,17 @@ class _LegacyRole:
|
|||||||
permissions: list[str] | None = None
|
permissions: list[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy Org class for SQL schema (has mutable permissions/roles lists)
|
||||||
|
@dataclass
|
||||||
|
class _LegacyOrg:
|
||||||
|
"""Org as stored in the old SQL schema with mutable permissions/roles."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
display_name: str
|
||||||
|
permissions: list[str] | None = None
|
||||||
|
roles: list[_LegacyRole] | None = None
|
||||||
|
|
||||||
|
|
||||||
# Legacy Session class for SQL schema (uses 'key' as field, 'user_uuid', 'credential_uuid')
|
# Legacy Session class for SQL schema (uses 'key' as field, 'user_uuid', 'credential_uuid')
|
||||||
@dataclass
|
@dataclass
|
||||||
class _LegacySession:
|
class _LegacySession:
|
||||||
@@ -112,8 +118,8 @@ def _normalize_dt(value: datetime | None) -> datetime | None:
|
|||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if value.tzinfo is None:
|
if value.tzinfo is None:
|
||||||
return value.replace(tzinfo=timezone.utc)
|
return value.replace(tzinfo=UTC)
|
||||||
return value.astimezone(timezone.utc)
|
return value.astimezone(UTC)
|
||||||
|
|
||||||
|
|
||||||
class Base(DeclarativeBase):
|
class Base(DeclarativeBase):
|
||||||
@@ -128,12 +134,13 @@ class OrgModel(Base):
|
|||||||
|
|
||||||
def as_dataclass(self):
|
def as_dataclass(self):
|
||||||
# Base Org without permissions/roles (filled by data accessors)
|
# Base Org without permissions/roles (filled by data accessors)
|
||||||
org = Org(display_name=self.display_name)
|
return _LegacyOrg(
|
||||||
org.uuid = UUID(bytes=self.uuid)
|
uuid=UUID(bytes=self.uuid),
|
||||||
return org
|
display_name=self.display_name,
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dataclass(org: Org):
|
def from_dataclass(org: _LegacyOrg):
|
||||||
return OrgModel(uuid=org.uuid.bytes, display_name=org.display_name)
|
return OrgModel(uuid=org.uuid.bytes, display_name=org.display_name)
|
||||||
|
|
||||||
|
|
||||||
@@ -172,7 +179,7 @@ class UserModel(Base):
|
|||||||
LargeBinary(16), ForeignKey("roles.uuid", ondelete="CASCADE"), nullable=False
|
LargeBinary(16), ForeignKey("roles.uuid", ondelete="CASCADE"), nullable=False
|
||||||
)
|
)
|
||||||
created_at: Mapped[datetime] = mapped_column(
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc)
|
DateTime(timezone=True), default=lambda: datetime.now(UTC)
|
||||||
)
|
)
|
||||||
last_seen: Mapped[datetime | None] = mapped_column(
|
last_seen: Mapped[datetime | None] = mapped_column(
|
||||||
DateTime(timezone=True), nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
@@ -195,7 +202,7 @@ class UserModel(Base):
|
|||||||
uuid=user.uuid.bytes,
|
uuid=user.uuid.bytes,
|
||||||
display_name=user.display_name,
|
display_name=user.display_name,
|
||||||
role_uuid=user.role_uuid.bytes,
|
role_uuid=user.role_uuid.bytes,
|
||||||
created_at=user.created_at or datetime.now(timezone.utc),
|
created_at=user.created_at or datetime.now(UTC),
|
||||||
last_seen=user.last_seen,
|
last_seen=user.last_seen,
|
||||||
visits=user.visits,
|
visits=user.visits,
|
||||||
)
|
)
|
||||||
@@ -215,7 +222,7 @@ class CredentialModel(Base):
|
|||||||
public_key: Mapped[bytes] = mapped_column(BLOB, nullable=False)
|
public_key: Mapped[bytes] = mapped_column(BLOB, nullable=False)
|
||||||
sign_count: Mapped[int] = mapped_column(Integer, nullable=False)
|
sign_count: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||||
created_at: Mapped[datetime] = mapped_column(
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc)
|
DateTime(timezone=True), default=lambda: datetime.now(UTC)
|
||||||
)
|
)
|
||||||
last_used: Mapped[datetime | None] = mapped_column(
|
last_used: Mapped[datetime | None] = mapped_column(
|
||||||
DateTime(timezone=True), nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
@@ -255,7 +262,7 @@ class SessionModel(Base):
|
|||||||
user_agent: Mapped[str] = mapped_column(String(512), nullable=False)
|
user_agent: Mapped[str] = mapped_column(String(512), nullable=False)
|
||||||
renewed: Mapped[datetime] = mapped_column(
|
renewed: Mapped[datetime] = mapped_column(
|
||||||
DateTime(timezone=True),
|
DateTime(timezone=True),
|
||||||
default=lambda: datetime.now(timezone.utc),
|
default=lambda: datetime.now(UTC),
|
||||||
nullable=False,
|
nullable=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -388,7 +395,7 @@ class DB:
|
|||||||
result = await session.execute(select(PermissionModel))
|
result = await session.execute(select(PermissionModel))
|
||||||
return [p.as_dataclass() for p in result.scalars().all()]
|
return [p.as_dataclass() for p in result.scalars().all()]
|
||||||
|
|
||||||
async def list_organizations(self) -> list[Org]:
|
async def list_organizations(self) -> list[_LegacyOrg]:
|
||||||
async with self.session() as session:
|
async with self.session() as session:
|
||||||
# Load all orgs
|
# Load all orgs
|
||||||
orgs_result = await session.execute(select(OrgModel))
|
orgs_result = await session.execute(select(OrgModel))
|
||||||
@@ -415,13 +422,13 @@ class DB:
|
|||||||
perms_by_role.setdefault(rp.role_uuid, []).append(rp.permission_id)
|
perms_by_role.setdefault(rp.role_uuid, []).append(rp.permission_id)
|
||||||
|
|
||||||
# Build org dataclasses with roles and permission IDs
|
# Build org dataclasses with roles and permission IDs
|
||||||
roles_by_org: dict[bytes, list[Role]] = {}
|
roles_by_org: dict[bytes, list[_LegacyRole]] = {}
|
||||||
for rm in role_models:
|
for rm in role_models:
|
||||||
r_dc = rm.as_dataclass()
|
r_dc = rm.as_dataclass()
|
||||||
r_dc.permissions = perms_by_role.get(rm.uuid, [])
|
r_dc.permissions = perms_by_role.get(rm.uuid, [])
|
||||||
roles_by_org.setdefault(rm.org_uuid, []).append(r_dc)
|
roles_by_org.setdefault(rm.org_uuid, []).append(r_dc)
|
||||||
|
|
||||||
orgs: list[Org] = []
|
orgs: list[_LegacyOrg] = []
|
||||||
for om in org_models:
|
for om in org_models:
|
||||||
o_dc = om.as_dataclass()
|
o_dc = om.as_dataclass()
|
||||||
o_dc.permissions = perms_by_org.get(om.uuid, [])
|
o_dc.permissions = perms_by_org.get(om.uuid, [])
|
||||||
|
|||||||
@@ -19,9 +19,9 @@ The first 3 words of the token serve as the pairing code for manual entry.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
|
from collections.abc import Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import UTC, datetime, timedelta
|
||||||
from typing import Callable
|
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from paskia.util import passphrase, pow
|
from paskia.util import passphrase, pow
|
||||||
@@ -94,7 +94,7 @@ class RemoteAuthManager:
|
|||||||
|
|
||||||
async def _cleanup_expired(self):
|
async def _cleanup_expired(self):
|
||||||
"""Remove expired requests and notify waiting clients."""
|
"""Remove expired requests and notify waiting clients."""
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
expired_keys = []
|
expired_keys = []
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
for key, req in self._requests.items():
|
for key, req in self._requests.items():
|
||||||
@@ -123,7 +123,7 @@ class RemoteAuthManager:
|
|||||||
Returns:
|
Returns:
|
||||||
(code, expiry) - The 3-word passphrase code and expiration time
|
(code, expiry) - The 3-word passphrase code and expiration time
|
||||||
"""
|
"""
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
expiry = now + REMOTE_AUTH_LIFETIME
|
expiry = now + REMOTE_AUTH_LIFETIME
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
@@ -160,7 +160,7 @@ class RemoteAuthManager:
|
|||||||
req = self._requests.get(normalized)
|
req = self._requests.get(normalized)
|
||||||
if req is None:
|
if req is None:
|
||||||
return None
|
return None
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
||||||
# Expired
|
# Expired
|
||||||
del self._requests[normalized]
|
del self._requests[normalized]
|
||||||
@@ -331,7 +331,7 @@ class RemoteAuthManager:
|
|||||||
req = self._requests.get(token)
|
req = self._requests.get(token)
|
||||||
if req is None:
|
if req is None:
|
||||||
return None
|
return None
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
||||||
del self._requests[token]
|
del self._requests[token]
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -0,0 +1,110 @@
|
|||||||
|
"""API response utilities using msgspec for JSON serialization.
|
||||||
|
|
||||||
|
msgspec handles UUID and datetime conversion automatically.
|
||||||
|
API structs inherit from db structs with kw_only=True to add uuid/key fields.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
import msgspec
|
||||||
|
|
||||||
|
from paskia.db.structs import Org, Permission, Role, User
|
||||||
|
from paskia.util import useragent
|
||||||
|
|
||||||
|
|
||||||
|
def _utc_datetime(dt: datetime | None) -> datetime | None:
|
||||||
|
"""Convert datetime to UTC, handling both aware and naive datetimes."""
|
||||||
|
if dt is None:
|
||||||
|
return None
|
||||||
|
if dt.tzinfo:
|
||||||
|
return dt.astimezone(UTC)
|
||||||
|
return dt.replace(tzinfo=UTC)
|
||||||
|
|
||||||
|
|
||||||
|
def format_datetime(dt: datetime | None) -> str | None:
|
||||||
|
"""Format a datetime to ISO 8601 string with Z suffix for UTC."""
|
||||||
|
if dt is None:
|
||||||
|
return None
|
||||||
|
utc_dt = _utc_datetime(dt)
|
||||||
|
return utc_dt.isoformat().replace("+00:00", "Z") if utc_dt else None
|
||||||
|
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
# API structs - inherit from db structs, add uuid for serialization
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class ApiUser(User, kw_only=True):
|
||||||
|
"""User with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, u: User) -> "ApiUser":
|
||||||
|
return cls(uuid=u.uuid, **msgspec.structs.asdict(u))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiOrg(Org, kw_only=True):
|
||||||
|
"""Org with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, o: Org) -> "ApiOrg":
|
||||||
|
return cls(uuid=o.uuid, **msgspec.structs.asdict(o))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiRole(Role, kw_only=True):
|
||||||
|
"""Role with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, r: Role) -> "ApiRole":
|
||||||
|
return cls(uuid=r.uuid, **msgspec.structs.asdict(r))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiPermission(Permission, kw_only=True):
|
||||||
|
"""Permission with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, p: Permission) -> "ApiPermission":
|
||||||
|
return cls(uuid=p.uuid, **msgspec.structs.asdict(p))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiSession(msgspec.Struct):
|
||||||
|
"""Session for API responses with computed fields."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
credential_uuid: UUID = msgspec.field(name="credential")
|
||||||
|
host: str
|
||||||
|
ip: str
|
||||||
|
user_agent: str
|
||||||
|
last_renewed: datetime
|
||||||
|
is_current: bool = False
|
||||||
|
is_current_host: bool = False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(
|
||||||
|
cls,
|
||||||
|
s, # Session
|
||||||
|
*,
|
||||||
|
current_key: str,
|
||||||
|
normalized_host: str | None,
|
||||||
|
expires_delta, # timedelta
|
||||||
|
) -> "ApiSession":
|
||||||
|
return cls(
|
||||||
|
id=s.key,
|
||||||
|
credential_uuid=s.credential_uuid,
|
||||||
|
host=s.host,
|
||||||
|
ip=s.ip,
|
||||||
|
user_agent=useragent.compact_user_agent(s.user_agent),
|
||||||
|
last_renewed=s.expiry - expires_delta,
|
||||||
|
is_current=s.key == current_key,
|
||||||
|
is_current_host=bool(
|
||||||
|
normalized_host and s.host and s.host == normalized_host
|
||||||
|
),
|
||||||
|
)
|
||||||
@@ -40,4 +40,4 @@ async def session_context(auth: str | None, host: str | None = None):
|
|||||||
if not auth:
|
if not auth:
|
||||||
return None
|
return None
|
||||||
normalized_host = normalize_host(host) if host else None
|
normalized_host = normalize_host(host) if host else None
|
||||||
return db.get_session_context(auth, normalized_host)
|
return db.data().session_ctx(auth, normalized_host)
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
"""Utility functions for session validation and checking."""
|
"""Utility functions for session validation and checking."""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
from paskia.db import SessionContext
|
from paskia.db import SessionContext
|
||||||
@@ -34,5 +34,5 @@ def check_session_age(ctx: SessionContext, max_age: str | None) -> bool:
|
|||||||
else:
|
else:
|
||||||
auth_time = ctx.session.expiry - EXPIRES
|
auth_time = ctx.session.expiry - EXPIRES
|
||||||
|
|
||||||
time_since_auth = datetime.now(timezone.utc) - auth_time
|
time_since_auth = datetime.now(UTC) - auth_time
|
||||||
return time_since_auth <= max_age_delta
|
return time_since_auth <= max_age_delta
|
||||||
|
|||||||
+35
-82
@@ -1,107 +1,60 @@
|
|||||||
"""User information formatting and retrieval logic."""
|
"""User information formatting and retrieval logic."""
|
||||||
|
|
||||||
from datetime import timezone
|
|
||||||
|
|
||||||
from paskia import aaguid, db
|
from paskia import aaguid, db
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
from paskia.db import SessionContext
|
from paskia.db import SessionContext
|
||||||
from paskia.util import hostutil, permutil, useragent
|
from paskia.util import hostutil, permutil
|
||||||
|
from paskia.util.apistructs import ApiSession
|
||||||
|
|
||||||
|
|
||||||
def _format_datetime(dt):
|
def build_session_context(ctx: SessionContext) -> dict:
|
||||||
"""Format a datetime object to ISO 8601 string with UTC timezone."""
|
"""Build session context dict from SessionContext."""
|
||||||
if dt is None:
|
|
||||||
return None
|
|
||||||
if dt.tzinfo:
|
|
||||||
return dt.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
else:
|
|
||||||
return dt.replace(tzinfo=timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
|
|
||||||
|
|
||||||
def format_session_context(ctx: SessionContext) -> dict:
|
|
||||||
"""Format SessionContext for JSON response."""
|
|
||||||
return {
|
return {
|
||||||
"user": {
|
"user": {"uuid": ctx.user.uuid, "display_name": ctx.user.display_name},
|
||||||
"uuid": str(ctx.user.uuid),
|
"org": {"uuid": ctx.org.uuid, "display_name": ctx.org.display_name},
|
||||||
"display_name": ctx.user.display_name,
|
"role": {"uuid": ctx.role.uuid, "display_name": ctx.role.display_name},
|
||||||
},
|
|
||||||
"org": {
|
|
||||||
"uuid": str(ctx.org.uuid),
|
|
||||||
"display_name": ctx.org.display_name,
|
|
||||||
},
|
|
||||||
"role": {
|
|
||||||
"uuid": str(ctx.role.uuid),
|
|
||||||
"display_name": ctx.role.display_name,
|
|
||||||
},
|
|
||||||
"permissions": [p.scope for p in ctx.permissions],
|
"permissions": [p.scope for p in ctx.permissions],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
async def format_user_info(
|
async def build_user_info(
|
||||||
*,
|
*,
|
||||||
user_uuid,
|
user_uuid,
|
||||||
auth: str,
|
auth: str,
|
||||||
session_record,
|
session_record,
|
||||||
request_host: str | None,
|
request_host: str | None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Format complete user information for authenticated users."""
|
"""Build user info dict for authenticated users."""
|
||||||
ctx = await permutil.session_context(auth, request_host)
|
ctx = await permutil.session_context(auth, request_host)
|
||||||
|
user = db.data().users[user_uuid]
|
||||||
|
normalized_host = hostutil.normalize_host(request_host)
|
||||||
|
|
||||||
# Fetch and format credentials
|
credentials = sorted(user.credentials, key=lambda c: c.created_at)
|
||||||
user_credentials = [
|
return {
|
||||||
c for c in db.data().credentials.values() if c.user == user_uuid
|
"ctx": build_session_context(ctx),
|
||||||
]
|
"created_at": ctx.user.created_at,
|
||||||
credentials: list[dict] = []
|
"last_seen": ctx.user.last_seen,
|
||||||
user_aaguids: set[str] = set()
|
"visits": ctx.user.visits,
|
||||||
|
"credentials": [
|
||||||
for c in user_credentials:
|
|
||||||
aaguid_str = str(c.aaguid)
|
|
||||||
user_aaguids.add(aaguid_str)
|
|
||||||
credentials.append(
|
|
||||||
{
|
{
|
||||||
"credential": str(c.uuid),
|
"credential": c.uuid,
|
||||||
"aaguid": aaguid_str,
|
"aaguid": c.aaguid,
|
||||||
"created_at": _format_datetime(c.created_at),
|
"created_at": c.created_at,
|
||||||
"last_used": _format_datetime(c.last_used),
|
"last_used": c.last_used,
|
||||||
"last_verified": _format_datetime(c.last_verified),
|
"last_verified": c.last_verified,
|
||||||
"sign_count": c.sign_count,
|
"sign_count": c.sign_count,
|
||||||
"is_current_session": session_record.credential == c.uuid,
|
"is_current_session": session_record.credential == c.uuid,
|
||||||
}
|
}
|
||||||
)
|
for c in credentials
|
||||||
|
],
|
||||||
credentials.sort(key=lambda cred: cred["created_at"])
|
"aaguid_info": aaguid.filter(c.aaguid for c in credentials),
|
||||||
aaguid_info = aaguid.filter(user_aaguids)
|
"sessions": [
|
||||||
|
ApiSession.from_db(
|
||||||
# Format sessions
|
s,
|
||||||
normalized_request_host = hostutil.normalize_host(request_host)
|
current_key=auth,
|
||||||
session_records = [s for s in db.data().sessions.values() if s.user == user_uuid]
|
normalized_host=normalized_host,
|
||||||
current_session_key = auth
|
expires_delta=EXPIRES,
|
||||||
sessions_payload: list[dict] = []
|
)
|
||||||
|
for s in user.sessions
|
||||||
for entry in session_records:
|
],
|
||||||
sessions_payload.append(
|
|
||||||
{
|
|
||||||
"id": entry.key,
|
|
||||||
"credential": str(entry.credential),
|
|
||||||
"host": entry.host,
|
|
||||||
"ip": entry.ip,
|
|
||||||
"user_agent": useragent.compact_user_agent(entry.user_agent),
|
|
||||||
"last_renewed": _format_datetime(entry.expiry - EXPIRES),
|
|
||||||
"is_current": entry.key == current_session_key,
|
|
||||||
"is_current_host": bool(
|
|
||||||
normalized_request_host
|
|
||||||
and entry.host
|
|
||||||
and entry.host == normalized_request_host
|
|
||||||
),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"ctx": format_session_context(ctx),
|
|
||||||
"created_at": _format_datetime(ctx.user.created_at),
|
|
||||||
"last_seen": _format_datetime(ctx.user.last_seen),
|
|
||||||
"visits": ctx.user.visits,
|
|
||||||
"credentials": credentials,
|
|
||||||
"aaguid_info": aaguid_info,
|
|
||||||
"sessions": sessions_payload,
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -74,10 +74,6 @@ filterwarnings = [
|
|||||||
"ignore::DeprecationWarning",
|
"ignore::DeprecationWarning",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.ruff]
|
|
||||||
target-version = "py39"
|
|
||||||
line-length = 88
|
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
select = ["E", "F", "I", "N", "W", "UP", "PLC0415"]
|
select = ["E", "F", "I", "N", "W", "UP", "PLC0415"]
|
||||||
ignore = ["E501"] # Line too long
|
ignore = ["E501"] # Line too long
|
||||||
|
|||||||
+35
-56
@@ -28,17 +28,14 @@ from paskia.db import (
|
|||||||
Permission,
|
Permission,
|
||||||
Role,
|
Role,
|
||||||
User,
|
User,
|
||||||
add_permission_to_org,
|
|
||||||
create_credential,
|
create_credential,
|
||||||
create_org,
|
|
||||||
create_permission,
|
|
||||||
create_reset_token,
|
create_reset_token,
|
||||||
create_role,
|
create_role,
|
||||||
create_session,
|
create_session,
|
||||||
create_user,
|
create_user,
|
||||||
)
|
)
|
||||||
from paskia.db.jsonl import JsonlStore
|
from paskia.db.jsonl import JsonlStore
|
||||||
from paskia.db.operations import DB, _create_token
|
from paskia.db.operations import DB
|
||||||
from paskia.fastapi.mainapp import app
|
from paskia.fastapi.mainapp import app
|
||||||
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
|
||||||
@@ -55,7 +52,13 @@ def event_loop():
|
|||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def test_db() -> AsyncGenerator[DB, None]:
|
async def test_db() -> AsyncGenerator[DB, None]:
|
||||||
"""Create an in-memory JSON database for testing."""
|
"""Create an in-memory JSON database for testing.
|
||||||
|
|
||||||
|
Uses bootstrap() to properly initialize the database with:
|
||||||
|
- auth:admin and auth:org:admin permissions
|
||||||
|
- A default organization with Administration role
|
||||||
|
- An admin user with the Administration role
|
||||||
|
"""
|
||||||
|
|
||||||
with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f:
|
with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f:
|
||||||
db = DB()
|
db = DB()
|
||||||
@@ -64,6 +67,11 @@ async def test_db() -> AsyncGenerator[DB, None]:
|
|||||||
await store.load()
|
await store.load()
|
||||||
ops_db._db = db
|
ops_db._db = db
|
||||||
ops_db._store = store
|
ops_db._store = store
|
||||||
|
# Bootstrap creates the initial permissions, org, role, and admin user
|
||||||
|
ops_db.bootstrap(
|
||||||
|
org_name="Test Organization",
|
||||||
|
admin_name="Test Admin",
|
||||||
|
)
|
||||||
yield db
|
yield db
|
||||||
ops_db._db = None
|
ops_db._db = None
|
||||||
ops_db._store = None
|
ops_db._store = None
|
||||||
@@ -82,49 +90,30 @@ async def passkey_instance() -> Passkey:
|
|||||||
paskia_globals.passkey._instance = None
|
paskia_globals.passkey._instance = None
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
|
||||||
async def test_org(test_db: DB, admin_permission: Permission) -> Org:
|
|
||||||
"""Create a test organization with admin permission."""
|
|
||||||
org = Org.create(display_name="Test Organization")
|
|
||||||
create_org(org)
|
|
||||||
# Grant admin permission to this org
|
|
||||||
add_permission_to_org(org.uuid, admin_permission.uuid)
|
|
||||||
return org
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def admin_permission(test_db: DB) -> Permission:
|
async def admin_permission(test_db: DB) -> Permission:
|
||||||
"""Create the auth:admin permission."""
|
"""Get the auth:admin permission created by bootstrap."""
|
||||||
perm = Permission.create(scope="auth:admin", display_name="Master Admin")
|
return next(p for p in test_db.permissions.values() if p.scope == "auth:admin")
|
||||||
create_permission(perm)
|
|
||||||
return perm
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def org_admin_permission(test_db: DB, test_org: Org) -> Permission:
|
async def org_admin_permission(test_db: DB) -> Permission:
|
||||||
"""Create the auth:org:admin permission."""
|
"""Get the auth:org:admin permission created by bootstrap."""
|
||||||
perm = Permission.create(scope="auth:org:admin", display_name="Organization Admin")
|
return next(p for p in test_db.permissions.values() if p.scope == "auth:org:admin")
|
||||||
create_permission(perm)
|
|
||||||
# Make it grantable by the org
|
|
||||||
add_permission_to_org(test_org.uuid, perm.uuid)
|
|
||||||
return perm
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def test_role(
|
async def test_org(test_db: DB) -> Org:
|
||||||
test_db: DB,
|
"""Get the test organization created by bootstrap."""
|
||||||
test_org: Org,
|
# Bootstrap creates exactly one org
|
||||||
admin_permission: Permission,
|
return next(iter(test_db.orgs.values()))
|
||||||
org_admin_permission: Permission,
|
|
||||||
) -> Role:
|
|
||||||
"""Create a test role with admin permission."""
|
@pytest_asyncio.fixture(scope="function")
|
||||||
role = Role.create(
|
async def test_role(test_db: DB) -> Role:
|
||||||
org=test_org.uuid,
|
"""Get the Administration role created by bootstrap."""
|
||||||
display_name="Test Admin Role",
|
# Bootstrap creates exactly one role (Administration)
|
||||||
permissions={admin_permission.uuid, org_admin_permission.uuid},
|
return next(iter(test_db.roles.values()))
|
||||||
)
|
|
||||||
create_role(role)
|
|
||||||
return role
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -139,14 +128,10 @@ async def user_role(test_db: DB, test_org: Org) -> Role:
|
|||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def test_user(test_db: DB, test_role: Role) -> User:
|
async def test_user(test_db: DB) -> User:
|
||||||
"""Create a test user with admin role."""
|
"""Get the admin user created by bootstrap."""
|
||||||
user = User.create(
|
# Bootstrap creates exactly one user (admin)
|
||||||
display_name="Test Admin",
|
return next(iter(test_db.users.values()))
|
||||||
role=test_role.uuid,
|
|
||||||
)
|
|
||||||
create_user(user)
|
|
||||||
return user
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -193,17 +178,14 @@ async def session_token(
|
|||||||
test_db: DB, test_user: User, test_credential: Credential
|
test_db: DB, test_user: User, test_credential: Credential
|
||||||
) -> 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()
|
return 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=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -211,17 +193,14 @@ async def regular_session_token(
|
|||||||
test_db: DB, regular_user: User, regular_credential: Credential
|
test_db: DB, regular_user: User, regular_credential: Credential
|
||||||
) -> 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()
|
return 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=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
|
|||||||
+8
-15
@@ -12,7 +12,8 @@ These tests cover:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
from datetime import datetime, timezone
|
import secrets
|
||||||
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -36,7 +37,7 @@ from paskia.db import (
|
|||||||
create_session,
|
create_session,
|
||||||
create_user,
|
create_user,
|
||||||
)
|
)
|
||||||
from paskia.db.operations import DB, _create_token
|
from paskia.db.operations import DB
|
||||||
from tests.conftest import auth_headers
|
from tests.conftest import auth_headers
|
||||||
|
|
||||||
# -------------------- Additional Fixtures --------------------
|
# -------------------- Additional Fixtures --------------------
|
||||||
@@ -97,17 +98,14 @@ async def second_org_session_token(
|
|||||||
test_db: DB, second_org_user: User, second_org_credential: Credential
|
test_db: DB, second_org_user: User, second_org_credential: Credential
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Create a session for the second org admin user."""
|
"""Create a session for the second org admin user."""
|
||||||
token = _create_token()
|
return 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=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -132,7 +130,7 @@ async def org_admin_user(test_db: DB, org_admin_role: Role) -> User:
|
|||||||
role=org_admin_role.uuid,
|
role=org_admin_role.uuid,
|
||||||
)
|
)
|
||||||
user.visits = 5
|
user.visits = 5
|
||||||
user.last_seen = datetime.now(timezone.utc)
|
user.last_seen = datetime.now(UTC)
|
||||||
create_user(user)
|
create_user(user)
|
||||||
return user
|
return user
|
||||||
|
|
||||||
@@ -157,17 +155,14 @@ async def org_admin_session_token(
|
|||||||
test_db: DB, org_admin_user: User, org_admin_credential: Credential
|
test_db: DB, org_admin_user: User, org_admin_credential: Credential
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Create a session for the org admin user."""
|
"""Create a session for the org admin user."""
|
||||||
token = _create_token()
|
return 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=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -1170,11 +1165,9 @@ class TestAdminSessions:
|
|||||||
):
|
):
|
||||||
"""Admin should be able to delete a user's session."""
|
"""Admin should be able to delete a user's session."""
|
||||||
# Create an additional session to delete
|
# Create an additional session to delete
|
||||||
extra_token = _create_token()
|
extra_token = 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_token,
|
|
||||||
host="other.host:4401",
|
host="other.host:4401",
|
||||||
ip="192.168.1.1",
|
ip="192.168.1.1",
|
||||||
user_agent="other-agent",
|
user_agent="other-agent",
|
||||||
@@ -1255,7 +1248,7 @@ class TestAdminSessions:
|
|||||||
):
|
):
|
||||||
"""Deleting non-existent session should fail."""
|
"""Deleting non-existent session should fail."""
|
||||||
# Use a valid format but non-existent key
|
# Use a valid format but non-existent key
|
||||||
fake_token = _create_token()
|
fake_token = secrets.token_urlsafe(12)
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/users/{test_user.uuid}/sessions/{fake_token}",
|
f"/auth/api/admin/orgs/{test_org.uuid}/users/{test_user.uuid}/sessions/{fake_token}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
|
|||||||
+5
-7
@@ -10,14 +10,14 @@ These tests cover:
|
|||||||
- /auth/api/set-session - Set session from bearer token
|
- /auth/api/set-session - Set session from bearer token
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timedelta, timezone
|
import secrets
|
||||||
|
from datetime import UTC, datetime, timedelta
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
from paskia.db import create_session, delete_session
|
from paskia.db import create_session, delete_session
|
||||||
from paskia.db.operations import _create_token
|
|
||||||
from paskia.util.passphrase import generate
|
from paskia.util.passphrase import generate
|
||||||
from tests.conftest import auth_headers
|
from tests.conftest import auth_headers
|
||||||
|
|
||||||
@@ -503,7 +503,7 @@ class TestValidateSessionRefresh:
|
|||||||
"""Validate should handle session expiry during refresh attempt."""
|
"""Validate should handle session expiry during refresh attempt."""
|
||||||
|
|
||||||
# Create a token but don't create a session for it
|
# Create a token but don't create a session for it
|
||||||
token = _create_token()
|
token = secrets.token_urlsafe(12)
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
"/auth/api/validate",
|
"/auth/api/validate",
|
||||||
headers={**auth_headers(token), "Host": "localhost:4401"},
|
headers={**auth_headers(token), "Host": "localhost:4401"},
|
||||||
@@ -522,12 +522,10 @@ class TestValidateSessionRefresh:
|
|||||||
"""Validate should return 401 if session disappears during refresh."""
|
"""Validate should return 401 if session disappears during refresh."""
|
||||||
|
|
||||||
# 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()
|
old_expiry = datetime.now(UTC) + EXPIRES - timedelta(minutes=10)
|
||||||
old_expiry = datetime.now(timezone.utc) + EXPIRES - timedelta(minutes=10)
|
token = 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=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
|
|||||||
Reference in New Issue
Block a user