Move get_session_context to DB.session_ctx().

This commit is contained in:
2026-01-28 14:52:18 +00:00
parent 29ea6426fe
commit c8d659b5ca
8 changed files with 76 additions and 101 deletions
+1 -1
View File
@@ -56,7 +56,7 @@ def refresh_session_token(token: str, *, ip: str, user_agent: str):
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 -4
View File
@@ -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
+1 -88
View File
@@ -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.
""" """
@@ -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
@@ -126,92 +125,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)
# ------------------------------------------------------------------------- # -------------------------------------------------------------------------
+64
View File
@@ -4,6 +4,8 @@ from uuid import UUID
import msgspec import msgspec
import uuid7 import uuid7
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)
@@ -270,3 +272,65 @@ 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 = self.users[s.user]
role = self.roles[user.role]
org = self.orgs[role.org]
credential = self.credentials[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 = {
pid for pid, p in self.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
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,
)
+3 -3
View File
@@ -233,7 +233,7 @@ 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")
@@ -250,7 +250,7 @@ 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 +263,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)
+3 -3
View File
@@ -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,7 +84,7 @@ 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"
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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)