Proxy to another Paskia #5

Open
LeoVasanko wants to merge 10 commits from feature/remote-satellite into main
27 changed files with 1471 additions and 85 deletions
+176
View File
@@ -0,0 +1,176 @@
# Remote Proxy / Satellite
Paskia can serve a configured domain (rp-id) from a **remote** paskia
instance instead of the local database, so latency-sensitive checks
(`/auth/api/forward`, `/auth/api/validate`) answer in ~1 ms even when the
auth server is on another continent. Example: `app2.example.com` runs on
our local host and needs fast local checks, while `app1.example.com` and
`auth.example.com` run far away — all sharing `example.com` as rp-id.
Client applications that used `https://auth.example.com` as their auth
backend (forward-auth checks) only repoint to the local satellite
(`http://127.0.0.1:4401`); both remain usable interchangeably, and the
satellite ultimately uses `auth.example.com`.
Status: **implemented**. The feature lives in `paskia/satellite.py`
(satellite side: replica, sync client, host dispatch, forwarding) and
`paskia/syncfeed.py` + `paskia/fastapi/sync.py` (remote side: change
feed and sync WebSocket). The design review comparing the rejected
alternatives is at the end of this document.
## Configuration
A domain becomes remote in the admin domains UI (master admin, on the
primary server's auth host — this configuration itself never touches a
remote): enable *Remote instance* and set the remote URL and sync token.
In the stored config (`DomainConfig.remote`):
```json
"remote": {
"url": "https://auth.example.com",
"token": "<sync token>",
"cache_ttl": 60,
"refresh_interval": 300
}
```
- `url` — the remote instance's base URL. The satellite connects to
`{url}/auth/api/sync/ws`; the connection is server-to-server and not
host-dispatched, so internal addresses work.
- `token` — bearer token for the sync channel. The **remote accepts tokens
via its `PASKIA_SYNC_TOKENS` environment variable** (comma-separated);
nothing is stored in the remote's database, and with the variable unset
the sync endpoint stays closed. The token is write-only over the admin
API (an empty field keeps the stored one).
- `cache_ttl` — seconds the replica remains trusted after the sync channel
goes down; then checks fail closed (503). Set it large (up to the 24 h
session lifetime) for fail-open behavior.
- `refresh_interval` — seconds between reconnects; every connect starts
from a full snapshot, which reconciles any drift.
A remote domain **must mark an auth host** (validated cross-domain and in
the UI): the profile, admin and sign-in pages live there, so browsers and
WebSockets go directly to the remote. Other domains on the same satellite
remain fully local — the multi-domain config mixes both kinds freely.
## How it works
**Dispatch is keyed by host, and only this module knows about stores.**
`satellite.store_for_host(host)` returns the local DB or the replica of
the remote backing the host's domain (raising 503 `HTTPException` when
the replica is unavailable). The session read path (`session_ctx`,
`authz.verify`, `build_user_info`, `/check`) just passes the host it
already has; writes dispatch likewise (`satellite.refresh_session`
write-behind for remote, `db.update_session` for local;
`satellite.evict_session` on logout). `satellite.forward_request(request)`
returns the proxied response for remote domains or `None` for local ones.
**The replica** is a plain `DB` struct instance in RAM, never persisted.
On connect the remote sends a snapshot of the replicated tables
(permissions, orgs, roles, users, credentials, sessions), then live
upsert/delete events emitted from the struct `store()`/`delete()` hooks
(which also cover cascade deletes) and field-mutating operations. A
single ordered WebSocket cannot gap; a slow subscriber is dropped and
resyncs. The feed carries no usable secrets: sessions are keyed by
`hash_secret` output, credentials carry public keys only, and the OIDC
signing key is never replicated.
## Endpoint behavior for remote domains
| Endpoint | Handling |
|---|---|
| `GET /auth/api/forward`, `GET /check`, `GET /user-info`, `GET /settings` | served from the replica (sub-ms) |
| `POST /auth/api/validate` | verified from the replica; the throttled refresh updates the replica and is written back over the sync channel; cookie renewed locally |
| `POST /auth/api/logout` | proxied to the remote (original Host preserved) and evicted from the replica immediately |
| `POST /auth/api/set-session`, `GET /token-info` | proxied (the exchange code/reset token lives on the remote); the session arrives via sync event |
| `/auth/oidc/*` | proxied (signing key and OIDC sessions stay on the remote) |
| `/auth/ws/*`, `/auth/remote-auth/*`, admin, profile | not served — the auth host requirement means these are reached on the remote directly |
Freshness hierarchy:
1. Changes made **through** the satellite: immediate (write-behind,
optimistic eviction).
2. Changes made **directly on the remote**: a sync event, ~1 network RTT.
3. Channel down: the replica stays authoritative until `cache_ttl` past
the disconnect (dead-peer detection is bounded by the ~10 s keepalive),
then 503. Every reconnect starts from a fresh snapshot.
## The remote side
Strictly additive and RAM-only: `syncfeed` (a subscriber set fed by the
commit hooks) and the token-gated `/auth/api/sync/ws` endpoint serving
snapshot + live events and accepting `session_refresh` write-backs. With
no satellites connected, the hooks are a no-op.
## Trust and caveats
- The satellite host holds a full copy of the remote's auth data in RAM
(minus the OIDC key) — treat it as trusted as the remote.
- Avatars are stored on the remote's disk; `user-info` from a replica
reports no avatar URL.
- OIDC sessions in a replica-backed `user-info` show the client UUID
rather than its name (OIDC clients are not replicated).
---
# Design review (the rejected alternatives)
## Option A — caching HTTP reverse proxy
A thin proxy caching `/auth/api/forward`, `/check`, `/user-info`,
`/settings` responses keyed by `(Host, cookie, query)` with
`TTL = min(configured TTL, Remote-Session-Expires now)`; everything
else forwarded verbatim, WebSockets tunneled, `/logout` intercepted for
eviction. **Option B** adds a remote change feed so eviction happens
within one RTT instead of at TTL.
- Remote changes: none for A; one additive endpoint for B.
- The proxy needs no credentials — requests are authenticated by the end
user's cookie, forwarded on a miss.
## What the read-only local state buys over the HTTP cache
- **Full `SessionContext` locally.** A replays the byte-response it once
saw; the satellite *computes* the answer. Query combinations never seen
before (new `perm`/`max_age`/`public` shapes) are served locally but
miss A's cache. The replica holds the *domain model*, so derived
answers (effective permissions per host, `max_age` against
`credential.last_used`, `Remote-*` composition) are correct without
having been witnessed.
- **One invalidation model.** A hand-builds invalidation rules per
endpoint (query-key mapping, cookie re-keying on renew, 401 variants).
Events mutate the replica (upsert/delete by table+key) and every
endpoint becomes consistent at once — including future ones.
- **Degradation behaves like a real instance.** With the remote down, the
satellite serves a coherent auth service from the replica (expiry
enforced locally, bounded by `cache_ttl`); A serves unrelated cached
responses with gaps wherever the cache was cold.
- **Multi-domain uniformity.** Remote backing is a property of a domain
in the existing registry; local and remote rp-ids coexist in one
instance. A is a separate component bolted in front of specific URLs.
- **User simplicity.** Configured once in domain config; A needs
deployment and cache-key discipline per frontend application.
## What it costs
- The read path must be honest about which DB it reads: `DB.session_ctx`
and `/check` were rewritten to use their own tables instead of struct
convenience properties that reach the global database. (A first draft's
contextvar-dependent `db.data()` was rejected: a global accessor whose
meaning shifts under the caller. Dispatch is instead keyed explicitly
by the request host.)
- A sync protocol (snapshot + live events + reconnect reconciliation).
- A trusted satellite host (full data copy in RAM).
- Additive remote code (sync endpoint + commit hooks), where A needs
none.
- Replica housekeeping (expiry sweeper, write-behind, optimistic
eviction).
## Summary
A(+B) is the right tool to "make forward-auth fast in front of an
untouched server". The satellite — implemented here — is the right tool
when it should *be* a paskia instance for its remote domains: one
consistency model, correct answers for un-cached query shapes, graceful
degradation, and per-domain mixing with local rp-ids, at the price of the
read-path cleanup, the sync protocol, and a trusted satellite host.
+15 -2
View File
@@ -480,6 +480,7 @@ function createDomain() {
origins: [], origins: [],
originValidation: [], originValidation: [],
wellKnownCheck: null, wellKnownCheck: null,
remote: null,
}) })
} }
@@ -495,6 +496,8 @@ function openDomain(domain) {
origins: rows.map(r => r.key), origins: rows.map(r => r.key),
originValidation: rows.map(() => null), originValidation: rows.map(() => null),
wellKnownCheck: null, wellKnownCheck: null,
// The sync token is write-only: an empty field keeps the stored one
remote: domain.remote ? { ...domain.remote, token: '' } : null,
}) })
} }
@@ -923,9 +926,19 @@ async function submitDialog() {
} }
closeDialog() closeDialog()
// remote is replaced wholesale when present; null clears it, an
// absent key (create without remote) leaves it unset.
const remote = d.remote?.url?.trim()
? {
url: d.remote.url.trim().replace(/\/+$/, ''),
token: d.remote.token || '',
cache_ttl: Number(d.remote.cache_ttl) || 60,
refresh_interval: Number(d.remote.refresh_interval) || 300,
}
: null
const req = d.isNew const req = d.isNew
? apiJson('/auth/api/admin/domains/', { method: 'POST', body: { rp_id, rp_name, origins } }) ? apiJson('/auth/api/admin/domains/', { method: 'POST', body: { rp_id, rp_name, origins, ...(remote ? { remote } : {}) } })
: apiJson(`/auth/api/admin/domains/${rp_id}`, { method: 'PATCH', body: { rp_name, origins } }) : apiJson(`/auth/api/admin/domains/${rp_id}`, { method: 'PATCH', body: { rp_name, origins, remote } })
req req
.then(() => { .then(() => {
authStore.showMessage(`Domain "${rp_id}" ${d.isNew ? 'created' : 'updated'}.`, 'success', 2500) authStore.showMessage(`Domain "${rp_id}" ${d.isNew ? 'created' : 'updated'}.`, 'success', 2500)
+1
View File
@@ -464,6 +464,7 @@ defineExpose({ focusFirstElement })
</div> </div>
<div class="perm-id-info"> <div class="perm-id-info">
<span class="id-text">{{ domain.rp_id }}</span> <span class="id-text">{{ domain.rp_id }}</span>
<span v-if="domain.remote" class="id-text" :title="`Served from remote ${domain.remote.url} (satellite mode)`">🛰 {{ domain.remote.url }}</span>
</div> </div>
</td> </td>
<td class="domain-origins"><span v-for="(e, i) in originDisplayEntries(domain)" :key="e.key">{{ i ? ', ' : '' }}{{ e.key }}{{ e.auth ? '🔑' : '' }}{{ e.related ? '🔗' : '' }}</span></td> <td class="domain-origins"><span v-for="(e, i) in originDisplayEntries(domain)" :key="e.key">{{ i ? ', ' : '' }}{{ e.key }}{{ e.auth ? '🔑' : '' }}{{ e.related ? '🔗' : '' }}</span></td>
@@ -18,6 +18,33 @@ const title = computed(() =>
// compares against it, and hosts are case-insensitive) // compares against it, and hosts are case-insensitive)
const dialogRpId = computed(() => (props.dialog.data?.rp_id || '').trim().toLowerCase()) const dialogRpId = computed(() => (props.dialog.data?.rp_id || '').trim().toLowerCase())
// --- Remote (satellite) backing ---
//
// A remote domain is served from another paskia instance: this one keeps a
// RAM-only read replica for fast local session checks and forwards
// mutations. The remote must accept our sync token via its
// PASKIA_SYNC_TOKENS environment variable. An auth host (the remote's) is
// required — profile, admin and sign-in pages live there.
const remoteEnabled = computed({
get: () => !!props.dialog.data?.remote,
set: on => {
const d = props.dialog.data
if (!d) return
d.remote = on ? { url: '', token: '', cache_ttl: 60, refresh_interval: 300 } : null
},
})
const remoteUrlInvalid = computed(() => {
const url = props.dialog.data?.remote?.url?.trim()
if (!url) return false
return !/^https?:\/\/[^\s/]+/.test(url)
})
// Remote domains must mark an auth host (the server rejects the save)
const remoteMissingAuthHost = computed(
() => !!props.dialog.data?.remote && !props.dialog.data?.auth_host
)
// Block submit on hard errors: malformed entries, an over-cap related // Block submit on hard errors: malformed entries, an over-cap related
// list (the server rejects the save), a save that would lock the admin // list (the server rejects the save), a save that would lock the admin
// out of the domain they are using, or validation still in flight. // out of the domain they are using, or validation still in flight.
@@ -31,6 +58,8 @@ const isValidationInvalid = computed(() => {
if (relatedEntries.value.length > 5) return true if (relatedEntries.value.length > 5) return true
if (d.isNew && !isWellFormedDomain(d.rp_id || '')) return true if (d.isNew && !isWellFormedDomain(d.rp_id || '')) return true
if (lockoutWarning.value) return true if (lockoutWarning.value) return true
if (remoteUrlInvalid.value || remoteMissingAuthHost.value) return true
if (d.remote && !d.remote.url?.trim()) return true
return false return false
}) })
@@ -515,6 +544,33 @@ function onRemoveOrigin(i) {
<p class="small muted"> <p class="small muted">
Only the listed sites may sign in with {{ dialog.data.rp_id }} passkeys. Wildcards may be used: <strong>**.{{ dialog.data.rp_id }}</strong> allows the whole domain, <strong>*.{{ dialog.data.rp_id }}</strong> only a single subdomain level.<template v-if="relatedEntries.length"> 🔗 means related host requiring WebAuthn ROR setup.</template><template v-if="dialog.data.auth_host"> 🔑 is the dedicated Paskia host for all account management.</template> Only the listed sites may sign in with {{ dialog.data.rp_id }} passkeys. Wildcards may be used: <strong>**.{{ dialog.data.rp_id }}</strong> allows the whole domain, <strong>*.{{ dialog.data.rp_id }}</strong> only a single subdomain level.<template v-if="relatedEntries.length"> 🔗 means related host requiring WebAuthn ROR setup.</template><template v-if="dialog.data.auth_host"> 🔑 is the dedicated Paskia host for all account management.</template>
</p> </p>
<div class="origin-label">
<label class="remote-toggle">
<input type="checkbox" v-model="remoteEnabled" />
Remote instance (satellite mode)
</label>
</div>
<template v-if="dialog.data.remote">
<label>Remote URL
<input v-model="dialog.data.remote.url" placeholder="https://auth.example.com" data-form-type="other" :class="{ 'input-error': remoteUrlInvalid }" />
</label>
<p v-if="remoteUrlInvalid" class="small error">Must be an http(s) URL.</p>
<p v-if="remoteMissingAuthHost" class="small error">A remote domain must mark an auth host above profile, admin and sign-in pages live there (typically the remote's own site).</p>
<label>Sync token
<input v-model="dialog.data.remote.token" type="password" placeholder="Token in the remote's PASKIA_SYNC_TOKENS" autocomplete="off" data-form-type="other" />
</label>
<p class="small muted">Accepted by the remote via its PASKIA_SYNC_TOKENS environment variable.<template v-if="!dialog.data.isNew"> Leave empty to keep the stored token.</template></p>
<label>Staleness limit (cache TTL, seconds)
<input v-model.number="dialog.data.remote.cache_ttl" type="number" min="1" />
</label>
<label>Full re-sync interval (seconds)
<input v-model.number="dialog.data.remote.refresh_interval" type="number" min="30" />
</label>
<p class="small muted">
Session checks run locally against a RAM replica of the remote (sub-millisecond). If the connection is down longer than the staleness limit, checks fail closed (503).
</p>
</template>
</AdminDialog> </AdminDialog>
</template> </template>
@@ -540,4 +596,7 @@ function onRemoveOrigin(i) {
border-color: var(--color-error); border-color: var(--color-error);
background: var(--color-error-bg, rgba(239, 68, 68, 0.05)); background: var(--color-error-bg, rgba(239, 68, 68, 0.05));
} }
.remote-toggle { display: flex; align-items: center; gap: var(--space-xs); font-weight: 600; font-size: 0.95rem; }
.remote-toggle input { width: auto; }
</style> </style>
+16 -24
View File
@@ -190,18 +190,6 @@ def cmd_migrate(args: argparse.Namespace) -> None:
print(f"{action} {db_file_path()} (domains: {', '.join(rp_ids)})") print(f"{action} {db_file_path()} (domains: {', '.join(rp_ids)})")
def _save_listen(db_path: Path, listen: list[str] | None) -> None:
"""Persist the listen endpoints to the stored configuration."""
kanta = Kanta(str(db_path), DB())
async def _write() -> None:
async with kanta:
with kanta.transaction("serve:save_listen"):
kanta.data.config.listen = listen
asyncio.run(_write())
def cmd_serve(args: argparse.Namespace) -> None: def cmd_serve(args: argparse.Namespace) -> None:
"""Open the combined database and serve all configured domains.""" """Open the combined database and serve all configured domains."""
db_path = db_file_path() db_path = db_file_path()
@@ -214,14 +202,20 @@ def cmd_serve(args: argparse.Namespace) -> None:
) )
raise SystemExit(f"Database {db_path} not found — run 'paskia init' first.") raise SystemExit(f"Database {db_path} not found — run 'paskia init' first.")
if args.save and args.listen is not None:
# '--listen ""' clears the stored endpoints (back to the default)
_save_listen(db_path, _split_multi(args.listen) or None)
config = _load_stored_config(db_path) config = _load_stored_config(db_path)
listen = _split_multi(args.listen) or config.listen # Effective serve parameters, teleported to the server process(es); the
configure_domains(listen=listen) # app persists the listen endpoints to the database when save is set.
cfg = serve_config()
cfg.save = bool(args.save and args.listen is not None)
if cfg.save:
# '--listen ""' clears the stored endpoints (back to the default)
cfg.listen = _split_multi(args.listen) or None
else:
cfg.listen = _split_multi(args.listen) or config.listen
teleport() # Serialize bound config before spawning workers
configure_domains(listen=cfg.listen)
try: try:
registry = build_registry(config) registry = build_registry(config)
except ValueError as e: except ValueError as e:
@@ -229,18 +223,16 @@ def cmd_serve(args: argparse.Namespace) -> None:
# Sanitization warnings (serving is best-effort; fixing the stored config # Sanitization warnings (serving is best-effort; fixing the stored config
# is the admin's job via the admin interface) are logged by build(). # is the admin's job via the admin interface) are logged by build().
# Pass process-global serve parameters to the server process(es) startupbox.print_startup_config(
serve_config().listen = listen registry, listen=cfg.listen, default_port=DEFAULT_PORT
teleport() # Serialize bound config before spawning workers )
startupbox.print_startup_config(registry, listen=listen, default_port=DEFAULT_PORT)
# Run the server (spawns processes in dev mode) # Run the server (spawns processes in dev mode)
# tracerite, access logging and log config are handled by fastapi_vue.server; # tracerite, access logging and log config are handled by fastapi_vue.server;
# we print our own startup config box, so disable the built-in one. # we print our own startup config box, so disable the built-in one.
server.run( server.run(
"paskia.fastapi.mainapp:app", "paskia.fastapi.mainapp:app",
listen=listen, listen=cfg.listen,
default_port=DEFAULT_PORT, default_port=DEFAULT_PORT,
server_header=False, server_header=False,
startup_box=None, startup_box=None,
+9 -2
View File
@@ -24,8 +24,15 @@ EXPIRES = SESSION_LIFETIME
def session_ctx(auth: str, host: str | None = None): def session_ctx(auth: str, host: str | None = None):
"""Get session context with normalized host.""" """Get session context with normalized host.
return db.data().session_ctx(auth, hostutil.normalize_host(host))
The store is dispatched by host: remote domains read their replica.
"""
from paskia import satellite # noqa: PLC0415 (import cycle)
return satellite.store_for_host(host).session_ctx(
auth, hostutil.normalize_host(host)
)
def expires() -> datetime: def expires() -> datetime:
+4 -1
View File
@@ -67,7 +67,10 @@ async def check_admin_credentials() -> bool:
# Check first admin user for credentials on any configured domain # Check first admin user for credentials on any configured domain
admin_user = admin_users[0] admin_user = admin_users[0]
reg = domains.registry() reg = domains.registry()
configured = sorted(d.rp_id for d in reg.domains) # Remote domains hold their credentials on the remote instance
configured = sorted(d.rp_id for d in reg.domains if d.remote is None)
if not configured:
return False
if not any(admin_user.credential_ids_for(rp_id) for rp_id in configured): if not any(admin_user.credential_ids_for(rp_id) for rp_id in configured):
# Admin exists but has no credential on any domain # Admin exists but has no credential on any domain
+2
View File
@@ -67,6 +67,7 @@ from paskia.db.structs import (
DomainConfig, DomainConfig,
Org, Org,
Permission, Permission,
RemoteConfig,
ResetToken, ResetToken,
Role, Role,
Session, Session,
@@ -90,6 +91,7 @@ __all__ = [
"Org", "Org",
"Permission", "Permission",
"DomainConfig", "DomainConfig",
"RemoteConfig",
"ResetToken", "ResetToken",
"Role", "Role",
"Session", "Session",
+27 -2
View File
@@ -13,7 +13,7 @@ from uuid import UUID
import uuid7 import uuid7
from paskia import oidc_notify from paskia import oidc_notify, syncfeed
from paskia.config import SESSION_LIFETIME from paskia.config import SESSION_LIFETIME
from paskia.db.structs import ( from paskia.db.structs import (
DB, DB,
@@ -23,6 +23,7 @@ from paskia.db.structs import (
Org, Org,
OriginEntry, OriginEntry,
Permission, Permission,
RemoteConfig,
ResetToken, ResetToken,
Role, Role,
Session, Session,
@@ -103,6 +104,7 @@ def update_permission(
_db.permissions[uuid].scope = scope _db.permissions[uuid].scope = scope
_db.permissions[uuid].display_name = display_name _db.permissions[uuid].display_name = display_name
_db.permissions[uuid].domain = domain _db.permissions[uuid].domain = domain
syncfeed.emit("permissions", str(uuid), _db.permissions[uuid])
def delete_permission(uuid: UUID, *, ctx: SessionContext | None = None) -> None: def delete_permission(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
@@ -155,6 +157,7 @@ def update_org_name(
raise ValueError(f"Organization {uuid} not found") raise ValueError(f"Organization {uuid} not found")
with _transaction("admin:update_org_name", ctx): with _transaction("admin:update_org_name", ctx):
_db.orgs[uuid].display_name = display_name _db.orgs[uuid].display_name = display_name
syncfeed.emit("orgs", str(uuid), _db.orgs[uuid])
def delete_org(uuid: UUID, *, ctx: SessionContext | None = None) -> None: def delete_org(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
@@ -180,6 +183,9 @@ def add_permission_to_org(
with _transaction("admin:add_permission_to_org", ctx): with _transaction("admin:add_permission_to_org", ctx):
_db.permissions[permission_uuid].orgs[org_uuid] = True _db.permissions[permission_uuid].orgs[org_uuid] = True
syncfeed.emit(
"permissions", str(permission_uuid), _db.permissions[permission_uuid]
)
def remove_permission_from_org( def remove_permission_from_org(
@@ -197,6 +203,9 @@ def remove_permission_from_org(
with _transaction("admin:remove_permission_from_org", ctx): with _transaction("admin:remove_permission_from_org", ctx):
_db.permissions[permission_uuid].orgs.pop(org_uuid, None) _db.permissions[permission_uuid].orgs.pop(org_uuid, None)
syncfeed.emit(
"permissions", str(permission_uuid), _db.permissions[permission_uuid]
)
def create_role(role: Role, *, ctx: SessionContext | None = None) -> None: def create_role(role: Role, *, ctx: SessionContext | None = None) -> None:
@@ -220,6 +229,7 @@ def update_role_name(
raise ValueError(f"Role {uuid} not found") raise ValueError(f"Role {uuid} not found")
with _transaction("admin:update_role_name", ctx): with _transaction("admin:update_role_name", ctx):
_db.roles[uuid].display_name = display_name _db.roles[uuid].display_name = display_name
syncfeed.emit("roles", str(uuid), _db.roles[uuid])
def add_permission_to_role( def add_permission_to_role(
@@ -235,6 +245,7 @@ def add_permission_to_role(
raise ValueError(f"Permission {permission_uuid} not found") raise ValueError(f"Permission {permission_uuid} not found")
with _transaction("admin:add_permission_to_role", ctx): with _transaction("admin:add_permission_to_role", ctx):
_db.roles[role_uuid].permissions[permission_uuid] = True _db.roles[role_uuid].permissions[permission_uuid] = True
syncfeed.emit("roles", str(role_uuid), _db.roles[role_uuid])
def remove_permission_from_role( def remove_permission_from_role(
@@ -248,6 +259,7 @@ def remove_permission_from_role(
raise ValueError(f"Role {role_uuid} not found") raise ValueError(f"Role {role_uuid} not found")
with _transaction("admin:remove_permission_from_role", ctx): with _transaction("admin:remove_permission_from_role", ctx):
_db.roles[role_uuid].permissions.pop(permission_uuid, None) _db.roles[role_uuid].permissions.pop(permission_uuid, None)
syncfeed.emit("roles", str(role_uuid), _db.roles[role_uuid])
def delete_role(uuid: UUID, *, ctx: SessionContext | None = None) -> None: def delete_role(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
@@ -302,6 +314,7 @@ def update_user_display_name(
slug = slugify_name(display_name) slug = slugify_name(display_name)
if slug and not is_username_taken(slug, exclude_uuid=uuid): if slug and not is_username_taken(slug, exclude_uuid=uuid):
user.preferred_username = slug user.preferred_username = slug
syncfeed.emit("users", str(uuid), user)
def update_user_info( def update_user_info(
@@ -380,6 +393,7 @@ def update_user_info(
user.preferred_username = preferred_username user.preferred_username = preferred_username
if telephone is not _UNSET: if telephone is not _UNSET:
user.telephone = telephone user.telephone = telephone
syncfeed.emit("users", str(uuid), user)
def update_user_role( def update_user_role(
@@ -395,6 +409,7 @@ def update_user_role(
raise ValueError(f"Role {role_uuid} not found") raise ValueError(f"Role {role_uuid} not found")
with _transaction("admin:update_user_role", ctx): with _transaction("admin:update_user_role", ctx):
_db.users[uuid].role_uuid = role_uuid _db.users[uuid].role_uuid = role_uuid
syncfeed.emit("users", str(uuid), _db.users[uuid])
def delete_user(uuid: UUID, *, ctx: SessionContext | None = None) -> None: def delete_user(uuid: UUID, *, ctx: SessionContext | None = None) -> None:
@@ -429,6 +444,7 @@ def update_credential_sign_count(
_db.credentials[uuid].sign_count = sign_count _db.credentials[uuid].sign_count = sign_count
if last_used: if last_used:
_db.credentials[uuid].last_used = last_used _db.credentials[uuid].last_used = last_used
syncfeed.emit("credentials", str(uuid), _db.credentials[uuid])
def delete_credential( def delete_credential(
@@ -476,6 +492,7 @@ def update_session(
s.validated = validated s.validated = validated
if issuer is not None: if issuer is not None:
s.issuer = issuer s.issuer = issuer
syncfeed.emit("sessions", key, s)
def delete_session( def delete_session(
@@ -598,6 +615,9 @@ def login(
# Update credential # Update credential
_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
syncfeed.emit(
"credentials", str(credential_uuid), _db.credentials[credential_uuid]
)
return token return token
@@ -625,6 +645,9 @@ def oidc_login(
# Update credential # Update credential
_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
syncfeed.emit(
"credentials", str(credential_uuid), _db.credentials[credential_uuid]
)
def create_credential_session( def create_credential_session(
@@ -714,9 +737,10 @@ def update_domain(
*, *,
rp_name: str | None, rp_name: str | None,
origins: dict[str, bool | OriginEntry], origins: dict[str, bool | OriginEntry],
remote: RemoteConfig | None = None,
ctx: SessionContext | None = None, ctx: SessionContext | None = None,
) -> None: ) -> None:
"""Replace a domain's rp_name and origins table (wholesale). """Replace a domain's rp_name, origins table and remote (wholesale).
The rp-id itself is immutable: credentials are stamped with it, so The rp-id itself is immutable: credentials are stamped with it, so
changing it would orphan them — delete and recreate the domain instead. changing it would orphan them — delete and recreate the domain instead.
@@ -728,6 +752,7 @@ def update_domain(
with _transaction("admin:update_domain", ctx): with _transaction("admin:update_domain", ctx):
domain.rp_name = rp_name domain.rp_name = rp_name
domain.origins = origins domain.origins = origins
domain.remote = remote
def delete_domain(rp_id: str, *, ctx: SessionContext | None = None) -> None: def delete_domain(rp_id: str, *, ctx: SessionContext | None = None) -> None:
+48 -8
View File
@@ -9,7 +9,7 @@ from uuid import UUID
import msgspec import msgspec
import uuid7 import uuid7
from paskia import db from paskia import db, syncfeed
from paskia.util import passphrase as passphrase_util from paskia.util import passphrase as passphrase_util
from paskia.util.crypto import hash_secret from paskia.util.crypto import hash_secret
@@ -51,6 +51,7 @@ class Permission(msgspec.Struct, dict=True, omit_defaults=True):
def store(self) -> None: def store(self) -> None:
"""Store this permission in the database. Must be called inside a transaction.""" """Store this permission in the database. Must be called inside a transaction."""
db.data().permissions[self.uuid] = self db.data().permissions[self.uuid] = self
syncfeed.emit("permissions", str(self.uuid), self)
def delete(self) -> None: def delete(self) -> None:
"""Delete this permission and remove it from all roles. """Delete this permission and remove it from all roles.
@@ -59,8 +60,10 @@ class Permission(msgspec.Struct, dict=True, omit_defaults=True):
""" """
_data = db.data() _data = db.data()
for role in _data.roles.values(): for role in _data.roles.values():
role.permissions.pop(self.uuid, None) if role.permissions.pop(self.uuid, None) is not None:
syncfeed.emit("roles", str(role.uuid), role)
del _data.permissions[self.uuid] del _data.permissions[self.uuid]
syncfeed.emit("permissions", str(self.uuid), None)
@classmethod @classmethod
def create( def create(
@@ -103,6 +106,7 @@ class Org(msgspec.Struct, dict=True):
def store(self) -> None: def store(self) -> None:
"""Store this organization in the database. Must be called inside a transaction.""" """Store this organization in the database. Must be called inside a transaction."""
db.data().orgs[self.uuid] = self db.data().orgs[self.uuid] = self
syncfeed.emit("orgs", str(self.uuid), self)
def delete(self) -> None: def delete(self) -> None:
"""Delete this org and cascade to roles, users. Remove from permissions. """Delete this org and cascade to roles, users. Remove from permissions.
@@ -111,12 +115,16 @@ class Org(msgspec.Struct, dict=True):
""" """
_data = db.data() _data = db.data()
for p in _data.permissions.values(): for p in _data.permissions.values():
p.orgs.pop(self.uuid, None) if p.orgs.pop(self.uuid, None) is not None:
syncfeed.emit("permissions", str(p.uuid), p)
for role in self.roles: for role in self.roles:
for user in role.users: for user in role.users:
del _data.users[user.uuid] del _data.users[user.uuid]
syncfeed.emit("users", str(user.uuid), None)
del _data.roles[role.uuid] del _data.roles[role.uuid]
syncfeed.emit("roles", str(role.uuid), None)
del _data.orgs[self.uuid] del _data.orgs[self.uuid]
syncfeed.emit("orgs", str(self.uuid), None)
@classmethod @classmethod
def create(cls, display_name: str, created_at: datetime | None = None) -> Org: def create(cls, display_name: str, created_at: datetime | None = None) -> Org:
@@ -170,10 +178,12 @@ class Role(msgspec.Struct, dict=True, omit_defaults=True):
def store(self) -> None: def store(self) -> None:
"""Store this role in the database. Must be called inside a transaction.""" """Store this role in the database. Must be called inside a transaction."""
db.data().roles[self.uuid] = self db.data().roles[self.uuid] = self
syncfeed.emit("roles", str(self.uuid), self)
def delete(self) -> None: def delete(self) -> None:
"""Delete this role from the database. Must be called inside a transaction.""" """Delete this role from the database. Must be called inside a transaction."""
del db.data().roles[self.uuid] del db.data().roles[self.uuid]
syncfeed.emit("roles", str(self.uuid), None)
@classmethod @classmethod
def create( def create(
@@ -254,6 +264,7 @@ class User(msgspec.Struct, dict=True, omit_defaults=True, kw_only=True):
def store(self) -> None: def store(self) -> None:
"""Store this user in the database. Must be called inside a transaction.""" """Store this user in the database. Must be called inside a transaction."""
db.data().users[self.uuid] = self db.data().users[self.uuid] = self
syncfeed.emit("users", str(self.uuid), self)
def delete(self) -> None: def delete(self) -> None:
"""Delete this user and cascade to credentials, sessions, reset tokens. """Delete this user and cascade to credentials, sessions, reset tokens.
@@ -263,11 +274,14 @@ class User(msgspec.Struct, dict=True, omit_defaults=True, kw_only=True):
_data = db.data() _data = db.data()
for cred in self.credentials: for cred in self.credentials:
del _data.credentials[cred.uuid] del _data.credentials[cred.uuid]
syncfeed.emit("credentials", str(cred.uuid), None)
for sess in self.sessions: for sess in self.sessions:
del _data.sessions[sess.key] del _data.sessions[sess.key]
syncfeed.emit("sessions", sess.key, None)
for token in self.reset_tokens: for token in self.reset_tokens:
del _data.reset_tokens[token.key] del _data.reset_tokens[token.key]
del _data.users[self.uuid] del _data.users[self.uuid]
syncfeed.emit("users", str(self.uuid), None)
@classmethod @classmethod
def create( def create(
@@ -331,6 +345,7 @@ class Credential(msgspec.Struct, dict=True):
def store(self) -> None: def store(self) -> None:
"""Store this credential in the database. Must be called inside a transaction.""" """Store this credential in the database. Must be called inside a transaction."""
db.data().credentials[self.uuid] = self db.data().credentials[self.uuid] = self
syncfeed.emit("credentials", str(self.uuid), self)
def delete(self) -> None: def delete(self) -> None:
"""Delete this credential and all its sessions. """Delete this credential and all its sessions.
@@ -340,7 +355,9 @@ class Credential(msgspec.Struct, dict=True):
_data = db.data() _data = db.data()
for sess in self.sessions: for sess in self.sessions:
del _data.sessions[sess.key] del _data.sessions[sess.key]
syncfeed.emit("sessions", sess.key, None)
del _data.credentials[self.uuid] del _data.credentials[self.uuid]
syncfeed.emit("credentials", str(self.uuid), None)
@classmethod @classmethod
def create( def create(
@@ -418,10 +435,13 @@ class Session(msgspec.Struct, dict=True, omit_defaults=True):
_data.sessions[self.key] = self _data.sessions[self.key] = self
_data.users[self.user_uuid].last_seen = last_seen _data.users[self.user_uuid].last_seen = last_seen
_data.users[self.user_uuid].visits += 1 _data.users[self.user_uuid].visits += 1
syncfeed.emit("sessions", self.key, self)
syncfeed.emit("users", str(self.user_uuid), _data.users[self.user_uuid])
def delete(self) -> None: def delete(self) -> None:
"""Delete this session from the database. Must be called inside a transaction.""" """Delete this session from the database. Must be called inside a transaction."""
del db.data().sessions[self.key] del db.data().sessions[self.key]
syncfeed.emit("sessions", self.key, None)
@classmethod @classmethod
def create( def create(
@@ -622,6 +642,21 @@ class OriginEntry(msgspec.Struct, omit_defaults=True):
auth_host: bool = False # This site hosts the account/admin interface auth_host: bool = False # This site hosts the account/admin interface
class RemoteConfig(msgspec.Struct, omit_defaults=True):
"""Upstream paskia instance backing a remote (satellite-served) domain.
The satellite keeps a RAM-only read replica of the remote's tables and
answers session-dependent reads locally; mutations are forwarded. The
token authenticates the sync channel (the remote reads accepted tokens
from its PASKIA_SYNC_TOKENS environment, never from its database).
"""
url: str # e.g. "https://auth.example.com"
token: str = ""
cache_ttl: int = 60 # staleness bound (seconds) while the sync channel is down
refresh_interval: int = 300 # full re-sync cadence (seconds)
class DomainConfig(msgspec.Struct, omit_defaults=True): class DomainConfig(msgspec.Struct, omit_defaults=True):
"""Configuration for one domain (one WebAuthn rp-id). """Configuration for one domain (one WebAuthn rp-id).
@@ -641,6 +676,7 @@ class DomainConfig(msgspec.Struct, omit_defaults=True):
rp_name: str | None = None rp_name: str | None = None
origins: dict[str, bool | OriginEntry] = {} origins: dict[str, bool | OriginEntry] = {}
remote: RemoteConfig | None = None
class Config(msgspec.Struct, omit_defaults=True): class Config(msgspec.Struct, omit_defaults=True):
@@ -728,17 +764,21 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
if s.host != host: if s.host != host:
return None return None
# Look up via this instance's own tables: a DB must be
# self-contained so that read replicas work unchanged.
try: try:
user = s.user user = self.users[s.user_uuid]
role = user.role role = self.roles[user.role_uuid]
org = role.org org = self.orgs[role.org_uuid]
credential = s.credential credential = self.credentials[s.credential_uuid]
except KeyError: except KeyError:
return None return None
# Effective permissions: role's permissions that the org can grant, # Effective permissions: role's permissions that the org can grant,
# filtered by domain restriction # filtered by domain restriction
org_perm_uuids = {p.uuid for p in org.permissions} org_perm_uuids = {
p.uuid for p in self.permissions.values() if org.uuid in p.orgs
}
effective_perms = [] effective_perms = []
for perm_uuid in role.permission_set: for perm_uuid in role.permission_set:
+50 -2
View File
@@ -14,13 +14,14 @@ domain are in-domain, entries outside it are related.
from __future__ import annotations from __future__ import annotations
import asyncio
import contextvars import contextvars
import logging import logging
import os import os
from fastapi_vue.hostutil import parse_endpoints from fastapi_vue.hostutil import parse_endpoints
from paskia.db.structs import Config, DomainConfig, OriginEntry from paskia.db.structs import Config, DomainConfig, OriginEntry, RemoteConfig
from paskia.sansio import Passkey from paskia.sansio import Passkey
from paskia.util import hostutil from paskia.util import hostutil
from paskia.util.constants import DEFAULT_PORT from paskia.util.constants import DEFAULT_PORT
@@ -101,6 +102,11 @@ class Domain:
def rp_name(self) -> str: def rp_name(self) -> str:
return self.passkey.rp_name return self.passkey.rp_name
@property
def remote(self) -> RemoteConfig | None:
"""Upstream config when this domain is served as a satellite."""
return self.config.remote
@property @property
def own_auth_host(self) -> str | None: def own_auth_host(self) -> str | None:
"""This domain's own auth host as host[:port], if configured.""" """This domain's own auth host as host[:port], if configured."""
@@ -207,6 +213,10 @@ def validate_config(
for rp_id, domain in config.domains.items(): for rp_id, domain in config.domains.items():
hostutil.validate_rp_id(rp_id) hostutil.validate_rp_id(rp_id)
if domain.remote is not None and not domain.remote.url.startswith(
("https://", "http://")
):
raise ValueError(f"Domain '{rp_id}': remote URL must be an http(s) URL")
domain_auth_host: str | None = None domain_auth_host: str | None = None
related_count = 0 related_count = 0
@@ -273,6 +283,11 @@ def validate_config(
f"Domain '{rp_id}' has {related_count} related origins " f"Domain '{rp_id}' has {related_count} related origins "
f"(maximum {related_origin_cap})" f"(maximum {related_origin_cap})"
) )
if domain.remote is not None and domain_auth_host is None:
raise ValueError(
f"Domain '{rp_id}' is remote — it must mark an auth host "
"(profile, admin and sign-in pages live there)"
)
rp_ids = set(config.domains) rp_ids = set(config.domains)
for hn, owner in auth_hosts.items(): for hn, owner in auth_hosts.items():
@@ -363,6 +378,20 @@ def sanitize_config(
auth_seen = True auth_seen = True
origins[key] = props origins[key] = props
if domain.remote is not None:
if not domain.remote.url.startswith(("https://", "http://")):
warn(f"Domain '{rp_id}': invalid remote URL — remote dropped")
remote = None
else:
remote = domain.remote
if not auth_seen:
warn(
f"Domain '{rp_id}': remote domain without an auth host — "
"profile, admin and sign-in pages have nowhere to live"
)
else:
remote = None
related = sorted(k for k in origins if is_related_key(rp_id, k)) related = sorted(k for k in origins if is_related_key(rp_id, k))
if len(related) > related_origin_cap: if len(related) > related_origin_cap:
warn( warn(
@@ -372,7 +401,9 @@ def sanitize_config(
for key in related[related_origin_cap:]: for key in related[related_origin_cap:]:
del origins[key] del origins[key]
domains[rp_id] = DomainConfig(rp_name=domain.rp_name, origins=origins) domains[rp_id] = DomainConfig(
rp_name=domain.rp_name, origins=origins, remote=remote
)
if not domains: if not domains:
raise ValueError("No servable domain in the stored configuration") raise ValueError("No servable domain in the stored configuration")
@@ -461,6 +492,17 @@ def _derive_site(
_registry: DomainRegistry | None = None _registry: DomainRegistry | None = None
_listen: list[str] | None = None _listen: list[str] | None = None
_rebuild_listeners: list = []
def add_rebuild_listener(fn) -> None:
"""Register fn(registry), called after every init_registry rebuild."""
_rebuild_listeners.append(fn)
def remove_rebuild_listener(fn) -> None:
if fn in _rebuild_listeners:
_rebuild_listeners.remove(fn)
def configure(*, listen: list[str] | None = None) -> None: def configure(*, listen: list[str] | None = None) -> None:
@@ -501,6 +543,12 @@ def init_registry(config: Config) -> DomainRegistry:
"""Build and install the global registry from a combined configuration.""" """Build and install the global registry from a combined configuration."""
global _registry global _registry
_registry = build(config) _registry = build(config)
for fn in _rebuild_listeners:
result = fn(_registry)
if asyncio.iscoroutine(result):
# init_registry runs within a running loop in every serving
# context (lifespan, tests, admin rebuild).
asyncio.get_running_loop().create_task(result)
return _registry return _registry
+39 -1
View File
@@ -12,7 +12,7 @@ immediately.
from fastapi import Body, FastAPI, Request from fastapi import Body, FastAPI, Request
from paskia import db, domains from paskia import db, domains
from paskia.db.structs import Config, DomainConfig, OriginEntry from paskia.db.structs import Config, DomainConfig, OriginEntry, RemoteConfig
from paskia.fastapi import authz from paskia.fastapi import authz
from paskia.fastapi.admin.errors import install_error_handlers from paskia.fastapi.admin.errors import install_error_handlers
from paskia.fastapi.response import MsgspecResponse from paskia.fastapi.response import MsgspecResponse
@@ -27,6 +27,14 @@ install_error_handlers(app)
def _domain_to_api(domain: domains.Domain) -> ApiDomain: def _domain_to_api(domain: domains.Domain) -> ApiDomain:
remote = domain.config.remote
if remote is not None:
# The sync token is a bearer secret: never echoed back
remote = RemoteConfig(
url=remote.url,
cache_ttl=remote.cache_ttl,
refresh_interval=remote.refresh_interval,
)
return ApiDomain( return ApiDomain(
rp_id=domain.rp_id, rp_id=domain.rp_id,
rp_name=domain.rp_name, rp_name=domain.rp_name,
@@ -34,6 +42,28 @@ def _domain_to_api(domain: domains.Domain) -> ApiDomain:
site_url=domain.site_url, site_url=domain.site_url,
auth_site_url=domain.auth_site_url, auth_site_url=domain.auth_site_url,
auth_host=domain.own_auth_host, auth_host=domain.own_auth_host,
remote=remote,
)
def _normalize_remote(
value, existing: RemoteConfig | None = None
) -> RemoteConfig | None:
"""Parse a remote object from the admin UI (raises on malformed).
An absent/empty token keeps the previously stored one — the token is
write-only over the API.
"""
if value is None:
return None
if not isinstance(value, dict) or not isinstance(value.get("url"), str):
raise ValueError("remote must be an object with a url")
token = str(value.get("token") or "") or (existing.token if existing else "")
return RemoteConfig(
url=value["url"].rstrip("/"),
token=token,
cache_ttl=int(value.get("cache_ttl") or 60),
refresh_interval=int(value.get("refresh_interval") or 300),
) )
@@ -123,6 +153,7 @@ async def admin_create_domain(
new = DomainConfig( new = DomainConfig(
rp_name=(payload.get("rp_name") or "").strip() or None, rp_name=(payload.get("rp_name") or "").strip() or None,
origins=_normalize_origins_map(payload.get("origins")), origins=_normalize_origins_map(payload.get("origins")),
remote=_normalize_remote(payload.get("remote")),
) )
config = db.data().config config = db.data().config
@@ -156,9 +187,15 @@ async def admin_update_domain(
if rp_id not in config.domains: if rp_id not in config.domains:
raise ValueError(f"Domain {rp_id} not found") raise ValueError(f"Domain {rp_id} not found")
current_remote = config.domains[rp_id].remote
updated = DomainConfig( updated = DomainConfig(
rp_name=(payload.get("rp_name") or "").strip() or None, rp_name=(payload.get("rp_name") or "").strip() or None,
origins=_normalize_origins_map(payload.get("origins")), origins=_normalize_origins_map(payload.get("origins")),
remote=(
_normalize_remote(payload["remote"], existing=current_remote)
if "remote" in payload
else current_remote
),
) )
would_be = Config( would_be = Config(
domains={k: updated if k == rp_id else v for k, v in config.domains.items()}, domains={k: updated if k == rp_id else v for k, v in config.domains.items()},
@@ -171,6 +208,7 @@ async def admin_update_domain(
rp_id, rp_id,
rp_name=updated.rp_name, rp_name=updated.rp_name,
origins=updated.origins, origins=updated.origins,
remote=updated.remote,
ctx=ctx, ctx=ctx,
) )
_rebuild_registry() _rebuild_registry()
+21 -9
View File
@@ -14,7 +14,7 @@ from fastapi import (
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from fastapi.security import HTTPBearer from fastapi.security import HTTPBearer
from paskia import authcode, db from paskia import authcode, db, satellite
from paskia._version import __version__ from paskia._version import __version__
from paskia.authsession import EXPIRES, get_reset, session_ctx from paskia.authsession import EXPIRES, get_reset, session_ctx
from paskia.domains import current_domain from paskia.domains import current_domain
@@ -122,8 +122,9 @@ async def validate_token(
if auth and renew: if auth and renew:
consumed = datetime.now(UTC) - ctx.session.validated consumed = datetime.now(UTC) - ctx.session.validated
if not timedelta(0) < consumed < _REFRESH_INTERVAL: if not timedelta(0) < consumed < _REFRESH_INTERVAL:
db.update_session( satellite.refresh_session(
ctx.session.key, ctx.session.key,
request.headers.get("host"),
ip=get_client_ip(request), ip=get_client_ip(request),
user_agent=request.headers.get("user-agent"), user_agent=request.headers.get("user-agent"),
validated=datetime.now(UTC), validated=datetime.now(UTC),
@@ -162,16 +163,16 @@ async def check_user(
No session cookie is read or written. Caller authentication is not required. No session cookie is read or written. Caller authentication is not required.
""" """
data = db.data() host = hostutil.normalize_host(request.headers.get("host"))
data = satellite.store_for_host(host)
try: try:
u = data.users[user_uuid] u = data.users[user_uuid]
role = u.role role = data.roles[u.role_uuid]
org = role.org org = data.orgs[role.org_uuid]
except KeyError: except KeyError:
raise HTTPException(status_code=404, detail="User not found") raise HTTPException(status_code=404, detail="User not found")
host = hostutil.normalize_host(request.headers.get("host")) org_perm_uuids = {p.uuid for p in data.permissions.values() if org.uuid in p.orgs}
org_perm_uuids = {p.uuid for p in org.permissions}
effective_perms = [] effective_perms = []
for perm_uuid in role.permission_set: for perm_uuid in role.permission_set:
@@ -212,7 +213,7 @@ def _remote_headers(ctx) -> dict[str, str]:
"Remote-Session-Expires": ( "Remote-Session-Expires": (
(ctx.session.validated + EXPIRES).isoformat().replace("+00:00", "Z") (ctx.session.validated + EXPIRES).isoformat().replace("+00:00", "Z")
), ),
"Remote-Credential": str(ctx.session.credential), "Remote-Credential": str(ctx.credential.uuid),
} }
@@ -351,8 +352,10 @@ async def api_user_info(
@app.get("/token-info") @app.get("/token-info")
async def token_info(credentials=Depends(bearer_auth)): async def token_info(request: Request, credentials=Depends(bearer_auth)):
"""Get reset/device-add token info. Pass token via Bearer header.""" """Get reset/device-add token info. Pass token via Bearer header."""
if (proxied := await satellite.forward_request(request)) is not None:
return proxied
if not credentials or not credentials.credentials: if not credentials or not credentials.credentials:
raise HTTPException(401, "Bearer token required") raise HTTPException(401, "Bearer token required")
token = credentials.credentials token = credentials.credentials
@@ -375,6 +378,10 @@ async def token_info(credentials=Depends(bearer_auth)):
@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 (proxied := await satellite.forward_request(request)) is not None:
if auth and proxied.status_code == 200:
satellite.evict_session(auth, request.headers.get("host"))
return proxied
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")
@@ -399,6 +406,11 @@ async def api_set_session(
if not auth or not auth.credentials: if not auth or not auth.credentials:
raise HTTPException(400, "Bearer token required") raise HTTPException(400, "Bearer token required")
if (proxied := await satellite.forward_request(request)) is not None:
# The exchange code lives in the remote's RAM; redeem it there. The
# session itself reaches the replica via the sync channel.
return proxied
host = hostutil.normalize_host(request.headers.get("host", "")) host = hostutil.normalize_host(request.headers.get("host", ""))
if not host: if not host:
raise HTTPException(400, "Host header required") raise HTTPException(400, "Host header required")
+5
View File
@@ -63,6 +63,11 @@ class DispatchMiddleware:
host = _header(scope, "host") host = _header(scope, "host")
host_domain = registry.resolve(host) host_domain = registry.resolve(host)
if host_domain is None: if host_domain is None:
# The sync endpoint is server-to-server and token-gated: the
# satellite may reach us via an address outside our domains.
if scope.get("path") == "/auth/api/sync/ws" and registry.domains:
await self._dispatch(scope, receive, send, registry.domains[0])
return
await send({"type": "websocket.close", "code": _WS_CLOSE_POLICY_VIOLATION}) await send({"type": "websocket.close", "code": _WS_CLOSE_POLICY_VIOLATION})
return return
+14 -6
View File
@@ -7,11 +7,11 @@ from fastapi import FastAPI, HTTPException, Request, Response
from fastapi.responses import FileResponse, RedirectResponse from fastapi.responses import FileResponse, RedirectResponse
from fastapi_vue import env from fastapi_vue import env
from paskia import authcode, db, domains, remoteauth from paskia import authcode, db, domains, remoteauth, satellite
from paskia.bootstrap import bootstrap_if_needed from paskia.bootstrap import bootstrap_if_needed
from paskia.db.background import start_background, stop_background from paskia.db.background import start_background, stop_background
from paskia.db.lifecycle import kanta from paskia.db.lifecycle import kanta
from paskia.fastapi import admin, api, auth_host, oid, ws from paskia.fastapi import admin, api, auth_host, oid, sync, ws
from paskia.fastapi.admin.adminapp import adminapp from paskia.fastapi.admin.adminapp import adminapp
from paskia.fastapi.dispatch import DispatchMiddleware from paskia.fastapi.dispatch import DispatchMiddleware
@@ -29,10 +29,12 @@ _EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples"
async def lifespan(app: FastAPI): # pragma: no cover - startup path async def lifespan(app: FastAPI): # pragma: no cover - startup path
"""Application lifespan: open the combined database and build the domain registry. """Application lifespan: open the combined database and build the domain registry.
Process-global serve parameters (listen endpoints) are passed via the Process-global serve parameters (listen endpoints, save flag) are passed
PASKIA_CONFIG JSON env variable (set by the CLI entrypoint) so that via the PASKIA_CONFIG JSON env variable (set by the CLI entrypoint) so
uvicorn reload / multiprocess workers derive site URLs the same way. that uvicorn reload / multiprocess workers derive site URLs the same
Domain configuration is read from the database. way. With the save flag set, the listen endpoints are persisted here —
the CLI never opens the database read-write. Domain configuration is
read from the database.
""" """
cfg = serve_config() cfg = serve_config()
domains.configure(listen=cfg.listen) domains.configure(listen=cfg.listen)
@@ -41,10 +43,14 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
Path(kanta.filename).parent.mkdir, parents=True, exist_ok=True Path(kanta.filename).parent.mkdir, parents=True, exist_ok=True
) )
async with kanta: async with kanta:
if cfg.save:
with kanta.transaction("serve:save_listen"):
db.data().config.listen = cfg.listen
try: try:
domains.init_registry(db.data().config) domains.init_registry(db.data().config)
await remoteauth.init() await remoteauth.init()
await authcode.start() await authcode.start()
await satellite.manager.start()
except ValueError as e: except ValueError as e:
logging.error(f"⚠️ {e}") logging.error(f"⚠️ {e}")
# Re-raise to fail fast # Re-raise to fail fast
@@ -55,6 +61,7 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
await start_background() await start_background()
yield yield
await stop_background() await stop_background()
await satellite.manager.stop()
await authcode.stop() await authcode.stop()
@@ -78,6 +85,7 @@ app.middleware("http")(auth_host.redirect_middleware)
app.add_middleware(DispatchMiddleware) app.add_middleware(DispatchMiddleware)
app.mount("/auth/api/admin/", admin.app) app.mount("/auth/api/admin/", admin.app)
app.mount("/auth/api/sync", sync.app)
app.mount("/auth/api/", api.app) app.mount("/auth/api/", api.app)
app.mount("/auth/ws/", ws.app) app.mount("/auth/ws/", ws.app)
app.mount("/auth/oidc/", oid.app) app.mount("/auth/oidc/", oid.app)
+9 -1
View File
@@ -20,7 +20,7 @@ from fastapi import Depends, FastAPI, Form, HTTPException, Request
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from fastapi.security import HTTPBearer from fastapi.security import HTTPBearer
from paskia import authcode, db from paskia import authcode, db, satellite
from paskia.db.structs import OIDC, Session from paskia.db.structs import OIDC, Session
from paskia.util import avatar, oidjwt from paskia.util import avatar, oidjwt
from paskia.util.crypto import hash_secret from paskia.util.crypto import hash_secret
@@ -30,6 +30,14 @@ _logger = logging.getLogger(__name__)
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None) app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
@app.middleware("http")
async def proxy_remote_domain(request: Request, call_next):
"""OIDC key material and sessions stay on the remote; proxy everything."""
if (proxied := await satellite.forward_request(request)) is not None:
return proxied
return await call_next(request)
def _provider() -> OIDC: def _provider() -> OIDC:
"""Return the instance-global OIDC provider state.""" """Return the instance-global OIDC provider state."""
return db.data().oidc return db.data().oidc
+94
View File
@@ -0,0 +1,94 @@
"""Sync WebSocket endpoint: serves snapshots and live events to satellites.
Token-gated via PASKIA_SYNC_TOKENS (env); closed when unset. All state is
RAM-only (syncfeed); the database schema is untouched. Protocol: snapshot
chunks per table, `ready`, then live upsert/delete events; the client sends
session_refresh write-backs. Reconnects always restart from a snapshot.
"""
import asyncio
import logging
from datetime import datetime
import msgspec
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from paskia import db, syncfeed
_logger = logging.getLogger(__name__)
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
async def _send(ws: WebSocket, message: dict) -> None:
await ws.send_bytes(syncfeed.encode(message))
async def _apply_client_message(message: dict) -> None:
"""Satellite write-behind: session refresh (validated/ip/user-agent)."""
if message.get("type") != "session_refresh":
return
key = message.get("key") or ""
session = db.data().sessions.get(key)
if session is None:
return
try:
validated = msgspec.convert(message.get("validated"), datetime)
except msgspec.ValidationError:
return
db.update_session(
key,
ip=message.get("ip") or None,
user_agent=message.get("user_agent") or None,
validated=validated,
)
@app.websocket("/ws")
async def sync_websocket(ws: WebSocket):
tokens = syncfeed.tokens_from_env()
auth = ws.headers.get("authorization", "")
if not tokens or auth.removeprefix("Bearer ").strip() not in tokens:
await ws.close(code=1008)
return
await ws.accept()
queue = syncfeed.subscribe()
try:
data = db.data()
for table in syncfeed.TABLES:
await _send(
ws,
{
"type": "snapshot",
"table": table,
"items": [
[str(key), msgspec.to_builtins(obj)]
for key, obj in getattr(data, table).items()
],
},
)
await _send(ws, {"type": "ready"})
sender = asyncio.create_task(_pump(ws, queue))
try:
while True:
await _apply_client_message(
msgspec.json.decode(await ws.receive_bytes())
)
finally:
sender.cancel()
except WebSocketDisconnect:
pass
except Exception:
_logger.exception("Sync WebSocket failed")
finally:
syncfeed.unsubscribe(queue)
async def _pump(ws: WebSocket, queue: asyncio.Queue) -> None:
try:
while True:
await _send(ws, await queue.get())
except WebSocketDisconnect, RuntimeError, asyncio.CancelledError:
pass
+344
View File
@@ -0,0 +1,344 @@
"""Satellite side of remote domains: RAM-only replicas + host dispatch.
Domains configured with ``DomainConfig.remote`` are backed by a remote
paskia instance. This module owns the whole feature: it resolves which
store serves a request host (local DB or the remote's read replica),
dispatches session writes (refresh write-behind, logout eviction), and
forwards requests the satellite cannot answer (exchange-code redemption,
OIDC, reset tokens) to the remote.
A replica is a plain DB instance, never persisted, fed by a sync
WebSocket (snapshot on connect, then live events) and swept for
expired sessions locally. While the channel is down the replica stays trusted for the
domain's cache_ttl, then reads fail with RemoteUnavailable (fail-closed;
a large cache_ttl gives fail-open behavior bounded by session expiry).
"""
import asyncio
import contextlib
import logging
import time
from datetime import UTC, datetime
from uuid import UUID
import httpx
import msgspec
import websockets
from fastapi import HTTPException, Request, Response
from paskia import db, domains
from paskia.config import SESSION_LIFETIME
from paskia.db.structs import (
DB,
Credential,
Org,
Permission,
RemoteConfig,
Role,
Session,
User,
)
from paskia.util.crypto import hash_secret
_logger = logging.getLogger(__name__)
_TABLES = {
"permissions": (Permission, True),
"orgs": (Org, True),
"roles": (Role, True),
"users": (User, True),
"credentials": (Credential, True),
"sessions": (Session, False),
}
_RECONNECT_DELAY = 5
_SWEEP_INTERVAL = 60
class RemoteReplica:
"""One remote instance's replica, its sync client and write-behind queue."""
def __init__(self, remote: RemoteConfig):
self.remote = remote
self.db = DB()
self.last_contact = 0.0 # monotonic time the feed last went down
self.connected = False
self._pending_refresh: dict[str, dict] = {}
self._refresh_signal = asyncio.Event()
self._task: asyncio.Task | None = None
self._sweeper: asyncio.Task | None = None
self._stopped = True
def available(self) -> bool:
"""Synced, and connected now or within cache_ttl of the disconnect."""
if not self.last_contact:
return False
return self.connected or (
time.monotonic() - self.last_contact <= self.remote.cache_ttl
)
def refresh_session(
self, key: str, validated, ip: str | None, user_agent: str | None
) -> None:
"""Apply a /validate refresh locally and queue it for the remote."""
session = self.db.sessions.get(key)
if session is not None:
session.validated = validated
if ip is not None:
session.ip = ip
if user_agent is not None:
session.user_agent = user_agent
self._pending_refresh[key] = {
"type": "session_refresh",
"key": key,
"validated": msgspec.to_builtins(validated),
"ip": ip,
"user_agent": user_agent,
}
self._refresh_signal.set()
async def start(self) -> None:
self._stopped = False
self._task = asyncio.create_task(self._run())
self._sweeper = asyncio.create_task(self._sweep())
async def stop(self) -> None:
self._stopped = True
for task in (self._task, self._sweeper):
if task:
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
async def _sweep(self) -> None:
while True:
await asyncio.sleep(_SWEEP_INTERVAL)
limit = datetime.now(UTC) - SESSION_LIFETIME
for key in [k for k, s in self.db.sessions.items() if s.validated < limit]:
del self.db.sessions[key]
async def _run(self) -> None:
while not self._stopped:
try:
await self._connect()
except asyncio.CancelledError:
raise
except Exception as e:
_logger.info("Sync to %s failed: %s", self.remote.url, e)
if self.connected:
# The TTL clock starts when the feed goes down, not at the
# last message — an idle connection is healthy.
self.connected = False
self.last_contact = time.monotonic()
if not self._stopped:
await asyncio.sleep(_RECONNECT_DELAY)
async def _connect(self) -> None:
ws_url = self.remote.url.replace("http", "ws", 1) + "/auth/api/sync/ws"
async with websockets.connect(
ws_url,
additional_headers={"Authorization": f"Bearer {self.remote.token}"},
# Prompt dead-peer detection: availability semantics count on it
ping_interval=5,
ping_timeout=5,
) as ws:
sender = asyncio.create_task(self._send_loop(ws))
staging: DB | None = None
ready_at = 0.0
try:
while True:
if ready_at:
# Periodic reconnects give full-snapshot reconciliation
remaining = self.remote.refresh_interval - (
time.monotonic() - ready_at
)
if remaining <= 0:
return
try:
message = msgspec.json.decode(
await asyncio.wait_for(ws.recv(), remaining)
)
except TimeoutError:
return # periodic resync: reconnect for a snapshot
else:
message = msgspec.json.decode(await ws.recv())
mtype = message.get("type")
if mtype == "snapshot":
staging = staging or DB()
for key, fields in message["items"]:
_apply(staging, message["table"], key, fields)
elif mtype == "event":
if staging is not None:
raise ValueError("sync: event before ready")
_apply(
self.db,
message["table"],
message["key"],
message.get("fields"),
)
elif mtype == "ready":
if staging is not None:
self.db = staging
staging = None
self.connected = True
self.last_contact = ready_at = time.monotonic()
finally:
sender.cancel()
with contextlib.suppress(asyncio.CancelledError):
await sender
async def _send_loop(self, ws) -> None:
while True:
self._refresh_signal.clear()
while self._pending_refresh:
_, message = self._pending_refresh.popitem()
await ws.send(msgspec.json.encode(message))
await self._refresh_signal.wait()
def _apply(replica: DB, table: str, key: str, fields: dict | None) -> None:
"""Apply an upsert (fields given) or delete (fields None) to a replica."""
cls, uuid_key = _TABLES[table]
store = getattr(replica, table)
store_key = UUID(key) if uuid_key else key
if fields is None:
store.pop(store_key, None)
return
obj = msgspec.convert(fields, cls)
if uuid_key:
obj.uuid = store_key
else:
obj.key = key
store[store_key] = obj
class SatelliteManager:
"""Replicas keyed by remote URL; domains sharing a remote share one."""
def __init__(self):
self.replicas: dict[str, RemoteReplica] = {}
async def start(self) -> None:
domains.add_rebuild_listener(self.reconcile)
await self.reconcile(domains.registry())
async def stop(self) -> None:
domains.remove_rebuild_listener(self.reconcile)
for replica in self.replicas.values():
await replica.stop()
self.replicas.clear()
async def reconcile(self, registry: domains.DomainRegistry) -> None:
"""Start/stop replicas to match the configured remote domains."""
wanted = {}
for domain in registry.domains:
if domain.remote is not None:
wanted.setdefault(domain.remote.url, domain.remote)
for url in list(self.replicas):
if url not in wanted:
await self.replicas.pop(url).stop()
for url, remote in wanted.items():
replica = self.replicas.get(url)
if replica is None or replica.remote != remote:
if replica is not None:
await replica.stop()
replica = RemoteReplica(remote)
self.replicas[url] = replica
await replica.start()
manager = SatelliteManager()
# -------------------------------------------------------------------------
# Host-keyed dispatch: the only interface the rest of the app uses
# -------------------------------------------------------------------------
def replica_for_host(host: str | None) -> RemoteReplica | None:
"""The replica serving this host, or None for locally served hosts."""
domain = domains.registry().resolve(host)
if domain is None or domain.remote is None:
return None
return manager.replicas.get(domain.remote.url)
def store_for_host(host: str | None) -> DB:
"""The data store to read for a request host: the local database, or
the replica of the remote backing the host's domain."""
replica = replica_for_host(host)
if replica is None:
return db.data()
if not replica.available():
raise HTTPException(503, "Remote authentication service unavailable")
return replica.db
def refresh_session(
key, host: str | None, ip: str, user_agent: str, validated, ctx=None
):
"""/validate refresh: write-behind for remote domains, else local DB."""
replica = replica_for_host(host)
if replica is not None:
replica.refresh_session(key, validated, ip, user_agent)
else:
db.update_session(
key, ip=ip, user_agent=user_agent, validated=validated, ctx=ctx
)
def evict_session(auth: str, host: str | None) -> None:
"""Drop a session from the replica (its remote deletion arrives via sync)."""
replica = replica_for_host(host)
if replica is not None:
replica.db.sessions.pop(hash_secret("cookie", auth), None)
_TIMEOUT = httpx.Timeout(15.0, connect=5.0)
_HOP_BY_HOP = {
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailers",
"transfer-encoding",
"upgrade",
"content-length",
"accept-encoding",
"content-encoding",
}
_clients: dict[str, httpx.AsyncClient] = {}
async def forward_request(request: Request) -> Response | None:
"""Forward the request to its domain's remote, or None when local.
The original Host header is preserved so the remote dispatches to the
same domain (sessions are host-bound). The user's cookie authenticates
the forwarded call; the satellite needs no credentials of its own.
"""
domain = domains.registry().resolve(request.headers.get("host"))
if domain is None or domain.remote is None:
return None
url = domain.remote.url
client = _clients.get(url)
if client is None:
client = _clients[url] = httpx.AsyncClient(base_url=url, timeout=_TIMEOUT)
upstream = await client.request(
request.method,
request.url.path,
params=request.url.query,
content=await request.body(),
headers={
k: v for k, v in request.headers.items() if k.lower() not in _HOP_BY_HOP
},
)
response = Response(content=upstream.content, status_code=upstream.status_code)
# Raw headers to preserve repeated Set-Cookie
response.raw_headers = [
(k, v) for k, v in upstream.headers.raw if k.decode().lower() not in _HOP_BY_HOP
]
return response
+60
View File
@@ -0,0 +1,60 @@
"""RAM-only change feed letting satellite instances mirror this server.
Nothing here touches the database file: committed mutations are pushed to
connected satellites over the sync WebSocket (fastapi/sync.py). Satellites
authenticate with a token from the PASKIA_SYNC_TOKENS environment variable
(comma-separated); with the variable unset the sync endpoint stays closed.
There is deliberately no replay log: snapshots are small, so a reconnecting
satellite simply takes a fresh one.
"""
import asyncio
import os
import msgspec
# Tables mirrored by satellites (reset tokens, OIDC data and domain config
# are instance-local and never replicated).
TABLES = ("permissions", "orgs", "roles", "users", "credentials", "sessions")
_subscribers: set[asyncio.Queue] = set()
def emit(table: str, key: str, obj) -> None:
"""Publish an upsert (obj given) or delete (obj None) to subscribers."""
event = {
"type": "event",
"table": table,
"key": key,
"fields": msgspec.to_builtins(obj) if obj is not None else None,
}
for queue in list(_subscribers):
try:
queue.put_nowait(event)
except asyncio.QueueFull:
# Slow consumer: drop it; the client reconnects and resyncs.
_subscribers.discard(queue)
def subscribe() -> asyncio.Queue:
queue: asyncio.Queue = asyncio.Queue(maxsize=1000)
_subscribers.add(queue)
return queue
def unsubscribe(queue: asyncio.Queue) -> None:
_subscribers.discard(queue)
def tokens_from_env() -> set[str]:
"""Accepted sync tokens (PASKIA_SYNC_TOKENS, comma-separated)."""
return {
t.strip()
for t in os.environ.get("PASKIA_SYNC_TOKENS", "").split(",")
if t.strip()
}
def encode(message: dict) -> bytes:
return msgspec.json.encode(message)
+10 -1
View File
@@ -14,7 +14,15 @@ import msgspec
from uarite import uaparse from uarite import uaparse
from paskia import db from paskia import db
from paskia.db.structs import Credential, Org, OriginEntry, Permission, Role, User from paskia.db.structs import (
Credential,
Org,
OriginEntry,
Permission,
RemoteConfig,
Role,
User,
)
# ------------------------------------------------------------------------- # -------------------------------------------------------------------------
# API structs - inherit from db structs, add uuid for serialization # API structs - inherit from db structs, add uuid for serialization
@@ -194,6 +202,7 @@ class ApiDomain(msgspec.Struct):
site_url: str site_url: str
auth_site_url: str auth_site_url: str
auth_host: str | None auth_host: str | None
remote: RemoteConfig | None = None
class ApiTokenInfo(msgspec.Struct, omit_defaults=True): class ApiTokenInfo(msgspec.Struct, omit_defaults=True):
+5 -3
View File
@@ -2,9 +2,10 @@
Domain configuration lives in the database (``Config.domains``); the Domain configuration lives in the database (``Config.domains``); the
``PASKIA_CONFIG`` environment variable only carries the effective listen ``PASKIA_CONFIG`` environment variable only carries the effective listen
endpoints so that child processes (uvicorn reload / workers) derive site endpoints and whether to persist them, so that child processes (uvicorn
URLs the same way the parent did. The CLI entry point mutates the bound reload / workers) derive site URLs the same way the parent did. The CLI
object before ``server.run()`` calls ``teleport()`` to pass it on. entry point mutates the bound object before ``server.run()`` calls
``teleport()`` to pass it on.
""" """
import msgspec import msgspec
@@ -15,6 +16,7 @@ class ServeConfig(msgspec.Struct):
"""Process-global serve parameters.""" """Process-global serve parameters."""
listen: list[str] | None = None listen: list[str] | None = None
save: bool = False # Persist listen to the stored config on startup
def serve_config() -> ServeConfig: def serve_config() -> ServeConfig:
+11 -5
View File
@@ -1,6 +1,6 @@
"""User information formatting and retrieval logic.""" """User information formatting and retrieval logic."""
from paskia import aaguid, db from paskia import aaguid, satellite
from paskia.db import SessionContext from paskia.db import SessionContext
from paskia.util import avatar, hostutil from paskia.util import avatar, hostutil
from paskia.util.apistructs import ( from paskia.util.apistructs import (
@@ -43,24 +43,30 @@ async def build_user_info(
ctx: SessionContext | None = None, ctx: SessionContext | None = None,
) -> ApiUserDetail: ) -> ApiUserDetail:
"""Build user info struct for authenticated users.""" """Build user info struct for authenticated users."""
user = db.data().users[user_uuid] data = satellite.store_for_host(request_host)
user = data.users[user_uuid]
normalized_host = hostutil.normalize_host(request_host) normalized_host = hostutil.normalize_host(request_host)
user_sessions = [s for s in data.sessions.values() if s.user_uuid == user_uuid]
user_credentials = [
c for c in data.credentials.values() if c.user_uuid == user_uuid
]
sessions = { sessions = {
s.key: ApiUserSession.from_db( s.key: ApiUserSession.from_db(
s, s,
current_key=session_key, current_key=session_key,
normalized_host=normalized_host, normalized_host=normalized_host,
) )
for s in user.sessions for s in user_sessions
} }
return ApiUserDetail( return ApiUserDetail(
user=ApiUser.from_db(user, avatar_url=avatar.avatar_browser_url(user.uuid)), user=ApiUser.from_db(user, avatar_url=avatar.avatar_browser_url(user.uuid)),
credentials={c.uuid: c for c in user.credentials}, credentials={c.uuid: c for c in user_credentials},
aaguid_info={ aaguid_info={
k: ApiAaguidInfo(**v) k: ApiAaguidInfo(**v)
for k, v in aaguid.filter(c.aaguid for c in user.credentials).items() for k, v in aaguid.filter(c.aaguid for c in user_credentials).items()
}, },
sessions=sessions, sessions=sessions,
permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions} permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions}
+1 -1
View File
@@ -21,7 +21,7 @@ dependencies = [
"pyjwt[crypto]>=2.11.0", "pyjwt[crypto]>=2.11.0",
"jsondiff>=2.2.1", "jsondiff>=2.2.1",
"msgspec>=0.20.0", "msgspec>=0.20.0",
"fastapi-vue~=1.7.1", "fastapi-vue~=1.7.2",
"kanta>=0.9.2", "kanta>=0.9.2",
"uarite>=0.2.1", "uarite>=0.2.1",
] ]
+10 -13
View File
@@ -10,23 +10,20 @@ from pathlib import Path
MIN_NODE_VERSION = 20 MIN_NODE_VERSION = 20
# Duplicated from fastapi_vue.logging because build environment is isolated
_LEVEL_EMOJI = {
logging.DEBUG: "🐛",
logging.INFO: "🔷",
logging.WARNING: "",
logging.ERROR: "🛑",
logging.CRITICAL: "🚨",
}
class _Formatter(logging.Formatter): class _Formatter(logging.Formatter):
"""Emoji level prefix formatter, mirroring fastapi_vue.logging.Formatter.""" """Prefix formatter, intentionally different from fastapi_vue.logging.
INFO and below pass through unprefixed so messages can use their own
markings (>>>, ###); WARNING and above get an emoji prefix.
"""
def format(self, record: logging.LogRecord) -> str: def format(self, record: logging.LogRecord) -> str:
emoji = _LEVEL_EMOJI.get(record.levelno) if record.levelno >= logging.ERROR:
prefix = f"{emoji} " if emoji else f"{record.levelname}: " return f"🛑 {record.getMessage()}"
return prefix + record.getMessage() if record.levelno >= logging.WARNING:
return f"💣 {record.getMessage()}"
return record.getMessage()
_handler = logging.StreamHandler() _handler = logging.StreamHandler()
+1 -1
View File
@@ -139,7 +139,7 @@ async def ready(url: str, path: str = "", max_attempts: int = 50) -> None:
for attempt in range(max_attempts): for attempt in range(max_attempts):
if await http_get_server(f"{url}{path}", timeout=1.0) is not None: if await http_get_server(f"{url}{path}", timeout=1.0) is not None:
logger.info(" Backend ready!") logger.info("🟢 Backend ready!")
return return
if attempt == max_attempts - 1: if attempt == max_attempts - 1:
logger.error("Backend at %s didn't start in time", url) logger.error("Backend at %s didn't start in time", url)
+13 -3
View File
@@ -172,6 +172,7 @@ def test_serve_uses_stored_config(run_cli, tmp_path):
assert calls["listen"] is None # stored listen (None) used assert calls["listen"] is None # stored listen (None) used
serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig) serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig)
assert serve.listen is None assert serve.listen is None
assert serve.save is False
def test_serve_listen_override_not_persisted(run_cli, tmp_path): def test_serve_listen_override_not_persisted(run_cli, tmp_path):
@@ -181,24 +182,33 @@ def test_serve_listen_override_not_persisted(run_cli, tmp_path):
assert calls["listen"] == ["4403"] assert calls["listen"] == ["4403"]
serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig) serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig)
assert serve.listen == ["4403"] assert serve.listen == ["4403"]
assert serve.save is False
# Stored config keeps the original listen value # Stored config keeps the original listen value
assert stored_config(tmp_path).listen == ["4402"] assert stored_config(tmp_path).listen == ["4402"]
def test_serve_listen_save_persists(run_cli, tmp_path): def test_serve_listen_save_persists(run_cli, tmp_path):
"""--save teleports the save flag; the app persists, the CLI is read-only."""
run_cli("init", "--listen", "4402") run_cli("init", "--listen", "4402")
calls = run_cli("--listen", "4403", "--save") calls = run_cli("--listen", "4403", "--save")
assert calls["listen"] == ["4403"] assert calls["listen"] == ["4403"]
assert stored_config(tmp_path).listen == ["4403"] serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig)
assert serve.listen == ["4403"]
assert serve.save is True
# The CLI itself does not write the database
assert stored_config(tmp_path).listen == ["4402"]
def test_serve_listen_save_clear(run_cli, tmp_path): def test_serve_listen_save_clear(run_cli, tmp_path):
"""--listen "" --save clears the stored endpoints (back to default).""" """--listen "" --save teleports a clear (back to default) for the app."""
run_cli("init", "--listen", "4402") run_cli("init", "--listen", "4402")
run_cli("--listen", "", "--save") run_cli("--listen", "", "--save")
assert stored_config(tmp_path).listen is None serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig)
assert serve.listen is None
assert serve.save is True
assert stored_config(tmp_path).listen == ["4402"]
def test_serve_suggests_migrate_when_legacy_present(run_cli, tmp_path): def test_serve_suggests_migrate_when_legacy_present(run_cli, tmp_path):
+427
View File
@@ -0,0 +1,427 @@
"""Tests for remote (satellite) domains: config, replica application, feed."""
import secrets
import time
from datetime import UTC, datetime
from uuid import UUID
import httpx
import msgspec
import pytest
import pytest_asyncio
from fastapi import Response
import paskia.db.operations as ops_db
from paskia import domains, satellite, syncfeed
from paskia.db.structs import (
DB,
Config,
Credential,
DomainConfig,
Org,
OriginEntry,
Permission,
RemoteConfig,
Role,
Session,
User,
)
from paskia.fastapi.mainapp import app
from paskia.fastapi.session import AUTH_COOKIE_NAME
from paskia.util.crypto import hash_secret
from .conftest import TEST_RP_ID
REMOTE_URL = "http://remote.test"
def _remote_domain_config(**kw) -> Config:
return Config(
domains={
TEST_RP_ID: DomainConfig(origins={f"**.{TEST_RP_ID}": True}),
"example.com": DomainConfig(
origins={
"**.example.com": True,
"auth.example.com": OriginEntry(auth_host=True),
},
remote=RemoteConfig(url=REMOTE_URL, token="t", **kw),
),
}
)
def test_remote_domain_valid():
domains.validate_config(_remote_domain_config())
def test_remote_domain_requires_auth_host():
config = _remote_domain_config()
config.domains["example.com"].origins = {"**.example.com": True}
with pytest.raises(ValueError, match="auth host"):
domains.validate_config(config)
def test_remote_domain_requires_http_url():
config = _remote_domain_config()
config.domains["example.com"].remote.url = "ftp://x"
with pytest.raises(ValueError, match="http"):
domains.validate_config(config)
def test_sanitize_preserves_remote():
config, warnings = domains.sanitize_config(_remote_domain_config())
assert not warnings
assert config.domains["example.com"].remote.url == REMOTE_URL
def test_apply_upsert_and_delete():
replica = DB()
user = User.create(display_name="U", role=UUID(int=1))
user.uuid = UUID(int=2)
satellite._apply(replica, "users", str(user.uuid), _builtins(user))
assert replica.users[user.uuid].display_name == "U"
satellite._apply(replica, "users", str(user.uuid), None)
assert not replica.users
def _builtins(obj):
return msgspec.to_builtins(obj)
def test_apply_session_roundtrip():
"""Sessions keep their string key and datetime/UUID fields."""
replica = DB()
session = Session.create(
user=UUID(int=1),
credential=UUID(int=2),
key=hash_secret("cookie", "sekret"),
host="app2.example.com",
ip="127.0.0.1",
user_agent="ua",
validated=datetime.now(UTC),
rp_id="example.com",
)
satellite._apply(replica, "sessions", session.key, _builtins(session))
stored = replica.sessions[session.key]
assert stored.host == "app2.example.com"
assert stored.validated == session.validated
assert stored.user_uuid == UUID(int=1)
def test_apply_credential_bytes_roundtrip():
"""credential_id/public_key are bytes over the wire (base64 in JSON)."""
replica = DB()
cred = Credential.create(
credential_id=secrets.token_bytes(32),
user=UUID(int=1),
aaguid=UUID(int=0),
public_key=secrets.token_bytes(64),
sign_count=3,
rp_id="example.com",
)
cred.uuid = UUID(int=9)
# Simulate the full wire path: builtins -> JSON -> builtins
wire = msgspec.json.decode(msgspec.json.encode(_builtins(cred)))
satellite._apply(replica, "credentials", str(cred.uuid), wire)
stored = replica.credentials[cred.uuid]
assert stored.credential_id == cred.credential_id
assert stored.public_key == cred.public_key
assert stored.sign_count == 3
def test_feed_emit_to_subscribers():
queue = syncfeed.subscribe()
try:
user = User.create(display_name="A", role=UUID(int=1))
syncfeed.emit("users", "k1", user)
syncfeed.emit("users", "k1", None)
assert queue.get_nowait()["fields"]["display_name"] == "A"
assert queue.get_nowait()["fields"] is None
finally:
syncfeed.unsubscribe(queue)
def test_feed_drops_full_queue():
queue = syncfeed.subscribe()
try:
for i in range(1001):
syncfeed.emit("users", f"k{i}", None)
assert queue.qsize() == 1000
syncfeed.emit("users", "k1001", None) # subscriber already dropped
assert queue.qsize() == 1000
finally:
syncfeed.unsubscribe(queue)
@pytest.mark.asyncio
async def test_operations_emit_events(test_db):
"""Writes through db.operations land on the sync feed."""
queue = syncfeed.subscribe()
try:
user = next(iter(test_db.users.values()))
ops_db.update_user_display_name(user.uuid, "Renamed")
event = queue.get_nowait()
assert event["table"] == "users"
assert event["key"] == str(user.uuid)
assert event["fields"]["display_name"] == "Renamed"
finally:
syncfeed.unsubscribe(queue)
@pytest.mark.asyncio
async def test_replica_refresh_and_evict():
replica = satellite.RemoteReplica(RemoteConfig(url=REMOTE_URL, token="t"))
token = secrets.token_urlsafe(12)
session = Session.create(
user=UUID(int=1),
credential=UUID(int=2),
key=hash_secret("cookie", token),
host="app2.example.com",
ip="1.1.1.1",
user_agent="ua",
validated=datetime(2020, 1, 1, tzinfo=UTC),
)
replica.db.sessions[session.key] = session
now = datetime.now(UTC)
replica.refresh_session(session.key, now, "2.2.2.2", "new-ua")
assert replica.db.sessions[session.key].validated == now
queued = replica._pending_refresh[session.key]
assert queued["type"] == "session_refresh"
assert queued["ip"] == "2.2.2.2"
# Host-keyed dispatch eviction (the replica's domain is resolved by host)
domains.configure(listen=["localhost:4401"])
domains.init_registry(_remote_domain_config())
satellite.manager.replicas[REMOTE_URL] = replica
satellite.evict_session(token, "app2.example.com")
assert not replica.db.sessions
satellite.manager.replicas.pop(REMOTE_URL)
def test_availability_gate():
replica = satellite.RemoteReplica(
RemoteConfig(url=REMOTE_URL, token="t", cache_ttl=60)
)
assert not replica.available() # never synced
replica.last_contact = time.monotonic()
assert replica.available()
# -------------------------------------------------------------------------
# API-level: endpoints served from an injected replica
# -------------------------------------------------------------------------
def _replica_db() -> tuple[DB, str]:
"""A replica DB holding one org/role/perm/user/credential/session."""
replica = DB()
org = Org.create(display_name="Org")
org.uuid = UUID(int=101)
replica.orgs[org.uuid] = org
perm = Permission.create(scope="auth:admin", display_name="Admin")
perm.uuid = UUID(int=102)
perm.orgs[org.uuid] = True
replica.permissions[perm.uuid] = perm
role = Role.create(org=org.uuid, display_name="Admins", permissions={perm.uuid})
role.uuid = UUID(int=103)
replica.roles[role.uuid] = role
user = User.create(display_name="Remote Admin", role=role.uuid)
user.uuid = UUID(int=104)
replica.users[user.uuid] = user
cred = Credential.create(
credential_id=b"cid",
user=user.uuid,
aaguid=UUID(int=0),
public_key=b"pk",
sign_count=0,
rp_id="example.com",
)
cred.uuid = UUID(int=105)
replica.credentials[cred.uuid] = cred
secret = secrets.token_urlsafe(12)
session = Session.create(
user=user.uuid,
credential=cred.uuid,
key=hash_secret("cookie", secret),
host="app2.example.com",
ip="127.0.0.1",
user_agent="pytest",
validated=datetime.now(UTC),
rp_id="example.com",
)
replica.sessions[session.key] = session
return replica, secret
@pytest_asyncio.fixture
async def remote_client(test_db):
"""ASGI client with example.com as a remote domain on a warm replica."""
config = _remote_domain_config()
domains.configure(listen=["localhost:4401"])
domains.init_registry(config)
replica_db, secret = _replica_db()
replica = satellite.RemoteReplica(RemoteConfig(url=REMOTE_URL, token="t"))
replica.db = replica_db
replica.last_contact = time.monotonic()
replica.connected = True
satellite.manager.replicas[REMOTE_URL] = replica
transport = httpx.ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport, base_url="http://localhost:4401"
) as client:
yield client, secret, replica
satellite.manager.replicas.pop(REMOTE_URL, None)
@pytest.mark.asyncio
async def test_forward_served_from_replica(remote_client):
client, secret, _ = remote_client
r = await client.get(
"/auth/api/forward?perm=auth:admin",
headers={"Host": "app2.example.com", "Cookie": f"{AUTH_COOKIE_NAME}={secret}"},
)
assert r.status_code == 204
assert r.headers["remote-name"] == "Remote Admin"
assert r.headers["remote-groups"] == "auth:admin"
@pytest.mark.asyncio
async def test_forward_replica_denies_missing_perm(remote_client):
client, secret, _ = remote_client
r = await client.get(
"/auth/api/forward?perm=other:scope",
headers={"Host": "app2.example.com", "Cookie": f"{AUTH_COOKIE_NAME}={secret}"},
)
assert r.status_code == 403
@pytest.mark.asyncio
async def test_validate_renews_locally_and_queues_writebehind(remote_client):
client, secret, replica = remote_client
session = next(iter(replica.db.sessions.values()))
session.validated = datetime(2020, 1, 1, tzinfo=UTC) # force refresh threshold
r = await client.post(
"/auth/api/validate",
headers={"Host": "app2.example.com", "Cookie": f"{AUTH_COOKIE_NAME}={secret}"},
)
assert r.status_code == 200
assert r.json()["renewed"] is True
assert session.validated.year > 2020 # applied to the replica
queued = replica._pending_refresh[session.key]
assert queued["type"] == "session_refresh"
@pytest.mark.asyncio
async def test_remote_domain_503_when_replica_stale(remote_client):
client, secret, replica = remote_client
replica.connected = False
replica.last_contact = 0
r = await client.get(
"/auth/api/forward",
headers={"Host": "app2.example.com", "Cookie": f"{AUTH_COOKIE_NAME}={secret}"},
)
assert r.status_code == 503
@pytest.mark.asyncio
async def test_logout_proxied_and_evicted(remote_client, monkeypatch):
client, secret, replica = remote_client
async def fake_forward(request):
return Response(status_code=200, content=b'{"message": "Logged out"}')
monkeypatch.setattr(satellite, "forward_request", fake_forward)
r = await client.post(
"/auth/api/logout",
headers={"Host": "app2.example.com", "Cookie": f"{AUTH_COOKIE_NAME}={secret}"},
)
assert r.status_code == 200
assert not replica.db.sessions # evicted optimistically
@pytest.mark.asyncio
async def test_admin_configures_remote_domain(client, session_token, test_db):
"""The admin domains API stores remote config and masks the token."""
r = await client.post(
"/auth/api/admin/domains/",
json={
"rp_id": "example.com",
"rp_name": "Example",
"origins": {
"**.example.com": True,
"auth.example.com": {"auth_host": True},
},
"remote": {"url": "http://remote.test", "token": "sekret", "cache_ttl": 30},
},
headers={
"Host": "localhost:4401",
"Cookie": f"{AUTH_COOKIE_NAME}={session_token}",
},
)
assert r.status_code == 200, r.text
stored = test_db.config.domains["example.com"]
assert stored.remote.url == "http://remote.test"
assert stored.remote.token == "sekret"
r = await client.get(
"/auth/api/admin/domains/",
headers={
"Host": "localhost:4401",
"Cookie": f"{AUTH_COOKIE_NAME}={session_token}",
},
)
entry = next(d for d in r.json() if d["rp_id"] == "example.com")
assert entry["remote"]["url"] == "http://remote.test"
assert "token" not in entry["remote"] # write-only
headers = {"Host": "localhost:4401", "Cookie": f"{AUTH_COOKIE_NAME}={session_token}"}
origins = {"**.example.com": True, "auth.example.com": {"auth_host": True}}
# PATCH without the remote key preserves it (and its token)
r = await client.patch(
"/auth/api/admin/domains/example.com",
json={"rp_name": "Ex", "origins": origins},
headers=headers,
)
assert r.status_code == 200, r.text
assert stored.remote.url == "http://remote.test"
assert stored.remote.token == "sekret"
# PATCH with a new URL but no token keeps the stored token
r = await client.patch(
"/auth/api/admin/domains/example.com",
json={"rp_name": "Ex", "origins": origins,
"remote": {"url": "http://other.test", "cache_ttl": 30}},
headers=headers,
)
assert r.status_code == 200, r.text
assert stored.remote.url == "http://other.test"
assert stored.remote.token == "sekret"
# PATCH with remote: null clears it
r = await client.patch(
"/auth/api/admin/domains/example.com",
json={"rp_name": "Ex", "origins": origins, "remote": None},
headers=headers,
)
assert r.status_code == 200, r.text
assert stored.remote is None
@pytest.mark.asyncio
async def test_admin_remote_domain_requires_auth_host(client, session_token):
r = await client.post(
"/auth/api/admin/domains/",
json={
"rp_id": "example.com",
"origins": {"**.example.com": True},
"remote": {"url": "http://remote.test"},
},
headers={
"Host": "localhost:4401",
"Cookie": f"{AUTH_COOKIE_NAME}={session_token}",
},
)
assert r.status_code == 400
assert "auth host" in r.json()["detail"]