Updated database to use async background worker, making changes lock-free synchronous ops.
This commit is contained in:
+314
-347
@@ -13,7 +13,6 @@ A background task periodically flushes queued changes to disk.
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from collections import deque
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timezone
|
||||
@@ -21,6 +20,7 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
import aiofiles
|
||||
import base64url
|
||||
import jsondiff
|
||||
import msgspec
|
||||
@@ -246,7 +246,7 @@ async def _background_loop():
|
||||
# Run cleanup immediately on startup to clear old expired items
|
||||
if _db is not None:
|
||||
_db.cleanup()
|
||||
_db.flush()
|
||||
await _db.flush()
|
||||
|
||||
last_cleanup = datetime.now(timezone.utc)
|
||||
|
||||
@@ -255,18 +255,18 @@ async def _background_loop():
|
||||
await asyncio.sleep(FLUSH_INTERVAL)
|
||||
if _db is not None:
|
||||
# Flush pending changes to disk
|
||||
_db.flush()
|
||||
await _db.flush()
|
||||
|
||||
# Run cleanup less frequently
|
||||
now = datetime.now(timezone.utc)
|
||||
if (now - last_cleanup).total_seconds() >= CLEANUP_INTERVAL:
|
||||
_db.cleanup()
|
||||
_db.flush() # Flush cleanup changes
|
||||
await _db.flush() # Flush cleanup changes
|
||||
last_cleanup = now
|
||||
except asyncio.CancelledError:
|
||||
# Final flush before exit
|
||||
if _db is not None:
|
||||
_db.flush()
|
||||
await _db.flush()
|
||||
break
|
||||
except Exception:
|
||||
_logger.exception("Error in database background loop")
|
||||
@@ -315,7 +315,7 @@ class DB:
|
||||
Changes are queued and periodically flushed to disk by a background task.
|
||||
Each change records the actor (user UUID or system identifier).
|
||||
|
||||
Thread-safety: Uses a lock for concurrent access to the data structure.
|
||||
Thread-safety: Not needed since the app is single-threaded.
|
||||
|
||||
Data structure:
|
||||
{
|
||||
@@ -335,7 +335,6 @@ class DB:
|
||||
self._data: _DatabaseData | None = None
|
||||
self._previous_builtins: dict[str, Any] = {} # For diffing (JSON-compatible)
|
||||
self._pending_changes: deque[_ChangeRecord] = deque()
|
||||
self._lock = threading.RLock() # Reentrant for nested calls
|
||||
self._current_actor: str = "system" # Default actor for changes
|
||||
|
||||
def _empty_data(self) -> _DatabaseData:
|
||||
@@ -357,31 +356,30 @@ class DB:
|
||||
schema evolution), then validates the final state against msgspec
|
||||
structs which become the working copy with proper datetime types.
|
||||
"""
|
||||
with self._lock:
|
||||
data_dict = msgspec.to_builtins(self._empty_data())
|
||||
if self.db_path.exists():
|
||||
try:
|
||||
# Read JSONL file line by line and apply diffs
|
||||
with open(self.db_path, encoding="utf-8") as f:
|
||||
for line_num, line in enumerate(f, 1):
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
change = msgspec.json.decode(line.encode("utf-8"))
|
||||
# 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}")
|
||||
data_dict = msgspec.to_builtins(self._empty_data())
|
||||
if self.db_path.exists():
|
||||
try:
|
||||
# Read JSONL file line by line and apply diffs
|
||||
with open(self.db_path, encoding="utf-8") as f:
|
||||
for line_num, line in enumerate(f, 1):
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
change = msgspec.json.decode(line.encode("utf-8"))
|
||||
# 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}")
|
||||
|
||||
# Validate and convert to msgspec struct (datetime strings -> datetime objects)
|
||||
self._data = _json_decoder.decode(_json_encoder.encode(data_dict))
|
||||
# Store builtins representation for diffing (to_builtins creates a copy)
|
||||
self._previous_builtins = msgspec.to_builtins(self._data)
|
||||
# Validate and convert to msgspec struct (datetime strings -> datetime objects)
|
||||
self._data = _json_decoder.decode(_json_encoder.encode(data_dict))
|
||||
# Store builtins representation for diffing (to_builtins creates a copy)
|
||||
self._previous_builtins = msgspec.to_builtins(self._data)
|
||||
|
||||
def _queue_change(self) -> None:
|
||||
"""Queue a change record for later flush. Must hold lock."""
|
||||
@@ -404,15 +402,14 @@ class DB:
|
||||
# Update previous builtins for next diff
|
||||
self._previous_builtins = current_builtins
|
||||
|
||||
def flush(self) -> None:
|
||||
async def flush(self) -> None:
|
||||
"""Write all pending changes to disk."""
|
||||
with self._lock:
|
||||
if not self._pending_changes:
|
||||
return
|
||||
if not self._pending_changes:
|
||||
return
|
||||
|
||||
# Collect all pending changes
|
||||
changes_to_write = list(self._pending_changes)
|
||||
self._pending_changes.clear()
|
||||
# Collect all pending changes
|
||||
changes_to_write = list(self._pending_changes)
|
||||
self._pending_changes.clear()
|
||||
|
||||
# Write outside the lock to avoid blocking other operations
|
||||
try:
|
||||
@@ -425,32 +422,32 @@ class DB:
|
||||
# Read existing content and append
|
||||
existing_content = ""
|
||||
if self.db_path.exists():
|
||||
existing_content = self.db_path.read_text("utf-8")
|
||||
async with aiofiles.open(self.db_path, encoding="utf-8") as f:
|
||||
existing_content = await f.read()
|
||||
|
||||
new_content = existing_content + "\n".join(lines) + "\n"
|
||||
|
||||
# Write atomically via temp file
|
||||
tmp_path = self.db_path.with_suffix(".tmp")
|
||||
tmp_path.write_text(new_content, "utf-8")
|
||||
async with aiofiles.open(tmp_path, "w", encoding="utf-8") as f:
|
||||
await f.write(new_content)
|
||||
tmp_path.replace(self.db_path)
|
||||
except OSError:
|
||||
_logger.exception("Failed to flush database changes")
|
||||
# Re-queue the changes on failure
|
||||
with self._lock:
|
||||
for change in reversed(changes_to_write):
|
||||
self._pending_changes.appendleft(change)
|
||||
for change in reversed(changes_to_write):
|
||||
self._pending_changes.appendleft(change)
|
||||
|
||||
@contextmanager
|
||||
def session(self, actor: str = "system"):
|
||||
"""Context manager for atomic operations with change queued on exit."""
|
||||
with self._lock:
|
||||
old_actor = self._current_actor
|
||||
self._current_actor = actor
|
||||
try:
|
||||
yield
|
||||
self._queue_change()
|
||||
finally:
|
||||
self._current_actor = old_actor
|
||||
old_actor = self._current_actor
|
||||
self._current_actor = actor
|
||||
try:
|
||||
yield
|
||||
self._queue_change()
|
||||
finally:
|
||||
self._current_actor = old_actor
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Internal helpers (caller must hold lock)
|
||||
@@ -550,11 +547,10 @@ class DB:
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
def get_user_by_uuid(self, user_uuid: UUID) -> User:
|
||||
with self._lock:
|
||||
key = str(user_uuid)
|
||||
if key not in self._data.users:
|
||||
raise ValueError("User not found")
|
||||
return self._build_user(key)
|
||||
key = str(user_uuid)
|
||||
if key not in self._data.users:
|
||||
raise ValueError("User not found")
|
||||
return self._build_user(key)
|
||||
|
||||
def create_user(self, user: User, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -622,11 +618,10 @@ class DB:
|
||||
del self._data.roles[key]
|
||||
|
||||
def get_role(self, role_uuid: UUID) -> Role:
|
||||
with self._lock:
|
||||
key = str(role_uuid)
|
||||
if key not in self._data.roles:
|
||||
raise ValueError("Role not found")
|
||||
return self._build_role(key)
|
||||
key = str(role_uuid)
|
||||
if key not in self._data.roles:
|
||||
raise ValueError("Role not found")
|
||||
return self._build_role(key)
|
||||
|
||||
def get_role_hidden_permissions(self, role_uuid: UUID) -> list[str]:
|
||||
"""Get permission scopes assigned to role but not grantable by its org.
|
||||
@@ -634,30 +629,25 @@ class DB:
|
||||
These are "hidden" permissions that should be preserved when updating
|
||||
the role, so they can become effective again if the org regains access.
|
||||
"""
|
||||
with self._lock:
|
||||
key = str(role_uuid)
|
||||
if key not in self._data.roles:
|
||||
return []
|
||||
key = str(role_uuid)
|
||||
if key not in self._data.roles:
|
||||
return []
|
||||
|
||||
role_data = self._data.roles[key]
|
||||
org_uuid = role_data.org
|
||||
role_data = self._data.roles[key]
|
||||
org_uuid = role_data.org
|
||||
|
||||
# Get org's grantable scopes
|
||||
if org_uuid not in self._data.orgs:
|
||||
return []
|
||||
# Get org's grantable scopes
|
||||
if org_uuid not in self._data.orgs:
|
||||
return []
|
||||
|
||||
org_allowed_scopes = {
|
||||
p.scope
|
||||
for pid, p in self._data.permissions.items()
|
||||
if org_uuid in p.orgs
|
||||
}
|
||||
org_allowed_scopes = {
|
||||
p.scope for pid, p in self._data.permissions.items() if org_uuid in p.orgs
|
||||
}
|
||||
|
||||
# Return scopes in role but not in org
|
||||
return [
|
||||
scope
|
||||
for scope in role_data.permissions
|
||||
if scope not in org_allowed_scopes
|
||||
]
|
||||
# Return scopes in role but not in org
|
||||
return [
|
||||
scope for scope in role_data.permissions if scope not in org_allowed_scopes
|
||||
]
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Credential operations
|
||||
@@ -678,22 +668,20 @@ class DB:
|
||||
)
|
||||
|
||||
def get_credential_by_id(self, credential_id: bytes) -> Credential:
|
||||
with self._lock:
|
||||
for key, c in self._data.credentials.items():
|
||||
if c.credential_id == credential_id:
|
||||
return self._build_credential(key)
|
||||
raise ValueError("Credential not found")
|
||||
for key, c in self._data.credentials.items():
|
||||
if c.credential_id == credential_id:
|
||||
return self._build_credential(key)
|
||||
raise ValueError("Credential not found")
|
||||
|
||||
def get_credentials_by_user_uuid(self, user_uuid: UUID) -> list[bytes]:
|
||||
with self._lock:
|
||||
user_key = str(user_uuid)
|
||||
result: list[bytes] = []
|
||||
for c in self._data.credentials.values():
|
||||
if c.user == user_key:
|
||||
cred_id = c.credential_id
|
||||
if cred_id is not None:
|
||||
result.append(cred_id)
|
||||
return result
|
||||
user_key = str(user_uuid)
|
||||
result: list[bytes] = []
|
||||
for c in self._data.credentials.values():
|
||||
if c.user == user_key:
|
||||
cred_id = c.credential_id
|
||||
if cred_id is not None:
|
||||
result.append(cred_id)
|
||||
return result
|
||||
|
||||
def update_credential(self, credential: Credential, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -745,11 +733,10 @@ class DB:
|
||||
)
|
||||
|
||||
def get_session(self, key: bytes) -> Session | None:
|
||||
with self._lock:
|
||||
key_b64 = _bytes_to_str(key)
|
||||
if key_b64 not in self._data.sessions:
|
||||
return None
|
||||
return self._build_session(key_b64)
|
||||
key_b64 = _bytes_to_str(key)
|
||||
if key_b64 not in self._data.sessions:
|
||||
return None
|
||||
return self._build_session(key_b64)
|
||||
|
||||
def delete_session(self, key: bytes, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -785,17 +772,16 @@ class DB:
|
||||
s.host = host
|
||||
|
||||
def list_sessions_for_user(self, user_uuid: UUID) -> list[Session]:
|
||||
with self._lock:
|
||||
user_key = str(user_uuid)
|
||||
sessions = []
|
||||
for key_b64, s in self._data.sessions.items():
|
||||
if s.user == user_key:
|
||||
key_bytes = _str_to_bytes(key_b64)
|
||||
if key_bytes and key_bytes.startswith(b"sess"):
|
||||
sessions.append(self._build_session(key_b64))
|
||||
# Sort by expiry desc (most recent expiry first)
|
||||
sessions.sort(key=lambda x: x.expiry, reverse=True)
|
||||
return sessions
|
||||
user_key = str(user_uuid)
|
||||
sessions = []
|
||||
for key_b64, s in self._data.sessions.items():
|
||||
if s.user == user_key:
|
||||
key_bytes = _str_to_bytes(key_b64)
|
||||
if key_bytes and key_bytes.startswith(b"sess"):
|
||||
sessions.append(self._build_session(key_b64))
|
||||
# Sort by expiry desc (most recent expiry first)
|
||||
sessions.sort(key=lambda x: x.expiry, reverse=True)
|
||||
return sessions
|
||||
|
||||
def delete_sessions_for_user(self, user_uuid: UUID, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -827,17 +813,16 @@ class DB:
|
||||
)
|
||||
|
||||
def get_reset_token(self, key: bytes) -> ResetToken | None:
|
||||
with self._lock:
|
||||
key_b64 = _bytes_to_str(key)
|
||||
if key_b64 not in self._data.reset_tokens:
|
||||
return None
|
||||
t = self._data.reset_tokens[key_b64]
|
||||
return ResetToken(
|
||||
key=_str_to_bytes(key_b64), # type: ignore[arg-type]
|
||||
user_uuid=UUID(t.user),
|
||||
expiry=t.expiry, # Already datetime
|
||||
token_type=t.token_type,
|
||||
)
|
||||
key_b64 = _bytes_to_str(key)
|
||||
if key_b64 not in self._data.reset_tokens:
|
||||
return None
|
||||
t = self._data.reset_tokens[key_b64]
|
||||
return ResetToken(
|
||||
key=_str_to_bytes(key_b64), # type: ignore[arg-type]
|
||||
user_uuid=UUID(t.user),
|
||||
expiry=t.expiry, # Already datetime
|
||||
token_type=t.token_type,
|
||||
)
|
||||
|
||||
def delete_reset_token(self, key: bytes, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -874,17 +859,15 @@ class DB:
|
||||
break
|
||||
|
||||
def get_organization(self, org_id: str) -> Org:
|
||||
with self._lock:
|
||||
if org_id not in self._data.orgs:
|
||||
raise ValueError("Organization not found")
|
||||
return self._build_org(org_id, include_roles=True)
|
||||
if org_id not in self._data.orgs:
|
||||
raise ValueError("Organization not found")
|
||||
return self._build_org(org_id, include_roles=True)
|
||||
|
||||
def list_organizations(self) -> list[Org]:
|
||||
with self._lock:
|
||||
return [
|
||||
self._build_org(org_uuid, include_roles=True)
|
||||
for org_uuid in self._data.orgs
|
||||
]
|
||||
return [
|
||||
self._build_org(org_uuid, include_roles=True)
|
||||
for org_uuid in self._data.orgs
|
||||
]
|
||||
|
||||
def update_organization(self, org: Org, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -949,52 +932,46 @@ class DB:
|
||||
raise ValueError("Users cannot be transferred to a different organization")
|
||||
|
||||
def get_user_organization(self, user_uuid: UUID) -> tuple[Org, str]:
|
||||
with self._lock:
|
||||
user_key = str(user_uuid)
|
||||
if user_key not in self._data.users:
|
||||
raise ValueError("User not found")
|
||||
role_uuid = self._data.users[user_key].role
|
||||
if role_uuid not in self._data.roles:
|
||||
raise ValueError("Role not found")
|
||||
r = self._data.roles[role_uuid]
|
||||
if r.org not in self._data.orgs:
|
||||
raise ValueError("Organization not found")
|
||||
return self._build_org(r.org), r.display_name
|
||||
user_key = str(user_uuid)
|
||||
if user_key not in self._data.users:
|
||||
raise ValueError("User not found")
|
||||
role_uuid = self._data.users[user_key].role
|
||||
if role_uuid not in self._data.roles:
|
||||
raise ValueError("Role not found")
|
||||
r = self._data.roles[role_uuid]
|
||||
if r.org not in self._data.orgs:
|
||||
raise ValueError("Organization not found")
|
||||
return self._build_org(r.org), r.display_name
|
||||
|
||||
def get_organization_users(self, org_id: str) -> list[tuple[User, str]]:
|
||||
with self._lock:
|
||||
# Get all roles for this org
|
||||
org_role_uuids = {
|
||||
role_uuid
|
||||
for role_uuid, r in self._data.roles.items()
|
||||
if r.org == org_id
|
||||
}
|
||||
return [
|
||||
(self._build_user(user_uuid), self._data.roles[u.role].display_name)
|
||||
for user_uuid, u in self._data.users.items()
|
||||
if u.role in org_role_uuids
|
||||
]
|
||||
# Get all roles for this org
|
||||
org_role_uuids = {
|
||||
role_uuid for role_uuid, r in self._data.roles.items() if r.org == org_id
|
||||
}
|
||||
return [
|
||||
(self._build_user(user_uuid), self._data.roles[u.role].display_name)
|
||||
for user_uuid, u in self._data.users.items()
|
||||
if u.role in org_role_uuids
|
||||
]
|
||||
|
||||
def get_roles_by_organization(self, org_id: str) -> list[Role]:
|
||||
with self._lock:
|
||||
return [
|
||||
self._build_role(role_uuid)
|
||||
for role_uuid, r in self._data.roles.items()
|
||||
if r.org == org_id
|
||||
]
|
||||
return [
|
||||
self._build_role(role_uuid)
|
||||
for role_uuid, r in self._data.roles.items()
|
||||
if r.org == org_id
|
||||
]
|
||||
|
||||
def get_user_role_in_organization(self, user_uuid: UUID, org_id: str) -> str | None:
|
||||
with self._lock:
|
||||
user_key = str(user_uuid)
|
||||
if user_key not in self._data.users:
|
||||
return None
|
||||
role_uuid = self._data.users[user_key].role
|
||||
if role_uuid not in self._data.roles:
|
||||
return None
|
||||
r = self._data.roles[role_uuid]
|
||||
if r.org != org_id:
|
||||
return None
|
||||
return r.display_name
|
||||
user_key = str(user_uuid)
|
||||
if user_key not in self._data.users:
|
||||
return None
|
||||
role_uuid = self._data.users[user_key].role
|
||||
if role_uuid not in self._data.roles:
|
||||
return None
|
||||
r = self._data.roles[role_uuid]
|
||||
if r.org != org_id:
|
||||
return None
|
||||
return r.display_name
|
||||
|
||||
def update_user_role_in_organization(
|
||||
self, user_uuid: UUID, new_role: str, actor: str = "system"
|
||||
@@ -1038,51 +1015,48 @@ class DB:
|
||||
- A UUID string (the primary key)
|
||||
- A scope string (searches for matching scope)
|
||||
"""
|
||||
with self._lock:
|
||||
# First try as UUID key
|
||||
if permission_id in self._data.permissions:
|
||||
p = self._data.permissions[permission_id]
|
||||
# First try as UUID key
|
||||
if permission_id in self._data.permissions:
|
||||
p = self._data.permissions[permission_id]
|
||||
return Permission(
|
||||
uuid=UUID(permission_id),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
# Fall back to scope search
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == permission_id:
|
||||
return Permission(
|
||||
uuid=UUID(permission_id),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
# Fall back to scope search
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == permission_id:
|
||||
return Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
raise ValueError("Permission not found")
|
||||
|
||||
def get_permission_by_scope(self, scope: str) -> Permission | None:
|
||||
"""Get a permission by its scope string."""
|
||||
with self._lock:
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == scope:
|
||||
return Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
return None
|
||||
|
||||
def list_permissions(self) -> list[Permission]:
|
||||
with self._lock:
|
||||
return [
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
for pid, p in self._data.permissions.items()
|
||||
]
|
||||
raise ValueError("Permission not found")
|
||||
|
||||
def get_permission_by_scope(self, scope: str) -> Permission | None:
|
||||
"""Get a permission by its scope string."""
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == scope:
|
||||
return Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
return None
|
||||
|
||||
def list_permissions(self) -> list[Permission]:
|
||||
return [
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
for pid, p in self._data.permissions.items()
|
||||
]
|
||||
|
||||
def update_permission(self, permission: Permission, actor: str = "system") -> None:
|
||||
with self.session(actor):
|
||||
@@ -1188,37 +1162,33 @@ class DB:
|
||||
del orgs[org_id]
|
||||
|
||||
def get_organization_permissions(self, org_id: str) -> list[Permission]:
|
||||
with self._lock:
|
||||
if org_id not in self._data.orgs:
|
||||
raise ValueError("Organization not found")
|
||||
permissions = []
|
||||
for pid, p in self._data.permissions.items():
|
||||
if org_id in p.orgs:
|
||||
permissions.append(
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
if org_id not in self._data.orgs:
|
||||
raise ValueError("Organization not found")
|
||||
permissions = []
|
||||
for pid, p in self._data.permissions.items():
|
||||
if org_id in p.orgs:
|
||||
permissions.append(
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
return permissions
|
||||
)
|
||||
return permissions
|
||||
|
||||
def get_permission_organizations(self, permission_id: str) -> list[Org]:
|
||||
"""Get organizations that can grant a permission.
|
||||
|
||||
permission_id can be a UUID string or a scope string.
|
||||
"""
|
||||
with self._lock:
|
||||
key = self._resolve_permission_key(permission_id)
|
||||
if not key or key not in self._data.permissions:
|
||||
return []
|
||||
org_ids = self._data.permissions[key].orgs
|
||||
return [
|
||||
self._build_org(org_id)
|
||||
for org_id in org_ids
|
||||
if org_id in self._data.orgs
|
||||
]
|
||||
key = self._resolve_permission_key(permission_id)
|
||||
if not key or key not in self._data.permissions:
|
||||
return []
|
||||
org_ids = self._data.permissions[key].orgs
|
||||
return [
|
||||
self._build_org(org_id) for org_id in org_ids if org_id in self._data.orgs
|
||||
]
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Role-permission operations
|
||||
@@ -1272,58 +1242,56 @@ class DB:
|
||||
can grant. Set to False to see all assigned permissions
|
||||
regardless of org restrictions.
|
||||
"""
|
||||
with self._lock:
|
||||
key = str(role_uuid)
|
||||
if key not in self._data.roles:
|
||||
return []
|
||||
role_data = self._data.roles[key]
|
||||
scopes = list(role_data.permissions.keys())
|
||||
key = str(role_uuid)
|
||||
if key not in self._data.roles:
|
||||
return []
|
||||
role_data = self._data.roles[key]
|
||||
scopes = list(role_data.permissions.keys())
|
||||
|
||||
# Get org permissions if filtering
|
||||
org_allowed_scopes = None
|
||||
if filter_by_org:
|
||||
org_uuid = role_data.org
|
||||
if org_uuid in self._data.orgs:
|
||||
org_allowed_scopes = {
|
||||
p.scope
|
||||
for pid, p in self._data.permissions.items()
|
||||
if org_uuid in p.orgs
|
||||
}
|
||||
# Get org permissions if filtering
|
||||
org_allowed_scopes = None
|
||||
if filter_by_org:
|
||||
org_uuid = role_data.org
|
||||
if org_uuid in self._data.orgs:
|
||||
org_allowed_scopes = {
|
||||
p.scope
|
||||
for pid, p in self._data.permissions.items()
|
||||
if org_uuid in p.orgs
|
||||
}
|
||||
|
||||
permissions = []
|
||||
for scope in scopes:
|
||||
# Skip if org filtering is enabled and scope not allowed by org
|
||||
if org_allowed_scopes is not None and scope not in org_allowed_scopes:
|
||||
continue
|
||||
permissions = []
|
||||
for scope in scopes:
|
||||
# Skip if org filtering is enabled and scope not allowed by org
|
||||
if org_allowed_scopes is not None and scope not in org_allowed_scopes:
|
||||
continue
|
||||
|
||||
# Find permission with this scope
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == scope:
|
||||
permissions.append(
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
# Find permission with this scope
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == scope:
|
||||
permissions.append(
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
break
|
||||
return permissions
|
||||
)
|
||||
break
|
||||
return permissions
|
||||
|
||||
def get_permission_roles(self, permission_id: str) -> list[Role]:
|
||||
"""Get roles that have a permission.
|
||||
|
||||
permission_id can be a UUID string or a scope string.
|
||||
"""
|
||||
with self._lock:
|
||||
scope = self._resolve_permission_scope(permission_id)
|
||||
if not scope:
|
||||
return []
|
||||
return [
|
||||
self._build_role(role_uuid)
|
||||
for role_uuid, r in self._data.roles.items()
|
||||
if scope in r.permissions
|
||||
]
|
||||
scope = self._resolve_permission_scope(permission_id)
|
||||
if not scope:
|
||||
return []
|
||||
return [
|
||||
self._build_role(role_uuid)
|
||||
for role_uuid, r in self._data.roles.items()
|
||||
if scope in r.permissions
|
||||
]
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Combined operations
|
||||
@@ -1467,84 +1435,83 @@ class DB:
|
||||
This is the primary method for validating sessions and getting all
|
||||
associated user/org/role/credential data in a single call.
|
||||
"""
|
||||
with self._lock:
|
||||
sess_key_b64 = _bytes_to_str(session_key)
|
||||
if sess_key_b64 not in self._data.sessions:
|
||||
sess_key_b64 = _bytes_to_str(session_key)
|
||||
if sess_key_b64 not in self._data.sessions:
|
||||
return None
|
||||
|
||||
s = self._data.sessions[sess_key_b64]
|
||||
|
||||
# Handle host binding
|
||||
if host is not None:
|
||||
if s.host is None:
|
||||
s.host = host
|
||||
self._queue_change() # Queue change for host binding
|
||||
elif s.host != host:
|
||||
return None
|
||||
|
||||
s = self._data.sessions[sess_key_b64]
|
||||
# Validate user exists
|
||||
user_key = s.user
|
||||
if user_key not in self._data.users:
|
||||
return None
|
||||
|
||||
# Handle host binding
|
||||
if host is not None:
|
||||
if s.host is None:
|
||||
s.host = host
|
||||
self._queue_change() # Queue change for host binding
|
||||
elif s.host != host:
|
||||
return None
|
||||
# Validate role exists
|
||||
role_uuid = self._data.users[user_key].role
|
||||
if role_uuid not in self._data.roles:
|
||||
return None
|
||||
|
||||
# Validate user exists
|
||||
user_key = s.user
|
||||
if user_key not in self._data.users:
|
||||
return None
|
||||
# Validate org exists
|
||||
org_uuid = self._data.roles[role_uuid].org
|
||||
if org_uuid not in self._data.orgs:
|
||||
return None
|
||||
|
||||
# Validate role exists
|
||||
role_uuid = self._data.users[user_key].role
|
||||
if role_uuid not in self._data.roles:
|
||||
return None
|
||||
# Build objects using helpers
|
||||
session_obj = self._build_session(sess_key_b64)
|
||||
user_obj = self._build_user(user_key)
|
||||
role_obj = self._build_role(role_uuid)
|
||||
org_obj = self._build_org(org_uuid)
|
||||
|
||||
# Validate org exists
|
||||
org_uuid = self._data.roles[role_uuid].org
|
||||
if org_uuid not in self._data.orgs:
|
||||
return None
|
||||
# Get credential (optional)
|
||||
cred_uuid = s.credential
|
||||
credential_obj = (
|
||||
self._build_credential(cred_uuid)
|
||||
if cred_uuid in self._data.credentials
|
||||
else None
|
||||
)
|
||||
|
||||
# Build objects using helpers
|
||||
session_obj = self._build_session(sess_key_b64)
|
||||
user_obj = self._build_user(user_key)
|
||||
role_obj = self._build_role(role_uuid)
|
||||
org_obj = self._build_org(org_uuid)
|
||||
# Effective permissions: role permissions (scopes) that the org can grant
|
||||
# role_obj.permissions contains scopes, org_obj.permissions contains scopes
|
||||
from paskia.util.hostutil import normalize_host
|
||||
|
||||
# Get credential (optional)
|
||||
cred_uuid = s.credential
|
||||
credential_obj = (
|
||||
self._build_credential(cred_uuid)
|
||||
if cred_uuid in self._data.credentials
|
||||
else None
|
||||
)
|
||||
|
||||
# Effective permissions: role permissions (scopes) that the org can grant
|
||||
# role_obj.permissions contains scopes, org_obj.permissions contains scopes
|
||||
from paskia.util.hostutil import normalize_host
|
||||
|
||||
normalized_host = normalize_host(host)
|
||||
# Strip port for domain matching (e.g., localhost:4401 -> localhost)
|
||||
host_without_port = (
|
||||
normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
||||
)
|
||||
effective_permissions = []
|
||||
for scope in role_obj.permissions:
|
||||
if scope not in org_obj.permissions:
|
||||
continue
|
||||
# Find the permission by scope
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == scope:
|
||||
# Check domain restriction (compare without port)
|
||||
if p.domain is not None and p.domain != host_without_port:
|
||||
continue
|
||||
effective_permissions.append(
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
normalized_host = normalize_host(host)
|
||||
# Strip port for domain matching (e.g., localhost:4401 -> localhost)
|
||||
host_without_port = (
|
||||
normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
||||
)
|
||||
effective_permissions = []
|
||||
for scope in role_obj.permissions:
|
||||
if scope not in org_obj.permissions:
|
||||
continue
|
||||
# Find the permission by scope
|
||||
for pid, p in self._data.permissions.items():
|
||||
if p.scope == scope:
|
||||
# Check domain restriction (compare without port)
|
||||
if p.domain is not None and p.domain != host_without_port:
|
||||
continue
|
||||
effective_permissions.append(
|
||||
Permission(
|
||||
uuid=UUID(pid),
|
||||
scope=p.scope,
|
||||
display_name=p.display_name,
|
||||
domain=p.domain,
|
||||
)
|
||||
break
|
||||
)
|
||||
break
|
||||
|
||||
return SessionContext(
|
||||
session=session_obj,
|
||||
user=user_obj,
|
||||
org=org_obj,
|
||||
role=role_obj,
|
||||
credential=credential_obj,
|
||||
permissions=effective_permissions or None,
|
||||
)
|
||||
return SessionContext(
|
||||
session=session_obj,
|
||||
user=user_obj,
|
||||
org=org_obj,
|
||||
role=role_obj,
|
||||
credential=credential_obj,
|
||||
permissions=effective_permissions or None,
|
||||
)
|
||||
|
||||
@@ -21,6 +21,7 @@ dependencies = [
|
||||
"user-agents>=2.2.0",
|
||||
"jsondiff>=2.2.1",
|
||||
"msgspec>=0.20.0",
|
||||
"aiofiles>=25.1.0",
|
||||
]
|
||||
requires-python = ">=3.10"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user