Proxy to another Paskia #5
@@ -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.
|
||||
@@ -480,6 +480,7 @@ function createDomain() {
|
||||
origins: [],
|
||||
originValidation: [],
|
||||
wellKnownCheck: null,
|
||||
remote: null,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -495,6 +496,8 @@ function openDomain(domain) {
|
||||
origins: rows.map(r => r.key),
|
||||
originValidation: rows.map(() => 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()
|
||||
// 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
|
||||
? apiJson('/auth/api/admin/domains/', { method: 'POST', body: { rp_id, rp_name, origins } })
|
||||
: apiJson(`/auth/api/admin/domains/${rp_id}`, { method: 'PATCH', body: { 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, remote } })
|
||||
req
|
||||
.then(() => {
|
||||
authStore.showMessage(`Domain "${rp_id}" ${d.isNew ? 'created' : 'updated'}.`, 'success', 2500)
|
||||
|
||||
@@ -464,6 +464,7 @@ defineExpose({ focusFirstElement })
|
||||
</div>
|
||||
<div class="perm-id-info">
|
||||
<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>
|
||||
</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)
|
||||
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
|
||||
// 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.
|
||||
@@ -31,6 +58,8 @@ const isValidationInvalid = computed(() => {
|
||||
if (relatedEntries.value.length > 5) return true
|
||||
if (d.isNew && !isWellFormedDomain(d.rp_id || '')) 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
|
||||
})
|
||||
|
||||
@@ -515,6 +544,33 @@ function onRemoveOrigin(i) {
|
||||
<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>
|
||||
</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>
|
||||
</template>
|
||||
|
||||
@@ -540,4 +596,7 @@ function onRemoveOrigin(i) {
|
||||
border-color: var(--color-error);
|
||||
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>
|
||||
|
||||
+16
-24
@@ -190,18 +190,6 @@ def cmd_migrate(args: argparse.Namespace) -> None:
|
||||
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:
|
||||
"""Open the combined database and serve all configured domains."""
|
||||
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.")
|
||||
|
||||
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)
|
||||
|
||||
listen = _split_multi(args.listen) or config.listen
|
||||
configure_domains(listen=listen)
|
||||
# Effective serve parameters, teleported to the server process(es); the
|
||||
# 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:
|
||||
registry = build_registry(config)
|
||||
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
|
||||
# is the admin's job via the admin interface) are logged by build().
|
||||
|
||||
# Pass process-global serve parameters to the server process(es)
|
||||
serve_config().listen = listen
|
||||
teleport() # Serialize bound config before spawning workers
|
||||
|
||||
startupbox.print_startup_config(registry, listen=listen, default_port=DEFAULT_PORT)
|
||||
startupbox.print_startup_config(
|
||||
registry, listen=cfg.listen, default_port=DEFAULT_PORT
|
||||
)
|
||||
|
||||
# Run the server (spawns processes in dev mode)
|
||||
# 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.
|
||||
server.run(
|
||||
"paskia.fastapi.mainapp:app",
|
||||
listen=listen,
|
||||
listen=cfg.listen,
|
||||
default_port=DEFAULT_PORT,
|
||||
server_header=False,
|
||||
startup_box=None,
|
||||
|
||||
@@ -24,8 +24,15 @@ EXPIRES = SESSION_LIFETIME
|
||||
|
||||
|
||||
def session_ctx(auth: str, host: str | None = None):
|
||||
"""Get session context with normalized host."""
|
||||
return db.data().session_ctx(auth, hostutil.normalize_host(host))
|
||||
"""Get session context with normalized 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:
|
||||
|
||||
+4
-1
@@ -67,7 +67,10 @@ async def check_admin_credentials() -> bool:
|
||||
# Check first admin user for credentials on any configured domain
|
||||
admin_user = admin_users[0]
|
||||
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):
|
||||
# Admin exists but has no credential on any domain
|
||||
|
||||
@@ -67,6 +67,7 @@ from paskia.db.structs import (
|
||||
DomainConfig,
|
||||
Org,
|
||||
Permission,
|
||||
RemoteConfig,
|
||||
ResetToken,
|
||||
Role,
|
||||
Session,
|
||||
@@ -90,6 +91,7 @@ __all__ = [
|
||||
"Org",
|
||||
"Permission",
|
||||
"DomainConfig",
|
||||
"RemoteConfig",
|
||||
"ResetToken",
|
||||
"Role",
|
||||
"Session",
|
||||
|
||||
+27
-2
@@ -13,7 +13,7 @@ from uuid import UUID
|
||||
|
||||
import uuid7
|
||||
|
||||
from paskia import oidc_notify
|
||||
from paskia import oidc_notify, syncfeed
|
||||
from paskia.config import SESSION_LIFETIME
|
||||
from paskia.db.structs import (
|
||||
DB,
|
||||
@@ -23,6 +23,7 @@ from paskia.db.structs import (
|
||||
Org,
|
||||
OriginEntry,
|
||||
Permission,
|
||||
RemoteConfig,
|
||||
ResetToken,
|
||||
Role,
|
||||
Session,
|
||||
@@ -103,6 +104,7 @@ def update_permission(
|
||||
_db.permissions[uuid].scope = scope
|
||||
_db.permissions[uuid].display_name = display_name
|
||||
_db.permissions[uuid].domain = domain
|
||||
syncfeed.emit("permissions", str(uuid), _db.permissions[uuid])
|
||||
|
||||
|
||||
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")
|
||||
with _transaction("admin:update_org_name", ctx):
|
||||
_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:
|
||||
@@ -180,6 +183,9 @@ def add_permission_to_org(
|
||||
|
||||
with _transaction("admin:add_permission_to_org", ctx):
|
||||
_db.permissions[permission_uuid].orgs[org_uuid] = True
|
||||
syncfeed.emit(
|
||||
"permissions", str(permission_uuid), _db.permissions[permission_uuid]
|
||||
)
|
||||
|
||||
|
||||
def remove_permission_from_org(
|
||||
@@ -197,6 +203,9 @@ def remove_permission_from_org(
|
||||
|
||||
with _transaction("admin:remove_permission_from_org", ctx):
|
||||
_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:
|
||||
@@ -220,6 +229,7 @@ def update_role_name(
|
||||
raise ValueError(f"Role {uuid} not found")
|
||||
with _transaction("admin:update_role_name", ctx):
|
||||
_db.roles[uuid].display_name = display_name
|
||||
syncfeed.emit("roles", str(uuid), _db.roles[uuid])
|
||||
|
||||
|
||||
def add_permission_to_role(
|
||||
@@ -235,6 +245,7 @@ def add_permission_to_role(
|
||||
raise ValueError(f"Permission {permission_uuid} not found")
|
||||
with _transaction("admin:add_permission_to_role", ctx):
|
||||
_db.roles[role_uuid].permissions[permission_uuid] = True
|
||||
syncfeed.emit("roles", str(role_uuid), _db.roles[role_uuid])
|
||||
|
||||
|
||||
def remove_permission_from_role(
|
||||
@@ -248,6 +259,7 @@ def remove_permission_from_role(
|
||||
raise ValueError(f"Role {role_uuid} not found")
|
||||
with _transaction("admin:remove_permission_from_role", ctx):
|
||||
_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:
|
||||
@@ -302,6 +314,7 @@ def update_user_display_name(
|
||||
slug = slugify_name(display_name)
|
||||
if slug and not is_username_taken(slug, exclude_uuid=uuid):
|
||||
user.preferred_username = slug
|
||||
syncfeed.emit("users", str(uuid), user)
|
||||
|
||||
|
||||
def update_user_info(
|
||||
@@ -380,6 +393,7 @@ def update_user_info(
|
||||
user.preferred_username = preferred_username
|
||||
if telephone is not _UNSET:
|
||||
user.telephone = telephone
|
||||
syncfeed.emit("users", str(uuid), user)
|
||||
|
||||
|
||||
def update_user_role(
|
||||
@@ -395,6 +409,7 @@ def update_user_role(
|
||||
raise ValueError(f"Role {role_uuid} not found")
|
||||
with _transaction("admin:update_user_role", ctx):
|
||||
_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:
|
||||
@@ -429,6 +444,7 @@ def update_credential_sign_count(
|
||||
_db.credentials[uuid].sign_count = sign_count
|
||||
if last_used:
|
||||
_db.credentials[uuid].last_used = last_used
|
||||
syncfeed.emit("credentials", str(uuid), _db.credentials[uuid])
|
||||
|
||||
|
||||
def delete_credential(
|
||||
@@ -476,6 +492,7 @@ def update_session(
|
||||
s.validated = validated
|
||||
if issuer is not None:
|
||||
s.issuer = issuer
|
||||
syncfeed.emit("sessions", key, s)
|
||||
|
||||
|
||||
def delete_session(
|
||||
@@ -598,6 +615,9 @@ def login(
|
||||
# Update credential
|
||||
_db.credentials[credential_uuid].sign_count = sign_count
|
||||
_db.credentials[credential_uuid].last_used = now
|
||||
syncfeed.emit(
|
||||
"credentials", str(credential_uuid), _db.credentials[credential_uuid]
|
||||
)
|
||||
return token
|
||||
|
||||
|
||||
@@ -625,6 +645,9 @@ def oidc_login(
|
||||
# Update credential
|
||||
_db.credentials[credential_uuid].sign_count = sign_count
|
||||
_db.credentials[credential_uuid].last_used = now
|
||||
syncfeed.emit(
|
||||
"credentials", str(credential_uuid), _db.credentials[credential_uuid]
|
||||
)
|
||||
|
||||
|
||||
def create_credential_session(
|
||||
@@ -714,9 +737,10 @@ def update_domain(
|
||||
*,
|
||||
rp_name: str | None,
|
||||
origins: dict[str, bool | OriginEntry],
|
||||
remote: RemoteConfig | None = None,
|
||||
ctx: SessionContext | 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
|
||||
changing it would orphan them — delete and recreate the domain instead.
|
||||
@@ -728,6 +752,7 @@ def update_domain(
|
||||
with _transaction("admin:update_domain", ctx):
|
||||
domain.rp_name = rp_name
|
||||
domain.origins = origins
|
||||
domain.remote = remote
|
||||
|
||||
|
||||
def delete_domain(rp_id: str, *, ctx: SessionContext | None = None) -> None:
|
||||
|
||||
+48
-8
@@ -9,7 +9,7 @@ from uuid import UUID
|
||||
import msgspec
|
||||
import uuid7
|
||||
|
||||
from paskia import db
|
||||
from paskia import db, syncfeed
|
||||
from paskia.util import passphrase as passphrase_util
|
||||
from paskia.util.crypto import hash_secret
|
||||
|
||||
@@ -51,6 +51,7 @@ class Permission(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
def store(self) -> None:
|
||||
"""Store this permission in the database. Must be called inside a transaction."""
|
||||
db.data().permissions[self.uuid] = self
|
||||
syncfeed.emit("permissions", str(self.uuid), self)
|
||||
|
||||
def delete(self) -> None:
|
||||
"""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()
|
||||
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]
|
||||
syncfeed.emit("permissions", str(self.uuid), None)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -103,6 +106,7 @@ class Org(msgspec.Struct, dict=True):
|
||||
def store(self) -> None:
|
||||
"""Store this organization in the database. Must be called inside a transaction."""
|
||||
db.data().orgs[self.uuid] = self
|
||||
syncfeed.emit("orgs", str(self.uuid), self)
|
||||
|
||||
def delete(self) -> None:
|
||||
"""Delete this org and cascade to roles, users. Remove from permissions.
|
||||
@@ -111,12 +115,16 @@ class Org(msgspec.Struct, dict=True):
|
||||
"""
|
||||
_data = db.data()
|
||||
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 user in role.users:
|
||||
del _data.users[user.uuid]
|
||||
syncfeed.emit("users", str(user.uuid), None)
|
||||
del _data.roles[role.uuid]
|
||||
syncfeed.emit("roles", str(role.uuid), None)
|
||||
del _data.orgs[self.uuid]
|
||||
syncfeed.emit("orgs", str(self.uuid), None)
|
||||
|
||||
@classmethod
|
||||
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:
|
||||
"""Store this role in the database. Must be called inside a transaction."""
|
||||
db.data().roles[self.uuid] = self
|
||||
syncfeed.emit("roles", str(self.uuid), self)
|
||||
|
||||
def delete(self) -> None:
|
||||
"""Delete this role from the database. Must be called inside a transaction."""
|
||||
del db.data().roles[self.uuid]
|
||||
syncfeed.emit("roles", str(self.uuid), None)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -254,6 +264,7 @@ class User(msgspec.Struct, dict=True, omit_defaults=True, kw_only=True):
|
||||
def store(self) -> None:
|
||||
"""Store this user in the database. Must be called inside a transaction."""
|
||||
db.data().users[self.uuid] = self
|
||||
syncfeed.emit("users", str(self.uuid), self)
|
||||
|
||||
def delete(self) -> None:
|
||||
"""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()
|
||||
for cred in self.credentials:
|
||||
del _data.credentials[cred.uuid]
|
||||
syncfeed.emit("credentials", str(cred.uuid), None)
|
||||
for sess in self.sessions:
|
||||
del _data.sessions[sess.key]
|
||||
syncfeed.emit("sessions", sess.key, None)
|
||||
for token in self.reset_tokens:
|
||||
del _data.reset_tokens[token.key]
|
||||
del _data.users[self.uuid]
|
||||
syncfeed.emit("users", str(self.uuid), None)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -331,6 +345,7 @@ class Credential(msgspec.Struct, dict=True):
|
||||
def store(self) -> None:
|
||||
"""Store this credential in the database. Must be called inside a transaction."""
|
||||
db.data().credentials[self.uuid] = self
|
||||
syncfeed.emit("credentials", str(self.uuid), self)
|
||||
|
||||
def delete(self) -> None:
|
||||
"""Delete this credential and all its sessions.
|
||||
@@ -340,7 +355,9 @@ class Credential(msgspec.Struct, dict=True):
|
||||
_data = db.data()
|
||||
for sess in self.sessions:
|
||||
del _data.sessions[sess.key]
|
||||
syncfeed.emit("sessions", sess.key, None)
|
||||
del _data.credentials[self.uuid]
|
||||
syncfeed.emit("credentials", str(self.uuid), None)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -418,10 +435,13 @@ class Session(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
_data.sessions[self.key] = self
|
||||
_data.users[self.user_uuid].last_seen = last_seen
|
||||
_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:
|
||||
"""Delete this session from the database. Must be called inside a transaction."""
|
||||
del db.data().sessions[self.key]
|
||||
syncfeed.emit("sessions", self.key, None)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -622,6 +642,21 @@ class OriginEntry(msgspec.Struct, omit_defaults=True):
|
||||
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):
|
||||
"""Configuration for one domain (one WebAuthn rp-id).
|
||||
|
||||
@@ -641,6 +676,7 @@ class DomainConfig(msgspec.Struct, omit_defaults=True):
|
||||
|
||||
rp_name: str | None = None
|
||||
origins: dict[str, bool | OriginEntry] = {}
|
||||
remote: RemoteConfig | None = None
|
||||
|
||||
|
||||
class Config(msgspec.Struct, omit_defaults=True):
|
||||
@@ -728,17 +764,21 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
if s.host != host:
|
||||
return None
|
||||
|
||||
# Look up via this instance's own tables: a DB must be
|
||||
# self-contained so that read replicas work unchanged.
|
||||
try:
|
||||
user = s.user
|
||||
role = user.role
|
||||
org = role.org
|
||||
credential = s.credential
|
||||
user = self.users[s.user_uuid]
|
||||
role = self.roles[user.role_uuid]
|
||||
org = self.orgs[role.org_uuid]
|
||||
credential = self.credentials[s.credential_uuid]
|
||||
except KeyError:
|
||||
return None
|
||||
|
||||
# Effective permissions: role's permissions that the org can grant,
|
||||
# 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 = []
|
||||
for perm_uuid in role.permission_set:
|
||||
|
||||
+50
-2
@@ -14,13 +14,14 @@ domain are in-domain, entries outside it are related.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import logging
|
||||
import os
|
||||
|
||||
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.util import hostutil
|
||||
from paskia.util.constants import DEFAULT_PORT
|
||||
@@ -101,6 +102,11 @@ class Domain:
|
||||
def rp_name(self) -> str:
|
||||
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
|
||||
def own_auth_host(self) -> str | None:
|
||||
"""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():
|
||||
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
|
||||
related_count = 0
|
||||
@@ -273,6 +283,11 @@ def validate_config(
|
||||
f"Domain '{rp_id}' has {related_count} related origins "
|
||||
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)
|
||||
for hn, owner in auth_hosts.items():
|
||||
@@ -363,6 +378,20 @@ def sanitize_config(
|
||||
auth_seen = True
|
||||
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))
|
||||
if len(related) > related_origin_cap:
|
||||
warn(
|
||||
@@ -372,7 +401,9 @@ def sanitize_config(
|
||||
for key in related[related_origin_cap:]:
|
||||
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:
|
||||
raise ValueError("No servable domain in the stored configuration")
|
||||
@@ -461,6 +492,17 @@ def _derive_site(
|
||||
|
||||
_registry: DomainRegistry | 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:
|
||||
@@ -501,6 +543,12 @@ def init_registry(config: Config) -> DomainRegistry:
|
||||
"""Build and install the global registry from a combined configuration."""
|
||||
global _registry
|
||||
_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
|
||||
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ immediately.
|
||||
from fastapi import Body, FastAPI, Request
|
||||
|
||||
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.admin.errors import install_error_handlers
|
||||
from paskia.fastapi.response import MsgspecResponse
|
||||
@@ -27,6 +27,14 @@ install_error_handlers(app)
|
||||
|
||||
|
||||
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(
|
||||
rp_id=domain.rp_id,
|
||||
rp_name=domain.rp_name,
|
||||
@@ -34,6 +42,28 @@ def _domain_to_api(domain: domains.Domain) -> ApiDomain:
|
||||
site_url=domain.site_url,
|
||||
auth_site_url=domain.auth_site_url,
|
||||
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(
|
||||
rp_name=(payload.get("rp_name") or "").strip() or None,
|
||||
origins=_normalize_origins_map(payload.get("origins")),
|
||||
remote=_normalize_remote(payload.get("remote")),
|
||||
)
|
||||
|
||||
config = db.data().config
|
||||
@@ -156,9 +187,15 @@ async def admin_update_domain(
|
||||
if rp_id not in config.domains:
|
||||
raise ValueError(f"Domain {rp_id} not found")
|
||||
|
||||
current_remote = config.domains[rp_id].remote
|
||||
updated = DomainConfig(
|
||||
rp_name=(payload.get("rp_name") or "").strip() or None,
|
||||
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(
|
||||
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_name=updated.rp_name,
|
||||
origins=updated.origins,
|
||||
remote=updated.remote,
|
||||
ctx=ctx,
|
||||
)
|
||||
_rebuild_registry()
|
||||
|
||||
+21
-9
@@ -14,7 +14,7 @@ from fastapi import (
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.security import HTTPBearer
|
||||
|
||||
from paskia import authcode, db
|
||||
from paskia import authcode, db, satellite
|
||||
from paskia._version import __version__
|
||||
from paskia.authsession import EXPIRES, get_reset, session_ctx
|
||||
from paskia.domains import current_domain
|
||||
@@ -122,8 +122,9 @@ async def validate_token(
|
||||
if auth and renew:
|
||||
consumed = datetime.now(UTC) - ctx.session.validated
|
||||
if not timedelta(0) < consumed < _REFRESH_INTERVAL:
|
||||
db.update_session(
|
||||
satellite.refresh_session(
|
||||
ctx.session.key,
|
||||
request.headers.get("host"),
|
||||
ip=get_client_ip(request),
|
||||
user_agent=request.headers.get("user-agent"),
|
||||
validated=datetime.now(UTC),
|
||||
@@ -162,16 +163,16 @@ async def check_user(
|
||||
|
||||
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:
|
||||
u = data.users[user_uuid]
|
||||
role = u.role
|
||||
org = role.org
|
||||
role = data.roles[u.role_uuid]
|
||||
org = data.orgs[role.org_uuid]
|
||||
except KeyError:
|
||||
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 org.permissions}
|
||||
org_perm_uuids = {p.uuid for p in data.permissions.values() if org.uuid in p.orgs}
|
||||
|
||||
effective_perms = []
|
||||
for perm_uuid in role.permission_set:
|
||||
@@ -212,7 +213,7 @@ def _remote_headers(ctx) -> dict[str, str]:
|
||||
"Remote-Session-Expires": (
|
||||
(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")
|
||||
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."""
|
||||
if (proxied := await satellite.forward_request(request)) is not None:
|
||||
return proxied
|
||||
if not credentials or not credentials.credentials:
|
||||
raise HTTPException(401, "Bearer token required")
|
||||
token = credentials.credentials
|
||||
@@ -375,6 +378,10 @@ async def token_info(credentials=Depends(bearer_auth)):
|
||||
|
||||
@app.post("/logout")
|
||||
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:
|
||||
return {"message": "Already logged out"}
|
||||
host = request.headers.get("host")
|
||||
@@ -399,6 +406,11 @@ async def api_set_session(
|
||||
if not auth or not auth.credentials:
|
||||
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", ""))
|
||||
if not host:
|
||||
raise HTTPException(400, "Host header required")
|
||||
|
||||
@@ -63,6 +63,11 @@ class DispatchMiddleware:
|
||||
host = _header(scope, "host")
|
||||
host_domain = registry.resolve(host)
|
||||
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})
|
||||
return
|
||||
|
||||
|
||||
@@ -7,11 +7,11 @@ from fastapi import FastAPI, HTTPException, Request, Response
|
||||
from fastapi.responses import FileResponse, RedirectResponse
|
||||
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.db.background import start_background, stop_background
|
||||
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.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
|
||||
"""Application lifespan: open the combined database and build the domain registry.
|
||||
|
||||
Process-global serve parameters (listen endpoints) are passed via the
|
||||
PASKIA_CONFIG JSON env variable (set by the CLI entrypoint) so that
|
||||
uvicorn reload / multiprocess workers derive site URLs the same way.
|
||||
Domain configuration is read from the database.
|
||||
Process-global serve parameters (listen endpoints, save flag) are passed
|
||||
via the PASKIA_CONFIG JSON env variable (set by the CLI entrypoint) so
|
||||
that uvicorn reload / multiprocess workers derive site URLs the same
|
||||
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()
|
||||
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
|
||||
)
|
||||
async with kanta:
|
||||
if cfg.save:
|
||||
with kanta.transaction("serve:save_listen"):
|
||||
db.data().config.listen = cfg.listen
|
||||
try:
|
||||
domains.init_registry(db.data().config)
|
||||
await remoteauth.init()
|
||||
await authcode.start()
|
||||
await satellite.manager.start()
|
||||
except ValueError as e:
|
||||
logging.error(f"⚠️ {e}")
|
||||
# Re-raise to fail fast
|
||||
@@ -55,6 +61,7 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
||||
await start_background()
|
||||
yield
|
||||
await stop_background()
|
||||
await satellite.manager.stop()
|
||||
await authcode.stop()
|
||||
|
||||
|
||||
@@ -78,6 +85,7 @@ app.middleware("http")(auth_host.redirect_middleware)
|
||||
app.add_middleware(DispatchMiddleware)
|
||||
|
||||
app.mount("/auth/api/admin/", admin.app)
|
||||
app.mount("/auth/api/sync", sync.app)
|
||||
app.mount("/auth/api/", api.app)
|
||||
app.mount("/auth/ws/", ws.app)
|
||||
app.mount("/auth/oidc/", oid.app)
|
||||
|
||||
@@ -20,7 +20,7 @@ from fastapi import Depends, FastAPI, Form, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
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.util import avatar, oidjwt
|
||||
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.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:
|
||||
"""Return the instance-global OIDC provider state."""
|
||||
return db.data().oidc
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -14,7 +14,15 @@ import msgspec
|
||||
from uarite import uaparse
|
||||
|
||||
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
|
||||
@@ -194,6 +202,7 @@ class ApiDomain(msgspec.Struct):
|
||||
site_url: str
|
||||
auth_site_url: str
|
||||
auth_host: str | None
|
||||
remote: RemoteConfig | None = None
|
||||
|
||||
|
||||
class ApiTokenInfo(msgspec.Struct, omit_defaults=True):
|
||||
|
||||
@@ -2,9 +2,10 @@
|
||||
|
||||
Domain configuration lives in the database (``Config.domains``); the
|
||||
``PASKIA_CONFIG`` environment variable only carries the effective listen
|
||||
endpoints so that child processes (uvicorn reload / workers) derive site
|
||||
URLs the same way the parent did. The CLI entry point mutates the bound
|
||||
object before ``server.run()`` calls ``teleport()`` to pass it on.
|
||||
endpoints and whether to persist them, so that child processes (uvicorn
|
||||
reload / workers) derive site URLs the same way the parent did. The CLI
|
||||
entry point mutates the bound object before ``server.run()`` calls
|
||||
``teleport()`` to pass it on.
|
||||
"""
|
||||
|
||||
import msgspec
|
||||
@@ -15,6 +16,7 @@ class ServeConfig(msgspec.Struct):
|
||||
"""Process-global serve parameters."""
|
||||
|
||||
listen: list[str] | None = None
|
||||
save: bool = False # Persist listen to the stored config on startup
|
||||
|
||||
|
||||
def serve_config() -> ServeConfig:
|
||||
|
||||
+11
-5
@@ -1,6 +1,6 @@
|
||||
"""User information formatting and retrieval logic."""
|
||||
|
||||
from paskia import aaguid, db
|
||||
from paskia import aaguid, satellite
|
||||
from paskia.db import SessionContext
|
||||
from paskia.util import avatar, hostutil
|
||||
from paskia.util.apistructs import (
|
||||
@@ -43,24 +43,30 @@ async def build_user_info(
|
||||
ctx: SessionContext | None = None,
|
||||
) -> ApiUserDetail:
|
||||
"""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)
|
||||
|
||||
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 = {
|
||||
s.key: ApiUserSession.from_db(
|
||||
s,
|
||||
current_key=session_key,
|
||||
normalized_host=normalized_host,
|
||||
)
|
||||
for s in user.sessions
|
||||
for s in user_sessions
|
||||
}
|
||||
|
||||
return ApiUserDetail(
|
||||
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={
|
||||
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,
|
||||
permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions}
|
||||
|
||||
+1
-1
@@ -21,7 +21,7 @@ dependencies = [
|
||||
"pyjwt[crypto]>=2.11.0",
|
||||
"jsondiff>=2.2.1",
|
||||
"msgspec>=0.20.0",
|
||||
"fastapi-vue~=1.7.1",
|
||||
"fastapi-vue~=1.7.2",
|
||||
"kanta>=0.9.2",
|
||||
"uarite>=0.2.1",
|
||||
]
|
||||
|
||||
@@ -10,23 +10,20 @@ from pathlib import Path
|
||||
|
||||
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):
|
||||
"""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:
|
||||
emoji = _LEVEL_EMOJI.get(record.levelno)
|
||||
prefix = f"{emoji} " if emoji else f"{record.levelname}: "
|
||||
return prefix + record.getMessage()
|
||||
if record.levelno >= logging.ERROR:
|
||||
return f"🛑 {record.getMessage()}"
|
||||
if record.levelno >= logging.WARNING:
|
||||
return f"💣 {record.getMessage()}"
|
||||
return record.getMessage()
|
||||
|
||||
|
||||
_handler = logging.StreamHandler()
|
||||
|
||||
@@ -139,7 +139,7 @@ async def ready(url: str, path: str = "", max_attempts: int = 50) -> None:
|
||||
|
||||
for attempt in range(max_attempts):
|
||||
if await http_get_server(f"{url}{path}", timeout=1.0) is not None:
|
||||
logger.info("✓ Backend ready!")
|
||||
logger.info("🟢 Backend ready!")
|
||||
return
|
||||
if attempt == max_attempts - 1:
|
||||
logger.error("Backend at %s didn't start in time", url)
|
||||
|
||||
+13
-3
@@ -172,6 +172,7 @@ def test_serve_uses_stored_config(run_cli, tmp_path):
|
||||
assert calls["listen"] is None # stored listen (None) used
|
||||
serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig)
|
||||
assert serve.listen is None
|
||||
assert serve.save is False
|
||||
|
||||
|
||||
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"]
|
||||
serve = msgspec.json.decode(os.environ["PASKIA_CONFIG"].encode(), type=ServeConfig)
|
||||
assert serve.listen == ["4403"]
|
||||
assert serve.save is False
|
||||
# Stored config keeps the original listen value
|
||||
assert stored_config(tmp_path).listen == ["4402"]
|
||||
|
||||
|
||||
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")
|
||||
calls = run_cli("--listen", "4403", "--save")
|
||||
|
||||
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):
|
||||
"""--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("--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):
|
||||
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user