diff --git a/paskia/authsession.py b/paskia/authsession.py index b5d7d68..34a1275 100644 --- a/paskia/authsession.py +++ b/paskia/authsession.py @@ -18,18 +18,21 @@ from paskia.db.structs import ResetToken from paskia.util import hostutil if TYPE_CHECKING: - from paskia.db import DB, ResetToken + from paskia.db import ResetToken EXPIRES = SESSION_LIFETIME -def session_ctx(auth: str, host: str | None = None, store: DB | None = None): +def session_ctx(auth: str, host: str | None = None): """Get session context with normalized host. - store defaults to the local database; remote-domain request paths pass - their domain's replica explicitly. + The store is dispatched by host: remote domains read their replica. """ - return (store or db.data()).session_ctx(auth, hostutil.normalize_host(host)) + from paskia import satellite # noqa: PLC0415 (import cycle) + + return satellite.store_for_host(host).session_ctx( + auth, hostutil.normalize_host(host) + ) def expires() -> datetime: diff --git a/paskia/domains.py b/paskia/domains.py index 140ec5c..3a56b6d 100644 --- a/paskia/domains.py +++ b/paskia/domains.py @@ -21,7 +21,6 @@ import os from fastapi_vue.hostutil import parse_endpoints -from paskia.db import operations from paskia.db.structs import Config, DomainConfig, OriginEntry, RemoteConfig from paskia.sansio import Passkey from paskia.util import hostutil @@ -92,7 +91,6 @@ class Domain: self.config = config self.site_url = site_url self.site_path = site_path - self._store = None self.passkey = Passkey( rp_id=rp_id, rp_name=config.rp_name, @@ -100,21 +98,6 @@ class Domain: related_origins=[origin_url(k) for k in related], ) - @property - def store(self): - """The DB instance this domain's request paths read. - - Defaults to the local database; remote domains get their read - replica attached by the sync client at startup. - """ - if self._store is not None: - return self._store - return operations._db - - @store.setter - def store(self, value) -> None: - self._store = value - @property def rp_name(self) -> str: return self.passkey.rp_name diff --git a/paskia/fastapi/api.py b/paskia/fastapi/api.py index 5091144..bee3b3c 100644 --- a/paskia/fastapi/api.py +++ b/paskia/fastapi/api.py @@ -18,7 +18,7 @@ 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 -from paskia.fastapi import authz, proxy, session, user +from paskia.fastapi import authz, session, user from paskia.fastapi.response import MsgspecResponse from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME, get_client_ip from paskia.util import hostutil, htmlutil, passphrase, permutil, userinfo @@ -97,20 +97,6 @@ def _parse_perm(perm: list[str]) -> list[tuple[str, ...]]: raise HTTPException(status_code=400, detail=str(e)) -def _store(request: Request): - """The dispatched domain's data store (local DB or read replica). - - Remote domains fail closed once the sync channel has been silent for - longer than their cache_ttl. - """ - domain = request.state.domain - if domain.remote is not None: - replica = satellite.manager.replica_for(domain) - if replica is None or not replica.available(): - raise HTTPException(503, "Remote authentication service unavailable") - return domain.store - - @app.post("/validate") async def validate_token( request: Request, @@ -128,7 +114,6 @@ async def validate_token( perm_groups, host=request.headers.get("host"), max_age=max_age, - store=_store(request), ) except HTTPException: # Global handler will clear cookie if 401 @@ -137,22 +122,14 @@ async def validate_token( if auth and renew: consumed = datetime.now(UTC) - ctx.session.validated if not timedelta(0) < consumed < _REFRESH_INTERVAL: - replica = satellite.manager.replica_for(request.state.domain) - if replica is not None: - replica.refresh_session( - ctx.session.key, - datetime.now(UTC), - get_client_ip(request), - request.headers.get("user-agent", ""), - ) - else: - db.update_session( - ctx.session.key, - ip=get_client_ip(request), - user_agent=request.headers.get("user-agent"), - validated=datetime.now(UTC), - ctx=ctx, - ) + 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), + ctx=ctx, + ) renewed = True _set_log_extra(request, ctx.session.key) resp = MsgspecResponse( @@ -186,7 +163,8 @@ async def check_user( No session cookie is read or written. Caller authentication is not required. """ - data = _store(request) + host = hostutil.normalize_host(request.headers.get("host")) + data = satellite.store_for_host(host) try: u = data.users[user_uuid] role = data.roles[u.role_uuid] @@ -194,7 +172,6 @@ async def check_user( 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 data.permissions.values() if org.uuid in p.orgs} effective_perms = [] @@ -291,7 +268,6 @@ async def forward_authentication( perm_groups, host=request.headers.get("host"), max_age=max_age, - store=_store(request), ) _set_log_extra(request, forwarded, ctx.session.key) remote_headers = _remote_headers(ctx) @@ -354,7 +330,7 @@ async def api_user_info( detail="Authentication required", mode="login", ) - ctx = session_ctx(auth, request.headers.get("host"), store=_store(request)) + ctx = session_ctx(auth, request.headers.get("host")) if not ctx: raise authz.AuthException( status_code=401, @@ -371,7 +347,6 @@ async def api_user_info( session_key=ctx.session.key, request_host=request.headers.get("host"), ctx=ctx, - store=_store(request), ) ) @@ -379,8 +354,8 @@ async def api_user_info( @app.get("/token-info") async def token_info(request: Request, credentials=Depends(bearer_auth)): """Get reset/device-add token info. Pass token via Bearer header.""" - if request.state.domain.remote is not None: - return await proxy.proxy_to_remote(request, request.state.domain.remote) + 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 @@ -403,12 +378,9 @@ async def token_info(request: Request, credentials=Depends(bearer_auth)): @app.post("/logout") async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE): - if request.state.domain.remote is not None: - proxied = await proxy.proxy_to_remote(request, request.state.domain.remote) + if (proxied := await satellite.forward_request(request)) is not None: if auth and proxied.status_code == 200: - replica = satellite.manager.replica_for(request.state.domain) - if replica is not None: - replica.evict_session(auth) + satellite.evict_session(auth, request.headers.get("host")) return proxied if not auth: return {"message": "Already logged out"} @@ -434,10 +406,10 @@ async def api_set_session( if not auth or not auth.credentials: raise HTTPException(400, "Bearer token required") - if request.state.domain.remote is not None: + 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 await proxy.proxy_to_remote(request, request.state.domain.remote) + return proxied host = hostutil.normalize_host(request.headers.get("host", "")) if not host: diff --git a/paskia/fastapi/authz.py b/paskia/fastapi/authz.py index 6815afd..71cec70 100644 --- a/paskia/fastapi/authz.py +++ b/paskia/fastapi/authz.py @@ -62,7 +62,6 @@ async def verify( match: Callable | None = None, host: str | None = None, max_age: str | None = None, - store=None, ): """Validate session token and optional list of required permissions. @@ -70,9 +69,6 @@ async def verify( scope patterns (OR semantics within a group). All entries must be satisfied (AND semantics). - store defaults to the local database; remote-domain request paths pass - their domain's replica explicitly. - Returns the session context. Raises AuthException on failure with metadata for UI rendering. @@ -84,7 +80,7 @@ async def verify( mode="login", ) - ctx = await permutil.session_context(auth, host, store=store) + ctx = await permutil.session_context(auth, host) if not ctx: raise AuthException( status_code=401, diff --git a/paskia/fastapi/oid.py b/paskia/fastapi/oid.py index 1936865..934975c 100644 --- a/paskia/fastapi/oid.py +++ b/paskia/fastapi/oid.py @@ -20,9 +20,8 @@ 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.fastapi import proxy from paskia.util import avatar, oidjwt from paskia.util.crypto import hash_secret @@ -34,9 +33,8 @@ 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.""" - remote = request.state.domain.remote - if remote is not None: - return await proxy.proxy_to_remote(request, remote) + if (proxied := await satellite.forward_request(request)) is not None: + return proxied return await call_next(request) diff --git a/paskia/fastapi/proxy.py b/paskia/fastapi/proxy.py deleted file mode 100644 index 5a51c2b..0000000 --- a/paskia/fastapi/proxy.py +++ /dev/null @@ -1,58 +0,0 @@ -"""HTTP forwarding for remote domains: mutations proxied to the remote. - -The original Host header is preserved so the remote dispatches the request -to the same domain (sessions are host-bound). User cookies authenticate the -forwarded call; no satellite credentials are involved. -""" - -import httpx -from fastapi import Request, Response - -from paskia.db.structs import RemoteConfig - -_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", -} - -_SKIP_RESPONSE_HEADERS = _HOP_BY_HOP | {"content-encoding"} - -_clients: dict[str, httpx.AsyncClient] = {} - - -def _client(base_url: str) -> httpx.AsyncClient: - client = _clients.get(base_url) - if client is None: - client = httpx.AsyncClient(base_url=base_url, timeout=_TIMEOUT) - _clients[base_url] = client - return client - - -async def proxy_to_remote(request: Request, remote: RemoteConfig) -> Response: - """Forward this request to the remote instance unchanged.""" - headers = {k: v for k, v in request.headers.items() if k.lower() not in _HOP_BY_HOP} - upstream = await _client(remote.url).request( - request.method, - request.url.path, - params=request.url.query, - content=await request.body(), - headers=headers, - ) - 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 _SKIP_RESPONSE_HEADERS - ] - return response diff --git a/paskia/fastapi/sync.py b/paskia/fastapi/sync.py index 5cd502f..d05a654 100644 --- a/paskia/fastapi/sync.py +++ b/paskia/fastapi/sync.py @@ -1,7 +1,9 @@ """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. +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 @@ -17,34 +19,9 @@ _logger = logging.getLogger(__name__) app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None) -_SNAPSHOT_TABLES = ( - ("permissions", "permissions"), - ("orgs", "orgs"), - ("roles", "roles"), - ("users", "users"), - ("credentials", "credentials"), - ("sessions", "sessions"), -) - -def _snapshot_messages() -> list[bytes]: - data = db.data() - messages = [] - for table, attr in _SNAPSHOT_TABLES: - items = [ - [str(key), msgspec.to_builtins(obj)] - for key, obj in getattr(data, attr).items() - ] - messages.append( - syncfeed.encode({"type": "snapshot", "table": table, "items": items}) - ) - return messages - - -async def _send(ws: WebSocket, message: dict | bytes) -> None: - await ws.send_bytes( - message if isinstance(message, bytes) else syncfeed.encode(message) - ) +async def _send(ws: WebSocket, message: dict) -> None: + await ws.send_bytes(syncfeed.encode(message)) async def _apply_client_message(message: dict) -> None: @@ -61,8 +38,8 @@ async def _apply_client_message(message: dict) -> None: return db.update_session( key, - ip=str(message.get("ip") or session.ip), - user_agent=str(message.get("user_agent") or session.user_agent), + ip=message.get("ip") or None, + user_agent=message.get("user_agent") or None, validated=validated, ) @@ -71,49 +48,42 @@ async def _apply_client_message(message: dict) -> None: async def sync_websocket(ws: WebSocket): tokens = syncfeed.tokens_from_env() auth = ws.headers.get("authorization", "") - token = auth.removeprefix("Bearer ").strip() - if not tokens or token not in tokens: + if not tokens or auth.removeprefix("Bearer ").strip() not in tokens: await ws.close(code=1008) return await ws.accept() - feed = syncfeed.feed - await _send(ws, {"type": "hello", "generation": feed.generation, "seq": feed.seq}) - + queue = syncfeed.subscribe() try: - # The client always speaks first: resume request (possibly null fields) - resume = msgspec.json.decode(await ws.receive_bytes()) - queue = feed.subscribe() - try: - replay = None - if ( - resume.get("type") == "resume" - and resume.get("generation") == feed.generation - and isinstance(resume.get("seq"), int) - ): - replay = feed.replay_since(resume["seq"]) - if replay is not None: - for event in replay: - await _send(ws, event) - else: - for chunk in _snapshot_messages(): - await _send(ws, chunk) - await _send(ws, {"type": "ready", "seq": feed.seq}) + 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() + sender = asyncio.create_task(_pump(ws, queue)) + try: + while True: + await _apply_client_message( + msgspec.json.decode(await ws.receive_bytes()) + ) finally: - feed.unsubscribe(queue) + 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: diff --git a/paskia/satellite.py b/paskia/satellite.py index 4309cd8..ba5ecaf 100644 --- a/paskia/satellite.py +++ b/paskia/satellite.py @@ -1,14 +1,17 @@ -"""Satellite side of remote domains: RAM-only read replicas. +"""Satellite side of remote domains: RAM-only replicas + host dispatch. -For each domain configured with ``DomainConfig.remote`` a replica of the -remote's tables (another plain DB instance, never persisted) is attached to -the runtime Domain as its store, fed by a sync WebSocket to the remote and -refreshed by periodic full snapshots. Session refreshes from /validate are -written back over the same channel. +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. -While the sync channel has been silent for longer than the domain's -cache_ttl the replica is considered unavailable (fail-closed; set a large -cache_ttl for fail-open behavior bounded by session expiry). +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 @@ -18,11 +21,13 @@ 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 domains -from paskia.authsession import EXPIRES +from paskia import db, domains +from paskia.config import SESSION_LIFETIME from paskia.db.structs import ( DB, Credential, @@ -50,30 +55,13 @@ _RECONNECT_DELAY = 5 _SWEEP_INTERVAL = 60 -def _apply(replica: DB, table: str, key: str, op: str, fields: dict | None) -> None: - cls, uuid_key = _TABLES[table] - store = getattr(replica, table) - store_key = UUID(key) if uuid_key else key - if op == "delete": - 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 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.generation: str | None = None - self.seq = 0 - self.last_contact = 0.0 # monotonic time of last snapshot/event + 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() @@ -82,25 +70,24 @@ class RemoteReplica: self._stopped = True def available(self) -> bool: - """Synced and either connected now or within cache_ttl of silence. - - The websockets keepalive drops a wedged connection, so a live - connection means events arrive within one round trip; after losing - it the replica remains trusted for cache_ttl. - """ + """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, user_agent: str) -> None: + 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 - session.ip = ip - session.user_agent = user_agent + 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, @@ -110,9 +97,6 @@ class RemoteReplica: } self._refresh_signal.set() - def evict_session(self, secret: str) -> None: - self.db.sessions.pop(hash_secret("cookie", secret), None) - async def start(self) -> None: self._stopped = False self._task = asyncio.create_task(self._run()) @@ -129,7 +113,7 @@ class RemoteReplica: async def _sweep(self) -> None: while True: await asyncio.sleep(_SWEEP_INTERVAL) - limit = datetime.now(UTC) - EXPIRES + 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] @@ -151,11 +135,6 @@ class RemoteReplica: async def _connect(self) -> None: ws_url = self.remote.url.replace("http", "ws", 1) + "/auth/api/sync/ws" - # Periodic full snapshots reconcile any drift; resume is cheaper. - full_resync = self.generation is None or ( - time.monotonic() - self.last_contact > self.remote.refresh_interval - ) - resume = {} if full_resync else {"generation": self.generation, "seq": self.seq} async with websockets.connect( ws_url, additional_headers={"Authorization": f"Bearer {self.remote.token}"}, @@ -163,44 +142,43 @@ class RemoteReplica: ping_interval=5, ping_timeout=5, ) as ws: - hello = msgspec.json.decode(await ws.recv()) - if hello.get("type") != "hello": - raise ValueError("sync: expected hello") - await ws.send(msgspec.json.encode({"type": "resume", **resume})) sender = asyncio.create_task(self._send_loop(ws)) staging: DB | None = None + ready_at = 0.0 try: while True: - message = msgspec.json.decode(await ws.recv()) - self.last_contact = time.monotonic() + if staging is None: + # Periodic reconnects give full-snapshot reconciliation + remaining = self.remote.refresh_interval - ( + time.monotonic() - ready_at + ) + if remaining <= 0: + return + message = msgspec.json.decode( + await asyncio.wait_for(ws.recv(), remaining) + ) + else: + message = msgspec.json.decode(await ws.recv()) mtype = message.get("type") if mtype == "snapshot": - if staging is None: - staging = DB() + staging = staging or DB() for key, fields in message["items"]: - _apply(staging, message["table"], key, "upsert", fields) + _apply(staging, message["table"], key, fields) elif mtype == "event": - if staging is not None or ( - self.generation is not None - and message["seq"] != self.seq + 1 - ): - raise ValueError("sync: event out of order") - self.seq = message["seq"] + if staging is not None: + raise ValueError("sync: event before ready") _apply( self.db, message["table"], message["key"], - message["op"], message.get("fields"), ) elif mtype == "ready": if staging is not None: self.db = staging staging = None - attach_stores() - self.generation = hello["generation"] - self.seq = message["seq"] self.connected = True + self.last_contact = ready_at = time.monotonic() finally: sender.cancel() with contextlib.suppress(asyncio.CancelledError): @@ -215,17 +193,28 @@ class RemoteReplica: 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] = {} - def replica_for(self, domain: domains.Domain) -> RemoteReplica | None: - if domain.remote is None: - return None - return self.replicas.get(domain.remote.url) - async def start(self) -> None: domains.add_rebuild_listener(self.reconcile) await self.reconcile(domains.registry()) @@ -237,7 +226,7 @@ class SatelliteManager: self.replicas.clear() async def reconcile(self, registry: domains.DomainRegistry) -> None: - """Attach stores and start/stop replicas to match the config.""" + """Start/stop replicas to match the configured remote domains.""" wanted = {} for domain in registry.domains: if domain.remote is not None: @@ -253,15 +242,100 @@ class SatelliteManager: replica = RemoteReplica(remote) self.replicas[url] = replica await replica.start() - attach_stores() manager = SatelliteManager() -def attach_stores() -> None: - """Attach each remote domain's store to its replica.""" - for domain in domains.registry().domains: - replica = manager.replica_for(domain) - if replica is not None: - domain.store = replica.db +# ------------------------------------------------------------------------- +# 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 diff --git a/paskia/syncfeed.py b/paskia/syncfeed.py index 4fcb9fc..2ce9fb2 100644 --- a/paskia/syncfeed.py +++ b/paskia/syncfeed.py @@ -1,85 +1,50 @@ """RAM-only change feed letting satellite instances mirror this server. -Nothing here touches the database file: events are held in a bounded ring -buffer and 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. +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 itertools -import logging import os -import secrets -from collections import deque import msgspec -_logger = logging.getLogger(__name__) - # Tables mirrored by satellites (reset tokens, OIDC data and domain config # are instance-local and never replicated). TABLES = ("permissions", "orgs", "roles", "users", "credentials", "sessions") -_RING_SIZE = 2000 - - -class SyncFeed: - """Sequenced change events with replay for reconnecting satellites.""" - - def __init__(self): - self.generation = secrets.token_hex(8) - self._seq = itertools.count(1) - self.seq = 0 - self.events: deque[dict] = deque(maxlen=_RING_SIZE) - self.subscribers: set[asyncio.Queue] = set() - - def emit(self, table: str, key: str, obj) -> None: - """Publish an upsert (obj given) or delete (obj None).""" - self.seq = next(self._seq) - event = { - "type": "event", - "seq": self.seq, - "table": table, - "key": key, - "op": "upsert" if obj is not None else "delete", - "fields": msgspec.to_builtins(obj) if obj is not None else None, - } - self.events.append(event) - for queue in self.subscribers: - try: - queue.put_nowait(event) - except asyncio.QueueFull: - # Slow consumer: drop it; the client reconnects and resyncs. - self.subscribers.discard(queue) - - def replay_since(self, seq: int) -> list[dict] | None: - """Events after seq, or None when the ring no longer reaches back - (or the claimed seq is ahead of us, which cannot be reconciled).""" - if seq > self.seq: - return None - if not self.events: - return [] if seq == self.seq else None - oldest = self.events[0]["seq"] - if seq < oldest - 1: - return None - return [e for e in self.events if e["seq"] > seq] - - def subscribe(self) -> asyncio.Queue: - queue: asyncio.Queue = asyncio.Queue(maxsize=1000) - self.subscribers.add(queue) - return queue - - def unsubscribe(self, queue: asyncio.Queue) -> None: - self.subscribers.discard(queue) - - -feed = SyncFeed() +_subscribers: set[asyncio.Queue] = set() def emit(table: str, key: str, obj) -> None: - feed.emit(table, key, obj) + """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]: diff --git a/paskia/util/permutil.py b/paskia/util/permutil.py index 589c30a..b1f335e 100644 --- a/paskia/util/permutil.py +++ b/paskia/util/permutil.py @@ -129,8 +129,8 @@ def has_all_scopes_groups(scopes: set[str], groups: Sequence[Sequence[str]]) -> return all(group_satisfied(scopes, g) for g in groups) -async def session_context(auth: str | None, host: str | None = None, store=None): +async def session_context(auth: str | None, host: str | None = None): if not auth: return None normalized_host = normalize_host(host) if host else None - return session_ctx(auth, normalized_host, store=store) + return session_ctx(auth, normalized_host) diff --git a/paskia/util/userinfo.py b/paskia/util/userinfo.py index 7e4a5d6..b754d5c 100644 --- a/paskia/util/userinfo.py +++ b/paskia/util/userinfo.py @@ -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 ( @@ -41,14 +41,9 @@ async def build_user_info( session_key: str, request_host: str | None, ctx: SessionContext | None = None, - store=None, ) -> ApiUserDetail: - """Build user info struct for authenticated users. - - store defaults to the local database; remote-domain request paths pass - their domain's replica explicitly. - """ - data = store or db.data() + """Build user info struct for authenticated users.""" + data = satellite.store_for_host(request_host) user = data.users[user_uuid] normalized_host = hostutil.normalize_host(request_host) diff --git a/tests/test_remote.py b/tests/test_remote.py index b38f548..14aad7a 100644 --- a/tests/test_remote.py +++ b/tests/test_remote.py @@ -1,6 +1,5 @@ """Tests for remote (satellite) domains: config, replica application, feed.""" -import collections import secrets import time from datetime import UTC, datetime @@ -79,9 +78,9 @@ 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), "upsert", _builtins(user)) + satellite._apply(replica, "users", str(user.uuid), _builtins(user)) assert replica.users[user.uuid].display_name == "U" - satellite._apply(replica, "users", str(user.uuid), "delete", None) + satellite._apply(replica, "users", str(user.uuid), None) assert not replica.users @@ -102,7 +101,7 @@ def test_apply_session_roundtrip(): validated=datetime.now(UTC), rp_id="example.com", ) - satellite._apply(replica, "sessions", session.key, "upsert", _builtins(session)) + satellite._apply(replica, "sessions", session.key, _builtins(session)) stored = replica.sessions[session.key] assert stored.host == "app2.example.com" assert stored.validated == session.validated @@ -123,47 +122,50 @@ def test_apply_credential_bytes_roundtrip(): 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), "upsert", wire) + 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_and_replay(): - feed = syncfeed.SyncFeed() - user = User.create(display_name="A", role=UUID(int=1)) - feed.emit("users", "k1", user) - feed.emit("users", "k1", None) - assert feed.seq == 2 - assert feed.replay_since(0)[0]["op"] == "upsert" - assert feed.replay_since(1)[0]["op"] == "delete" - assert feed.replay_since(2) == [] - assert feed.replay_since(99) is None +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_ring_overflow_replay_none(): - feed = syncfeed.SyncFeed() - feed.events = collections.deque(maxlen=3) - for i in range(5): - feed.emit("users", f"k{i}", None) - assert feed.replay_since(0) is None # fell off the ring - assert [e["seq"] for e in feed.replay_since(4)] == [5] - assert feed.replay_since(5) == [] +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.""" - syncfeed.feed.events.clear() - syncfeed.feed.seq = 0 - user = next(iter(test_db.users.values())) - ops_db.update_user_display_name(user.uuid, "Renamed") - tables = {e["table"] for e in syncfeed.feed.events} - assert "users" in tables - key = syncfeed.feed.events[-1]["key"] - assert syncfeed.feed.events[-1]["fields"]["display_name"] == "Renamed" - assert key == str(user.uuid) + 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 @@ -188,8 +190,13 @@ async def test_replica_refresh_and_evict(): assert queued["type"] == "session_refresh" assert queued["ip"] == "2.2.2.2" - replica.evict_session(token) + # 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(): @@ -259,7 +266,6 @@ async def remote_client(test_db): replica.last_contact = time.monotonic() replica.connected = True satellite.manager.replicas[REMOTE_URL] = replica - satellite.attach_stores() transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://localhost:4401" @@ -322,10 +328,10 @@ async def test_remote_domain_503_when_replica_stale(remote_client): async def test_logout_proxied_and_evicted(remote_client, monkeypatch): client, secret, replica = remote_client - async def fake_proxy(request, remote): + async def fake_forward(request): return Response(status_code=200, content=b'{"message": "Logged out"}') - monkeypatch.setattr("paskia.fastapi.proxy.proxy_to_remote", fake_proxy) + 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}"},