Move get_session_context to DB.session_ctx().
This commit is contained in:
@@ -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,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
@@ -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)
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user