Remove most remaining DB getters. Add ws auth chat helper function to avoid repetition, along with the existing register chat in wschat.py.

This commit is contained in:
2026-01-27 20:01:17 +00:00
parent 968964c4c9
commit e8247a2c7f
11 changed files with 144 additions and 287 deletions
+37 -32
View File
@@ -91,7 +91,7 @@ async def admin_list_orgs(request: Request, auth=AUTH_COOKIE):
match=permutil.has_any,
host=request.headers.get("host"),
)
orgs = db.list_organizations()
orgs = list(db.data().orgs.values())
if not is_global_admin(ctx):
# Org admins can only see their own organization
orgs = [o for o in orgs if o.uuid == ctx.org.uuid]
@@ -194,7 +194,7 @@ async def admin_delete_org(org_uuid: UUID, request: Request, auth=AUTH_COOKIE):
# Delete organization-specific permissions
org_perm_pattern = f"org:{str(org_uuid).lower()}"
all_permissions = db.list_permissions()
all_permissions = list(db.data().permissions.values())
for perm in all_permissions:
perm_scope_lower = perm.scope.lower()
# Check if permission contains "org:{uuid}" separated by colons or at boundaries
@@ -226,7 +226,10 @@ async def admin_add_org_permission(
permission_uuid = UUID(permission_id)
except ValueError:
# It's a scope - look up the UUID
perm = db.get_permission_by_scope(permission_id)
perm = next(
(p for p in db.data().permissions.values() if p.scope == permission_id),
None,
)
if not perm:
raise HTTPException(status_code=404, detail="Permission not found")
permission_uuid = perm.uuid
@@ -251,13 +254,16 @@ async def admin_remove_org_permission(
permission_uuid = UUID(permission_id)
except ValueError:
# It's a scope - look up the UUID
perm = db.get_permission_by_scope(permission_id)
perm = next(
(p for p in db.data().permissions.values() if p.scope == permission_id),
None,
)
if not perm:
raise HTTPException(status_code=404, detail="Permission not found")
permission_uuid = perm.uuid
# Guard rail: prevent removing auth:admin from your own org if it would lock you out
perm = db.get_permission(permission_uuid)
perm = db.data().permissions.get(permission_uuid)
if perm and perm.scope == "auth:admin" and ctx.org.uuid == org_uuid:
# Check if any other org grants auth:admin that we're a member of
# (we only know our current org, so this effectively means we can't remove it from our own org)
@@ -294,15 +300,14 @@ async def admin_create_role(
display_name = payload.get("display_name") or "New Role"
perms = payload.get("permissions") or []
org = db.get_organization(org_uuid)
if not org:
if org_uuid not in db.data().orgs:
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}
# Normalize permission IDs to UUIDs
permission_uuids: set[UUID] = set()
for pid in perms:
perm = db.get_permission(UUID(pid))
perm = db.data().permissions.get(UUID(pid))
if not perm:
raise ValueError(f"Permission {pid} not found")
if perm.uuid not in grantable:
@@ -337,8 +342,8 @@ async def admin_update_role_name(
raise authz.AuthException(
status_code=403, detail="Insufficient permissions", mode="forbidden"
)
role = db.get_role(role_uuid)
if role.org != org_uuid:
role = db.data().roles.get(role_uuid)
if not role or role.org != org_uuid:
raise HTTPException(status_code=404, detail="Role not found in organization")
display_name = payload.get("display_name")
@@ -369,12 +374,12 @@ async def admin_add_role_permission(
status_code=403, detail="Insufficient permissions", mode="forbidden"
)
role = db.get_role(role_uuid)
if role.org != org_uuid:
role = db.data().roles.get(role_uuid)
if not role or role.org != org_uuid:
raise HTTPException(status_code=404, detail="Role not found in organization")
# Verify permission exists and org can grant it
perm = db.get_permission(permission_uuid)
perm = db.data().permissions.get(permission_uuid)
if not perm:
raise HTTPException(status_code=404, detail="Permission not found")
if org_uuid not in perm.orgs:
@@ -404,19 +409,19 @@ async def admin_remove_role_permission(
status_code=403, detail="Insufficient permissions", mode="forbidden"
)
role = db.get_role(role_uuid)
if role.org != org_uuid:
role = db.data().roles.get(role_uuid)
if not role or role.org != org_uuid:
raise HTTPException(status_code=404, detail="Role not found in organization")
# Sanity check: prevent admin from removing their own access
perm = db.get_permission(permission_uuid)
perm = db.data().permissions.get(permission_uuid)
if ctx.org.uuid == org_uuid and ctx.role.uuid == role_uuid:
if perm and perm.scope in ["auth:admin", "auth:org:admin"]:
# Check if removing this permission would leave no admin access
remaining_perms = role.permission_set - {permission_uuid}
has_admin = False
for rp_uuid in remaining_perms:
rp = db.get_permission(rp_uuid)
rp = db.data().permissions.get(rp_uuid)
if rp and rp.scope in ["auth:admin", "auth:org:admin"]:
has_admin = True
break
@@ -445,8 +450,8 @@ async def admin_delete_role(
raise authz.AuthException(
status_code=403, detail="Insufficient permissions", mode="forbidden"
)
role = db.get_role(role_uuid)
if role.org != org_uuid:
role = db.data().roles.get(role_uuid)
if not role or role.org != org_uuid:
raise HTTPException(status_code=404, detail="Role not found in organization")
# Sanity check: prevent admin from deleting their own role
@@ -483,7 +488,7 @@ async def admin_create_user(
raise ValueError("display_name and role are required")
from ..db import User as UserDC
roles = db.get_roles_by_organization(org_uuid)
roles = [r for r in db.data().roles.values() if r.org == org_uuid]
role_obj = next((r for r in roles if r.display_name == role_name), None)
if not role_obj:
raise ValueError("Role not found in organization")
@@ -522,7 +527,7 @@ async def admin_update_user_role(
raise ValueError("User not found")
if user_org.uuid != org_uuid:
raise ValueError("User does not belong to this organization")
roles = db.get_roles_by_organization(org_uuid)
roles = [r for r in db.data().roles.values() if r.org == org_uuid]
if not any(r.display_name == new_role for r in roles):
raise ValueError("Role not found in organization")
@@ -533,7 +538,7 @@ async def admin_update_user_role(
# Check if any permission in the new role is an admin permission
has_admin_access = False
for perm_uuid in new_role_obj.permissions:
perm = db.get_permission(perm_uuid)
perm = db.data().permissions.get(perm_uuid)
if perm and perm.scope in ["auth:admin", "auth:org:admin"]:
has_admin_access = True
break
@@ -572,8 +577,8 @@ async def admin_create_user_registration_link(
)
# Check if user has existing credentials
credentials = db.get_credentials_by_user_uuid(user_uuid)
token_type = "user registration" if not credentials else "account recovery"
has_credentials = db.get_user_credential_ids(user_uuid)
token_type = "user registration" if not has_credentials else "account recovery"
token = passphrase.generate()
expiry = reset_expires()
@@ -618,8 +623,8 @@ async def admin_get_user_detail(
raise authz.AuthException(
status_code=403, detail="Insufficient permissions", mode="forbidden"
)
user = db.get_user_by_uuid(user_uuid)
user_creds = db.get_credentials_by_user_uuid(user_uuid)
user = db.data().users.get(user_uuid)
user_creds = [c for c in db.data().credentials.values() if c.user == user_uuid]
creds: list[dict] = []
aaguids: set[str] = set()
for c in user_creds:
@@ -871,7 +876,7 @@ def _check_admin_lockout(
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
# Get all auth:admin permissions
all_perms = db.list_permissions()
all_perms = list(db.data().permissions.values())
admin_perms = [p for p in all_perms if p.scope == "auth:admin"]
# Check if at least one auth:admin would remain accessible
@@ -906,7 +911,7 @@ def _check_admin_lockout_on_delete(perm_uuid: str, current_host: str | None) ->
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
# Get all auth:admin permissions
all_perms = db.list_permissions()
all_perms = list(db.data().permissions.values())
admin_perms = [p for p in all_perms if p.scope == "auth:admin"]
# Check if at least one auth:admin would remain accessible after deletion
@@ -938,7 +943,7 @@ async def admin_list_permissions(request: Request, auth=AUTH_COOKIE):
match=permutil.has_any,
host=request.headers.get("host"),
)
perms = db.list_permissions()
perms = list(db.data().permissions.values())
# Global admins see all permissions
if is_global_admin(ctx):
@@ -997,7 +1002,7 @@ async def admin_update_permission(
)
# Get existing permission
perm = db.get_permission(permission_uuid)
perm = db.data().permissions.get(permission_uuid)
# Update fields that were provided
new_scope = scope if scope is not None else perm.scope
@@ -1045,7 +1050,7 @@ async def admin_rename_permission(
raise ValueError("new_scope required")
# Sanity check: prevent renaming critical permissions
perm = db.get_permission(permission_uuid)
perm = db.data().permissions.get(permission_uuid)
if perm.scope == "auth:admin":
raise ValueError("Cannot rename the master admin permission")
@@ -1086,7 +1091,7 @@ async def admin_delete_permission(
)
# Get the permission to check its scope
perm = db.get_permission(permission_uuid)
perm = db.data().permissions.get(permission_uuid)
# Sanity check: prevent deleting critical permissions if it would lock out admin
if perm.scope == "auth:admin":
+1 -1
View File
@@ -124,7 +124,7 @@ async def token_info(credentials=Depends(bearer_auth)):
except ValueError as e:
raise HTTPException(401, str(e))
u = db.get_user_by_uuid(reset_token.user)
u = db.data().users.get(reset_token.user)
return {
"token_type": reset_token.token_type,
"display_name": u.display_name,
+10 -30
View File
@@ -17,8 +17,8 @@ from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from paskia import db, remoteauth
from paskia.fastapi.session import infodict
from paskia.fastapi.wschat import authenticate_chat
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
from paskia.globals import passkey
from paskia.util import passphrase, pow
# Create a FastAPI subapp for remote auth WebSocket endpoints
@@ -311,30 +311,10 @@ async def websocket_remote_auth_permit(ws: WebSocket):
# Handle authenticate request (no PoW needed - already validated during lookup)
if msg.get("authenticate") and request is not None:
# Generate authentication options
options, webauthn_challenge = passkey.instance.auth_generate_options(
credential_ids=None
)
await ws.send_json({"optionsJSON": options})
# Wait for WebAuthn response
credential = passkey.instance.auth_parse(await ws.receive_json())
# Fetch and verify credential
try:
stored_cred = db.get_credential_by_id(credential.raw_id)
except ValueError:
raise ValueError(
f"This passkey is no longer registered with {passkey.instance.rp_name}"
)
# Verify the credential
passkey.instance.auth_verify(
credential, webauthn_challenge, stored_cred, origin
)
cred = await authenticate_chat(ws, origin)
# Create a session for the REQUESTING device
assert stored_cred.uuid is not None
assert cred.uuid is not None
session_token = None
reset_token = None
@@ -347,7 +327,7 @@ async def websocket_remote_auth_permit(ws: WebSocket):
token_str = passphrase.generate()
expiry = expires()
db.create_reset_token(
user_uuid=stored_cred.user,
user_uuid=cred.user,
passphrase=token_str,
expiry=expiry,
token_type="device addition",
@@ -356,8 +336,8 @@ async def websocket_remote_auth_permit(ws: WebSocket):
# Also create a session so the device is logged in
normalized_host = hostutil.normalize_host(request.host)
session_token = db.login(
user_uuid=stored_cred.user,
credential=stored_cred,
user_uuid=cred.user,
credential=cred,
host=normalized_host,
ip=request.ip,
user_agent=request.user_agent,
@@ -370,8 +350,8 @@ async def websocket_remote_auth_permit(ws: WebSocket):
normalized_host = hostutil.normalize_host(request.host)
session_token = db.login(
user_uuid=stored_cred.user,
credential=stored_cred,
user_uuid=cred.user,
credential=cred,
host=normalized_host,
ip=request.ip,
user_agent=request.user_agent,
@@ -382,8 +362,8 @@ async def websocket_remote_auth_permit(ws: WebSocket):
completed = await remoteauth.instance.complete_request(
token=request.key,
session_token=session_token,
user_uuid=stored_cred.user,
credential_uuid=stored_cred.uuid,
user_uuid=cred.user,
credential_uuid=cred.uuid,
reset_token=reset_token,
)
+10 -48
View File
@@ -1,11 +1,10 @@
from uuid import UUID
from fastapi import FastAPI, WebSocket
from paskia import db
from paskia.authsession import expires, get_reset, get_session
from paskia.fastapi import authz, remote
from paskia.fastapi.session import AUTH_COOKIE, infodict
from paskia.fastapi.wschat import authenticate_chat, register_chat
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
from paskia.globals import passkey
from paskia.util import hostutil, passphrase
@@ -17,24 +16,6 @@ app = FastAPI()
app.mount("/remote-auth", remote.app)
async def register_chat(
ws: WebSocket,
user_uuid: UUID,
user_name: str,
origin: str,
credential_ids: list[bytes] | None = None,
):
"""Generate registration options and send them to the client."""
options, challenge = passkey.instance.reg_generate_options(
user_id=user_uuid,
user_name=user_name,
credential_ids=credential_ids,
)
await ws.send_json({"optionsJSON": options})
response = await ws.receive_json()
return passkey.instance.reg_verify(response, challenge, user_uuid, origin=origin)
@app.websocket("/register")
@websocket_error_handler
async def websocket_register_add(
@@ -65,14 +46,13 @@ async def websocket_register_add(
s = ctx.session
# Get user information and determine effective user_name for this registration
user = db.get_user_by_uuid(user_uuid)
user = db.data().users.get(user_uuid)
user_name = user.display_name
if name is not None:
stripped = name.strip()
if stripped:
user_name = stripped
credentials = db.get_credentials_by_user_uuid(user_uuid)
credential_ids = [c.credential_id for c in credentials] if credentials else None
credential_ids = db.get_user_credential_ids(user_uuid) or None
# WebAuthn registration
credential = await register_chat(ws, user_uuid, user_name, origin, credential_ids)
@@ -114,36 +94,18 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
try:
session = await get_session(auth, host=host)
session_user_uuid = session.user
credentials = db.get_credentials_by_user_uuid(session_user_uuid)
credential_ids = (
[c.credential_id for c in credentials] if credentials else None
)
credential_ids = db.get_user_credential_ids(session_user_uuid) or None
except ValueError:
pass # Invalid/expired session - allow normal authentication
options, challenge = passkey.instance.auth_generate_options(
credential_ids=credential_ids
)
await ws.send_json({"optionsJSON": options})
# Wait for the client to use his authenticator to authenticate
credential = passkey.instance.auth_parse(await ws.receive_json())
# Fetch from the database by credential ID
try:
stored_cred = db.get_credential_by_id(credential.raw_id)
except ValueError:
raise ValueError(
f"This passkey is no longer registered with {passkey.instance.rp_name}"
)
cred = await authenticate_chat(ws, origin, credential_ids)
# If reauth mode, verify the credential belongs to the session's user
if session_user_uuid and stored_cred.user != session_user_uuid:
if session_user_uuid and cred.user != session_user_uuid:
raise ValueError("This passkey belongs to a different account")
# Verify the credential matches the stored data
passkey.instance.auth_verify(credential, challenge, stored_cred, origin)
# Create session and update user/credential in a single transaction
assert stored_cred.uuid is not None
assert cred.uuid is not None
metadata = infodict(ws, "auth")
normalized_host = hostutil.normalize_host(host)
if not normalized_host:
@@ -154,8 +116,8 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
raise ValueError(f"Host must be the same as or a subdomain of {rp_id}")
token = db.login(
user_uuid=stored_cred.user,
credential=stored_cred,
user_uuid=cred.user,
credential=cred,
host=normalized_host,
ip=metadata.get("ip") or "",
user_agent=metadata.get("user_agent") or "",
@@ -164,7 +126,7 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
await ws.send_json(
{
"user": str(stored_cred.user),
"user": str(cred.user),
"session_token": token,
}
)
+58
View File
@@ -0,0 +1,58 @@
"""
WebSocket chat functions for WebAuthn registration and authentication flows.
"""
from uuid import UUID
from fastapi import WebSocket
from paskia import db
from paskia.db import Credential
from paskia.globals import passkey
async def register_chat(
ws: WebSocket,
user_uuid: UUID,
user_name: str,
origin: str,
credential_ids: list[bytes] | None = None,
):
"""Run WebAuthn registration flow and return the verified credential."""
options, challenge = passkey.instance.reg_generate_options(
user_id=user_uuid,
user_name=user_name,
credential_ids=credential_ids,
)
await ws.send_json({"optionsJSON": options})
response = await ws.receive_json()
return passkey.instance.reg_verify(response, challenge, user_uuid, origin=origin)
async def authenticate_chat(
ws: WebSocket,
origin: str,
credential_ids: list[bytes] | None = None,
) -> Credential:
"""Run WebAuthn authentication flow and return the verified credential."""
options, challenge = passkey.instance.auth_generate_options(
credential_ids=credential_ids
)
await ws.send_json({"optionsJSON": options})
authcred = passkey.instance.auth_parse(await ws.receive_json())
cred = next(
(
c
for c in db.data().credentials.values()
if c.credential_id == authcred.raw_id
),
None,
)
if not cred:
raise ValueError(
f"This passkey is no longer registered with {passkey.instance.rp_name}"
)
passkey.instance.auth_verify(authcred, challenge, cred, origin)
return cred