diff --git a/.github/workflows/deploy-crabbox-coordinator.yml b/.github/workflows/deploy-crabbox-coordinator.yml index e49355cd..a1b4d89b 100644 --- a/.github/workflows/deploy-crabbox-coordinator.yml +++ b/.github/workflows/deploy-crabbox-coordinator.yml @@ -20,7 +20,7 @@ jobs: uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: repository: openclaw/crabbox - ref: main + ref: ec073fdcf868b2d5c450d600d900cdf41f958de9 # crabbox main, 2026-07-11 - name: Set up Node uses: actions/setup-node@48b55a011bda9f5d6aeb4c2d9c7362e8dae4041e # v6.4.0 diff --git a/.gitignore b/.gitignore index 91afe0bc..b7e803ca 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ node_modules src/spec.generated.ts +# Rebuilt from canonical sources by pnpm check/build/deploy; do not commit. src/generated.ts dist/ .wrangler/ diff --git a/CHANGELOG.md b/CHANGELOG.md index 31e76755..6b816cd1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,20 @@ ## Unreleased +- Resolve final audit blockers by preserving unknown GitHub Actions input outcomes across runner replacement, disconnect, upstream close or error, bounded PTY-write failure, and post-write confirmation loss, making terminal confirmation races prefer completed delivery, retiring terminal connections after late detached-writer failures, replaying control grants across attachment gaps, retaining rollback recovery and authorized namespace retirement after a staged credential write begins, reserving expired write-started credential rows for recovery, requiring active credential generations to match the exact current lookup set, proving interrupted pre-fence migration recovery, blocking session deletion while staged credential rows remain, validating ARD safe-prime groups and nonzero key material with bounded accepted-prime caching, deferring partial-fence ZRLE resets to a trailing synchronization boundary, rejecting legacy mutations and deletions of token-owned desktop registrations, discarding definitively unowned publication recovery without tokenless cleanup, documenting the opaque publication-ID contract, completing ambiguous legacy publication cleanup despite persistence failures, and holding generated-asset tests under a crash-released lock through module consumption. +- Close the final review blockers by requiring an exact live credential-policy lease at promotion, negotiating strict runtime-adapter deletion tombstones without breaking legacy `404` release semantics, disconnecting VNC sessions when pixel-format capability probes cannot establish a safe ZRLE boundary, classifying aborted JSON body streams as bad requests, and preserving live terminal subscriptions when browser history restores the sessions grid. +- Complete the final audit follow-up by fencing rollback against newer legacy credential generations, keeping viewer acknowledgement timeouts from detaching live GitHub Actions sessions, fencing queued documented-runner input after relay replacement, generating ignored embedded assets before parity tests import them, resetting ZRLE exactly at pixel-format boundaries, preventing delayed tokenless desktop cleanup from deleting replacement publishers, skipping idle recovery I/O when no local state exists, and stopping canceled auto-share preflight. +- Preserve upgrade and teardown authority by leaving pre-lookup-migration credential registrations recoverable from current and historical runtime identities, retaining GitHub Actions runner generations after queues drain, scoping Share This Mac recovery state to the normalized API origin and stable owner, allowing idle termination after unavailable recovery lookup while retaining active cleanup vetoes, terminating timed-out or canceled Tailscale descendant process groups, clearing only definitive failed publication intent, quiescing remote-input producers before final release, rejecting UltraVNC Diffie-Hellman elements at `p - 1`, and synchronizing complete RFB pixel-format and encoding transitions with protocol-required ZRLE resets. +- Finish the audited terminal, credential, runtime, and native-app lifecycle boundaries by explicitly negotiating the generation-fenced GitHub Actions runner protocol, retiring stale and overflowing runner input queues, capturing relay replacement during viewer authorization, fencing retired Go terminal attachments, repairing credential lookup namespaces with rollback-compatible staging, requiring replayable runtime-adapter deletion tombstones, validating Apple Remote Desktop Diffie-Hellman groups, keying desktop publication cleanup by the API host ID, and requiring exact retained identity before uncertain publication recovery; the runner guide now preserves raw fallback, bounds admission before serialized restricted steering, and distinguishes unknown delivery from rejection. +- Harden session lifecycle concurrency with atomic card claims and duplicate-first claim results, single-winner GitHub Actions credential rotation, monotonic exact authenticated-revision fences on runner writes, terminal session-status fences, revision-fenced lifecycle updates and grant revocation, exclusively claimed rollback recovery for Sandbox credential rotation, rejection of live legacy registration claims during migration without staging renewable legacy claims, staged-rotation fences that block legacy policy mutation while new-worker claims are live but release abandoned rows for rollback compatibility, persisted staged lookup identities across namespace changes with exact current-identity fallback for mixed-version rows, ownership-fenced repair of incomplete and rotated lookup sets before credential rotation, explicit retirement of obsolete durable identities, idempotent recovery after ambiguous committed promotion, R2-clean reservation rollback, preserved Sandbox attachment state, retained registration data for superseded runtime workspace cleanup, and durable observed-deletion markers that terminate cleanup after post-delete crashes. +- Close final terminal and desktop publication race windows by carrying the initial GitHub Actions runner generation through viewer authorization, translating generation-fenced acknowledgements for legacy framed viewers, serializing and bounding per-runner PTY input by frames, bytes, and age, matching generation-fenced local send failures, ordering raw and confirmed Go client acknowledgements, bounding shutdown when terminal writers block, rejecting malformed desktop recovery IDs as client errors, preserving idempotent publication retries across mixed worker versions, and retaining uncertain Share This Mac publications when older servers lack the recovery route. +- Make terminal input delivery durable across multiplex subscribers, apply backpressure until an attachment owns initial output while waking blocked readers to discard and acknowledge output for one-shot confirmed messages, retain pre-attach closure, prefer completed input over later transport shutdown, reject denied or unacknowledged input explicitly, prevent delivery into retired attachments and wait for their frame consumers even when input reads cannot be canceled, bound serialized browser input backlog by frame count and bytes while preserving one ordered completion per dropped frame, enforce relay-owned runner generations before forwarding GitHub Actions input and acknowledgements, snapshot SSH connection limits before launching handlers, make confirmation serialization cancelable, bound attachment confirmation waits, serialize every acknowledgement-aware input source and per-subscription completion event, wake uncancelable attachments when their context ends, retire connections after ambiguous confirmation timeouts, negotiate bounded one-shot input acknowledgements across rolling upgrades without waiting on empty payloads, keep configured HTTP timeouts from canceling established sockets, scope rejected writes without dropping live subscriptions, and send attributed commands atomically to prevent interleaving. +- Add connection-query-negotiated `CFR1` input, output, and acknowledgement frames across GitHub Actions runner and internal viewer relay boundaries, confirm viewer negotiation before leaving raw fallback during mixed deployments, document the independent legacy viewer fallback, buffer split UTF-8 within byte, frame, and age bounds until the string-only Node adapter delivers it to the PTY before acknowledging every contributing frame, define that adapter's UTF-8-only output contract while preserving opaque bytes for byte-oriented adapters, close the runner socket when its PTY exits, prevent binary PTY output from colliding with relay control frames, and keep multiplex dispatch responsive while framed input acknowledgements are pending. +- Preserve bounded opaque profile IDs for fixed runtime adapters while rejecting unroutable or ambiguous adapter routes only when provisioning depends on them, so mixed migration configuration cannot break unrelated control-plane reads; durably claim and retry superseded workspace cleanup without touching the replacement workspace; also reject malformed encoded session routes, numeric literals that become integers only after precision loss, and invalid-Unicode JSON event values, and reconcile browser history drawers and focus on back/forward navigation. +- Harden Share This Mac against stale starts and responses, canceled starts stranded in transition, self-connections, leaked subprocess environment, stalled Tailscale and RFB handshakes, successful commands whose descendants retain output pipes, stale desktop registrations including listener-failure, ambiguous committed publication reconciled by stable publication identity, and application-termination races with durably retained cleanup retries, legacy publishers mutating or deleting token-owned registrations, concurrent teardown calls that could outpace application termination, completed teardown operations coalescing a later stop, dropped auto-starts, stuck remote input including releases retained through revoked Accessibility trust and teardown with bounded retries and no retry when no input is held, and destructive non-text or empty clipboard changes while preserving repeated remote text, X11 Unicode input, proxy, custom-CA networking, and validated Crabbox config/state paths. +- Fence Share This Mac registry cleanup with explicitly negotiated per-registration ownership tokens, require a valid publication identity before selecting token ownership, and return the exact atomically written registration row so delayed or overlapping current publishers cannot displace cleanup authority, while preserving tokenless registration and cleanup for rolling upgrades with legacy clients or servers. +- Harden the bundled RoyalVNCKit fork across ARD and UltraVNC authentication and parameter validation, composed Unicode keysyms with legacy ASCII scalars, Tight and ZRLE parsing, non-trapping bounded zlib streams, RFB Fence-synchronized color-depth transitions with atomic capability publication, premature-response rejection, and fail-closed legacy handling, CopyRect, EOF handling, reconnects, credential cancellation and release after handoff, and cursor channel preservation. +- Pin Crabbox coordinator deployment to an immutable source revision and verify downloaded Crabbox release archives before image assembly, always enforcing the repository digest for the default version and requiring an explicit architecture checksum for non-default versions. - Add a VideoToolbox-backed Open H.264 RFB pipeline for Share This Mac with up to 60 fps capture, adaptive 1.5–30 Mbit/s rate control, automatic Tight/JPEG fallback, live stream stats, larger resize limits, and a persisted host-enforced view-only mode. - Exchange full UTF-8 clipboard text between the native Mac viewer, Share This Mac hosts, and any Extended Clipboard-capable VNC server by completing the RoyalVNCKit fork's extension stub, keeping Latin-1 cut text as the fallback and dropping malformed extension bodies without tearing down the connection. - Add persisted send-only and receive-only clipboard directions to the native viewer's focus toolbar; automatic sync respects the direction while the explicit Send and Get actions keep working. diff --git a/Dockerfile b/Dockerfile index 2e919eb9..78a4f335 100644 --- a/Dockerfile +++ b/Dockerfile @@ -3,6 +3,8 @@ FROM docker.io/cloudflare/sandbox:0.10.1 USER root ARG CRABBOX_VERSION=0.17.1 +ARG CRABBOX_SHA256_AMD64= +ARG CRABBOX_SHA256_ARM64= RUN set -eux; \ apt-get update; \ @@ -61,10 +63,34 @@ RUN set -eux; \ RUN set -eux; \ arch="$(dpkg --print-architecture)"; \ - case "$arch" in amd64|arm64) ;; *) echo "unsupported arch: $arch" >&2; exit 1 ;; esac; \ + case "$arch" in \ + amd64) \ + override_checksum="$CRABBOX_SHA256_AMD64"; \ + checksum_arg="CRABBOX_SHA256_AMD64"; \ + pinned_checksum="3c41839257e4622e28bcec8b0f0153f19d78d436fd548894a7c7d7726d922611" \ + ;; \ + arm64) \ + override_checksum="$CRABBOX_SHA256_ARM64"; \ + checksum_arg="CRABBOX_SHA256_ARM64"; \ + pinned_checksum="4bf87a0d2365441ee2f8cb34183cfd9ebeb065111697eb2d8dc867b3a627fdd2" \ + ;; \ + *) echo "unsupported arch: $arch" >&2; exit 1 ;; \ + esac; \ + archive="crabbox_${CRABBOX_VERSION}_linux_${arch}.tar.gz"; \ + if [ "$CRABBOX_VERSION" = "0.17.1" ]; then \ + checksum="$pinned_checksum"; \ + else \ + checksum="$override_checksum"; \ + if [ -z "$checksum" ]; then \ + echo "an explicit $checksum_arg is required for non-default versions" >&2; \ + exit 1; \ + fi; \ + fi; \ + printf '%s\n' "$checksum" | grep -Eq '^[0-9a-f]{64}$'; \ curl -fsSL \ - "https://github.com/openclaw/crabbox/releases/download/v${CRABBOX_VERSION}/crabbox_${CRABBOX_VERSION}_linux_${arch}.tar.gz" \ + "https://github.com/openclaw/crabbox/releases/download/v${CRABBOX_VERSION}/${archive}" \ -o /tmp/crabbox.tar.gz; \ + echo "$checksum /tmp/crabbox.tar.gz" | sha256sum -c -; \ tar -xzf /tmp/crabbox.tar.gz -C /usr/local/bin crabbox; \ chmod +x /usr/local/bin/crabbox; \ rm -f /tmp/crabbox.tar.gz; \ diff --git a/README.md b/README.md index d4c228b9..aace4a72 100644 --- a/README.md +++ b/README.md @@ -94,14 +94,27 @@ Content-Type: application/json {"workKey":"openclaw/crabfleet:pr:42","workKind":"pr_repair","repo":"openclaw/crabfleet","branch":"fix/pr-42","owner":"operator@example.test","sourceUrl":"https://github.com/openclaw/crabfleet/pull/42","runUrl":"https://github.com/openclaw/crabfleet/actions/runs/123","purpose":"repair PR 42","summary":"starting repair"} ``` -The response contains `{session, agentToken, runnerPtyUrl, browserUrl}`. New registrations and resumes require `owner` to resolve to one active Crabfleet user; resumes must prove the same stable owner subject already recorded on the `workKey`. The stable subject owns browser visibility while the OpenClaw service retains lifecycle authority for its session. `runnerPtyUrl` includes the rotated session-scoped query credential and works directly with Node's global `WebSocket`: +The response contains `{session, agentToken, runnerPtyUrl, browserUrl}`. New registrations and resumes require `owner` to resolve to one active Crabfleet user; resumes must prove the same stable owner subject already recorded on the `workKey`. The stable subject owns browser visibility while the OpenClaw service retains lifecycle authority for its session. `runnerPtyUrl` includes the rotated session-scoped query credential and can be opened unchanged as the legacy raw duplex byte stream. New runners offer the framed protocol as a WebSocket subprotocol: ```js -const terminal = new WebSocket(runnerPtyUrl); -terminal.onmessage = (event) => process.stdout.write(Buffer.from(event.data)); -process.stdin.on("data", (chunk) => terminal.send(chunk)); +const terminal = new WebSocket(runnerPtyUrl, "cfr1-framed-io-v2"); +terminal.binaryType = "arraybuffer"; +terminal.addEventListener("open", () => { + const framed = terminal.protocol === "cfr1-framed-io-v2"; + // Use CFR1 only when framed is true; otherwise retain raw compatibility. +}); ``` +Existing runners retain raw input and output by opening the returned URL +unchanged. A new runner opts into correlated binary `CFR1` input, output, and +acknowledgement frames only when the relay selects the +`cfr1-framed-io-v2` WebSocket subprotocol in the upgrade response. Relays that +ignore the offer leave `WebSocket.protocol` empty, so new runners retain raw +compatibility. The relay fences negotiated runner input with its connection +generation. Framed runners acknowledge only after their PTY accepts the input. +The complete byte-safe encoder, decoder, and Node PTY runner are in +[`docs/github-actions-sessions.md`](docs/github-actions-sessions.md#runner-pty). + The runner reports heartbeat and durable progress with bearer `agentToken` to `POST /api/agent/interactive-sessions/:id/work-state`. Terminal states are `completed`, `blocked`, `failed`, and `canceled`; active work uses `registered` or `running` plus a specific `phase`. The full registration, relay, resumption, steering, heartbeat, completion, diff --git a/cmd/crabbox-ssh-gateway/main.go b/cmd/crabbox-ssh-gateway/main.go index a4970a83..17e07501 100644 --- a/cmd/crabbox-ssh-gateway/main.go +++ b/cmd/crabbox-ssh-gateway/main.go @@ -98,6 +98,22 @@ type sessionPTY struct { resizes chan fleetapi.TerminalSize } +type sshConnectionSettings struct { + handshakeTimeout time.Duration + connectionIdle time.Duration + sessionIdle time.Duration + sessionChannels int +} + +func currentSSHConnectionSettings() sshConnectionSettings { + return sshConnectionSettings{ + handshakeTimeout: sshHandshakeTimeout, + connectionIdle: sshConnectionIdle, + sessionIdle: sshSessionIdleTimer, + sessionChannels: sshSessionChannels, + } +} + func main() { var addr string var apiURL string @@ -176,6 +192,7 @@ func main() { } func acceptConn(raw net.Conn, config *ssh.ServerConfig, client *apiClient) { + settings := currentSSHConnectionSettings() if !sshConnectionSlots.acquire() { log.Printf("connection limit reached for %s", raw.RemoteAddr()) raw.Close() @@ -187,17 +204,25 @@ func acceptConn(raw net.Conn, config *ssh.ServerConfig, client *apiClient) { raw.Close() return } - go handleConnWithRelease(raw, config, client, sshHandshakeSlots.release, sshConnectionSlots.release) + go handleConnWithRelease( + raw, + config, + client, + settings, + sshHandshakeSlots.release, + sshConnectionSlots.release, + ) } func handleConn(raw net.Conn, config *ssh.ServerConfig, client *apiClient) { - handleConnWithRelease(raw, config, client, nil, nil) + handleConnWithRelease(raw, config, client, currentSSHConnectionSettings(), nil, nil) } func handleConnWithRelease( raw net.Conn, config *ssh.ServerConfig, client *apiClient, + settings sshConnectionSettings, releaseHandshake func(), releaseConnection func(), ) { @@ -212,7 +237,7 @@ func handleConnWithRelease( } }() defer raw.Close() - if err := raw.SetDeadline(time.Now().Add(sshHandshakeTimeout)); err != nil { + if err := raw.SetDeadline(time.Now().Add(settings.handshakeTimeout)); err != nil { log.Printf("handshake deadline %s: %v", raw.RemoteAddr(), err) } conn, chans, reqs, err := ssh.NewServerConn(raw, config) @@ -234,12 +259,12 @@ func handleConnWithRelease( if permissions == nil { permissions = &ssh.Permissions{Extensions: map[string]string{}} } - sessionSlots := newConnectionLimiter(sshSessionChannels) + sessionSlots := newConnectionLimiter(settings.sessionChannels) sessionDone := make(chan struct{}) connectionClosed := make(chan struct{}) defer close(connectionClosed) activeSessions := 0 - connectionIdleTimer, connectionIdle := newConnectionIdleTimer() + connectionIdleTimer, connectionIdle := newIdleTimer(settings.connectionIdle) defer stopTimer(connectionIdleTimer) for { select { @@ -251,7 +276,7 @@ func handleConnWithRelease( } if activeSessions == 0 { stopTimer(connectionIdleTimer) - connectionIdleTimer, connectionIdle = newConnectionIdleTimer() + connectionIdleTimer, connectionIdle = newIdleTimer(settings.connectionIdle) } case ch, ok := <-chans: if !ok { @@ -277,12 +302,12 @@ func handleConnWithRelease( activeSessions-- } if activeSessions == 0 { - connectionIdleTimer, connectionIdle = newConnectionIdleTimer() + connectionIdleTimer, connectionIdle = newIdleTimer(settings.connectionIdle) } log.Printf("channel accept: %v", err) continue } - go handleSession(channel, requests, permissions, client, func() { + go handleSession(channel, requests, permissions, client, settings.sessionIdle, func() { sessionSlots.release() select { case sessionDone <- struct{}{}: @@ -298,6 +323,7 @@ func handleSession( requests <-chan *ssh.Request, perms *ssh.Permissions, client *apiClient, + idleTimeout time.Duration, release func(), ) { if release != nil { @@ -306,7 +332,7 @@ func handleSession( defer channel.Close() ctx, cancel := context.WithCancel(context.Background()) defer cancel() - idleTimer, idle := newSessionIdleTimer() + idleTimer, idle := newIdleTimer(idleTimeout) defer stopTimer(idleTimer) pty := sessionPTY{ cols: 120, @@ -398,19 +424,11 @@ func handleSession( } } -func newSessionIdleTimer() (*time.Timer, <-chan time.Time) { - if sshSessionIdleTimer <= 0 { - return nil, nil - } - timer := time.NewTimer(sshSessionIdleTimer) - return timer, timer.C -} - -func newConnectionIdleTimer() (*time.Timer, <-chan time.Time) { - if sshConnectionIdle <= 0 { +func newIdleTimer(timeout time.Duration) (*time.Timer, <-chan time.Time) { + if timeout <= 0 { return nil, nil } - timer := time.NewTimer(sshConnectionIdle) + timer := time.NewTimer(timeout) return timer, timer.C } diff --git a/docs/api.md b/docs/api.md index f35f50d2..651cb784 100644 --- a/docs/api.md +++ b/docs/api.md @@ -507,7 +507,7 @@ Crabfleet authenticates every adapter request with `Authorization: Bearer CRABBO - `POST /v1/workspaces`: idempotent create. Crabfleet persists the deterministic adapter identity, TTL, idle timeout, requested capabilities, and exact serialized create payload before the request, then sends the same namespaced DNS-safe lowercase `id` and `Idempotency-Key`, plus repo, branch, runtime, opaque profile, command, prompt, ownership/lineage, and lifecycle settings. A definitive non-2xx response to the initial request is read once, sanitized, and durably recorded as the failure reason before provider release begins. After an ambiguous result, a bounded reconciliation pass retries only that immutable payload and key before any inspect; later edits to session metadata do not alter it. Replay-time authentication, routing, validation, or other non-success responses cannot prove the original request failed and therefore keep create ambiguity pending. - An adapter that finds the requested ID already bound to a different immutable request returns `409` with `error.code = "workspace_id_conflict"`. Crabfleet marks only its local session failed and atomically drops that adapter identity when the exact pending create attempt still owns the lifecycle revision and reconciliation claim; a stale conflict response is ignored. It never adopts, inspects, or deletes the pre-existing workspace. Other `409` responses remain ambiguous and retryable. - `GET /v1/workspaces/:id`: inspect current status, capabilities, terminal URL, expiry, and provider resource identity. Status-only responses preserve previously stored capabilities and expiry; explicit `null` clears those fields. Active external sessions are reconciled in bounded batches; state responses wait only for a short foreground budget while remaining work continues in the Worker background. -- `DELETE /v1/workspaces/:id`: stop/release. Crabfleet enters `stopping` before calling the adapter and marks the session stopped only after `204`, `404`, or a valid exact-ID terminal response confirms release; malformed successful bodies remain `stopping`. Plain-text and malformed-JSON responses are read once and sanitized before their evidence is retained. An explicit stop whose ownership claim loses returns success only when the exact workspace is already stopping or terminal; otherwise it returns a lifecycle conflict. +- `DELETE /v1/workspaces/:id`: idempotent stop/release. Crabfleet advertises `delete-tombstone-v1` in `X-Crabfleet-Runtime-Adapter-Capabilities`; adapters that support the stricter contract echo that token in the response header. Before provider release, a supporting adapter must durably retain the exact workspace identity and stopping intent. Once a DELETE is accepted, every retry for that immutable ID must return `204` or a valid exact-ID `stopping`, `stopped`, or `expired` response, including after provider deletion or adapter restart; it must not collapse that lifecycle tombstone into `404`. This lets a caller recover when deletion commits but the response is lost without treating a pre-visibility `404` from an ambiguous create as release proof. Crabfleet enters `stopping` before calling the adapter and marks the session stopped only after `204`, a `404` when create ambiguity is absent, prior deletion evidence is durable, or the adapter does not echo the capability during a rolling upgrade, or a valid exact-ID terminal response confirms release; malformed successful bodies remain `stopping`. Plain-text and malformed-JSON responses are read once and sanitized before their evidence is retained. An explicit stop whose ownership claim loses returns success only when the exact workspace is already stopping or terminal; otherwise it returns a lifecycle conflict. - `POST /v1/workspaces/:id/connections/desktop`: mint a current transient desktop URL. The request has no body. `expiresAt` is optional; when present it must be in the future and no more than 15 minutes away. Accepted HTTPS URLs are treated as opaque signed connection material and redirected byte-for-byte without URL normalization. After minting, Crabfleet re-reads the exact current session status, control grant, capabilities, and registered adapter identity before redirecting; a concurrent stop, revocation, capability withdrawal, or lifecycle replacement discards the URL and denies access. - `POST /v1/workspaces/:id/connections/native-vnc`: mint a short-lived, single-use native VNC grant. The response must use the `crabbox/native-vnc-grant/v1` schema, an HTTPS broker URL (literal loopback HTTP is allowed for development), the exact opaque lease ID, a 32-byte-hex `native_vnc_` ticket, and an expiry no more than two minutes away. Crabfleet never exposes the provider lease ID in Fleet state and requests this grant only after revalidating current session control and the persisted adapter identity. @@ -598,11 +598,68 @@ Response: } ``` -Every new registration and every resume requires `owner`; it must resolve to exactly one active Crabfleet user by login, email, or stable subject. Existing work keys resume only when the supplied owner resolves to the same stable owner subject already stored on the work key. Ownerless resumes fail closed before token rotation, and a work key cannot transfer to a different stable owner. `runnerPtyUrl` is directly usable with Node's global `WebSocket`; no custom headers are required. The query credential is session-scoped, rotates on registration, is stored only as a hash, and is not exposed through viewer/session APIs. +Every new registration and every resume requires `owner`; it must resolve to exactly one active Crabfleet user by login, email, or stable subject. Existing work keys resume only when the supplied owner resolves to the same stable owner subject already stored on the work key. Ownerless resumes fail closed before token rotation, and a work key cannot transfer to a different stable owner. `runnerPtyUrl` can be opened with Node's global `WebSocket` without custom headers. Existing runners retain raw input/output by opening it unchanged; new runners offer `cfr1-framed-io-v2` as a WebSocket subprotocol and enter the generation-fenced contract below only when the relay selects it. The query credential is session-scoped, rotates on registration, is stored only as a hash, and is not exposed through viewer/session APIs. ### GET /api/agent/interactive-sessions/:id/runner-pty -WebSocket endpoint for the outbound GitHub Actions runner. Authentication uses the scoped `agentToken` query parameter embedded in `runnerPtyUrl`. The runner sends raw terminal output bytes and receives raw viewer input bytes. One runner is current; a reconnect replaces the previous runner while browser viewers remain attached. +WebSocket endpoint for the outbound GitHub Actions runner. Authentication uses the scoped `agentToken` query parameter embedded in `runnerPtyUrl`. One runner is current; a reconnect replaces the previous runner while browser viewers remain attached. + +Opening the returned URL unchanged selects legacy raw input and output. Adding +`cfr1-framed-io-v2` to the `WebSocket` constructor's protocol list offers +generation-fenced input, output, acknowledgements, and relay control traffic. +The runner switches formats only when `WebSocket.protocol` confirms that exact +selection. An older relay that ignores the offer therefore remains a raw +connection. `SessionControlDO` stores the selected mode and a relay-owned runner +generation on the server socket before accepting it. Viewer framing is +negotiated independently: v2 viewers receive +`CFR1` output and generation-bearing control frames, while unnegotiated viewers +retain raw output and legacy JSON notices. The relay translates the earlier +`cfr1-framed-io-v1` format and raw sockets at each boundary during rolling +upgrades. Arbitrary raw PTY bytes cannot be consumed as control traffic by +framed viewers. + +The `runnerProtocol` query remains accepted for compatibility with already +deployed query-aware runners. New runners must use subprotocol negotiation so +they do not switch formats against a relay that did not explicitly confirm v2. + +Each `CFR1` frame occupies one binary WebSocket message and starts with: + +| Offset | Size | Value | +| ------ | -------- | ---------------------------------------- | +| 0 | 4 | ASCII `CFR1` (`43 46 52 31` hexadecimal) | +| 4 | 1 | frame type | +| 5 | 1 | input ID byte length | +| 6 | variable | input ID, then type-specific payload | + +Input IDs are nonempty ASCII `[A-Za-z0-9_-]` values of at most 80 bytes. + +| Type | Direction | Payload | +| ---------------------- | ---------------- | --------------------------------------------------------------------------------------------- | +| `0x05` input | relay to runner | generation-length byte, generation, then raw terminal input bytes | +| `0x06` acknowledgement | runner to relay | generation envelope, then `1` accepted or `0` rejected, followed by optional UTF-8 error text | +| `0x07` lifecycle event | relay to viewers | empty input ID, generation envelope, and one event-code byte | +| `0x04` output | runner to relay | empty input ID followed by raw terminal output bytes | + +Lifecycle event codes are `0x01` runner connected, `0x02` runner disconnected, +and `0x03` runner waiting. + +The runner must copy the input frame's ID and generation into its +acknowledgement. It must send an accepted acknowledgement only after its PTY +write API has accepted the payload. Queueing the frame in `WebSocket.send()` is +not acceptance. Crabfleet rejects stale-generation input before forwarding it, +and rejects input when no current runner is available or the relay send fails. +Stale or mismatched acknowledgement IDs or generations do not complete another +pending input. + +The v1 `0x01`, `0x02`, and `0x03` frames omit generations. They remain accepted +for mixed-version deployments and are translated by the relay. + +For legacy connections, the relay unwraps viewer input to raw bytes and reports +acceptance once the runner socket accepts the send. Framed connections provide +PTY-level completion through correlated acknowledgements. + +See [GitHub Actions Sessions](/github-actions-sessions/#runner-pty) for a +complete Node runner integration. ### POST /api/agent/interactive-sessions/:id/work-state @@ -1071,10 +1128,45 @@ The host is visible only to the same stable user, regardless of shared or private session-tenancy mode. Re-registering the same ID updates its name, address, port, and timestamp while preserving its creation time. +Clients opt into fenced registration by sending +`X-Crabfleet-Ownership-Mode: token-v1` and a stable +`X-Crabfleet-Publication-ID`. The publication ID identifies one client's +attempt across retries and restarts; it is an opaque 1-200 byte value that +cannot contain whitespace or control characters. The response includes an `ownershipToken` +required for deletion. Omitting the ownership-mode header preserves the legacy +`{ "host": ... }` response and stores a tokenless registration so older clients +can still clean up during rolling upgrades. Current clients tolerate a legacy +server response without `ownershipToken` and retain enough state for guarded +legacy cleanup. + +### POST /api/desktop-hosts/:id?recover=1 + +Recovers the ownership token after a fenced `PUT` may have committed but its +response was lost. The signed-in viewer and host ID must match the original +registration, and the JSON body carries the same stable publication ID: + +```json +{ + "publicationID": "01JZDESKTOPPUBLICATION" +} +``` + +The response is `{ "ownershipToken": "..." }` when that publication still owns +the host, or `{ "ownershipToken": null }` when it does not. Recovery never +reassigns ownership and cannot replace a newer publisher. A route-level `404` +means the server predates publication recovery, not that the original `PUT` +definitely failed; rolling-upgrade clients must preserve the uncertain +registration or use guarded legacy cleanup rather than silently discarding it. + ### DELETE /api/desktop-hosts/:id Removes one registered desktop owned by the signed-in viewer. The route cannot -remove another user's record with the same ID. +remove another user's record with the same ID. Fenced registrations require the +exact `X-Crabfleet-Ownership-Token` returned by `PUT`. Legacy clients may omit +the header only to remove a registration whose stored ownership token is empty; +omission never removes a tokenized registration. Legacy writes or deletes that +target a tokenized registration fail explicitly instead of reporting success +without changing the row. ## Static Routes diff --git a/docs/architecture.md b/docs/architecture.md index 0a35ed29..52dc6317 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -95,7 +95,7 @@ D1 is canonical for product metadata: ### Durable Objects - `Sandbox` runs first-party Cloudflare Sandbox workspaces. -- `SessionControlDO` stores generation-fenced Sandbox credential/checkpoint state and relays one current GitHub Actions runner to multiple viewers. +- `SessionControlDO` stores generation-fenced Sandbox credential/checkpoint state and relays one current GitHub Actions runner to multiple viewers. Existing runners retain raw input/output at the runner boundary; new runners use relay-generation-fenced binary `CFR1` input, lifecycle, and acknowledgement frames only after the upgrade response selects the offered v2 WebSocket subprotocol. Ignored offers remain raw-compatible. Viewer framing is negotiated independently: opted-in viewers receive `CFR1` terminal, lifecycle, and acknowledgement frames, legacy viewers retain raw terminal output and JSON control-message fallbacks, and the relay translates v1 framed peers during rolling upgrades. There is no `BoardDO` or `RunDO`. General Board/Fleet state is D1 plus REST polling. @@ -137,7 +137,7 @@ Interactive sessions are the live execution plane. Supported paths: - **Built-in Sandbox:** Worker provisions a Cloudflare Sandbox, prepares the repo, starts a Codex-capable shell, and proxies PTY traffic. - **Versioned runtime adapter:** Worker durably registers a tenant-namespaced workspace ID, creates and reconciles the provider workspace, proxies PTY access, mints transient desktop links, and confirms provider release before terminal state. -- **GitHub Actions:** OpenClaw automation registers a logical work key; an Actions runner connects outbound to `SessionControlDO`, reports work state, and receives browser steering. +- **GitHub Actions:** OpenClaw automation registers a logical work key; an Actions runner connects outbound to `SessionControlDO`, reports work state, and either retains legacy raw terminal traffic or opts into correlated `CFR1` browser input/output through the connection URL. Framed runners acknowledge each PTY write before Crabfleet reports input acceptance. Viewers separately negotiate framed output and control traffic, with raw terminal output and JSON notices preserved for legacy viewers. Sessions can carry a stable tenant owner, parent/root lineage, purpose, summary, named grants, public share state, delegated control, multiplayer mode, archive metadata, and runtime-specific capability state. @@ -157,7 +157,7 @@ An optional `CRABFLEET_RUNTIME_PROFILES_JSON` allowlist exposes generic Crabbox - `POST /v1/workspaces`: idempotent create using an immutable namespaced ID and request snapshot. - `GET /v1/workspaces/:id`: inspect status, capabilities, expiry, provider identity, and terminal connection. Native-only VNC uses a separate `nativeVnc` capability and does not imply the browser desktop endpoint. -- `DELETE /v1/workspaces/:id`: release the provider workspace. +- `DELETE /v1/workspaces/:id`: idempotently release the provider workspace while retaining an exact-ID stopping or terminal tombstone for retries after response loss. - `POST /v1/workspaces/:id/connections/desktop`: mint a short-lived desktop URL. Important invariants: @@ -167,6 +167,7 @@ Important invariants: - Redirects are rejected so bearer credentials cannot cross origins. - Request/response bodies are bounded before parsing. - Create ambiguity replays the exact original idempotent request before inspection. +- Delete retries replay the retained exact-ID stopping or terminal lifecycle instead of degrading to an ambiguous 404. - An explicit workspace-ID conflict never adopts or deletes the existing provider workspace. - Provider failure is not terminal until DELETE confirms release. - Status, capabilities, expiry, terminal state, and cleanup use compare-and-swap ownership fences. diff --git a/docs/github-actions-sessions.md b/docs/github-actions-sessions.md index bc095a93..a7e20370 100644 --- a/docs/github-actions-sessions.md +++ b/docs/github-actions-sessions.md @@ -41,7 +41,8 @@ flowchart LR E[GitHub Actions runner] -->|outbound WebSocket| D F[Browser Ghostty viewer] -->|terminal hub| D F -->|input| D - D -->|raw input bytes| E + D -->|CFR1 input frame| E + E -->|CFR1 acknowledgement| D E -->|Codex turn/steer| G[Codex app-server] E -->|heartbeat and work state| B B --> H[(R2 event archives)] @@ -261,35 +262,377 @@ returns only the sanitized event. ## Runner PTY -The Action connects outbound to the returned `runnerPtyUrl`: +The Action connects outbound to the returned `runnerPtyUrl`. Node's global +`WebSocket` can open the URL without custom headers. + +The returned URL opens a legacy raw-input/raw-output socket. A runner offers +`cfr1-framed-io-v2` as a WebSocket subprotocol and switches to collision-free +framed I/O only when the upgrade response selects it. The relay records that +mode and a relay-owned runner generation before accepting the connection. +Viewer input then arrives in a binary `CFR1` frame carrying a correlation ID +and that generation, and runner output uses a distinct `CFR1` output frame. The +runner copies the generation into its correlated acknowledgement only after its +restricted steering handler accepts the input. + +Complete Node framing adapter around a restricted Codex steering handler: ```js -const terminal = new WebSocket(runnerPtyUrl); +import { + closeSteering, + deliverSteeringInput, + subscribeSteeringExit, + subscribeSteeringOutput, +} from "./restricted-codex-steering.js"; + +const runnerPtyUrl = process.env.CRABFLEET_RUNNER_PTY_URL; +if (!runnerPtyUrl) throw new Error("CRABFLEET_RUNNER_PTY_URL is required"); + +const magic = new Uint8Array([0x43, 0x46, 0x52, 0x31]); // CFR1 +const inputIdDecoder = new TextDecoder(); +const encoder = new TextEncoder(); +const maxAdmittedInputBytes = 16 * 1024; +const maxAdmittedInputFrames = 32; +const maxPendingInputAgeMs = 1_000; +let admittedInputBytes = 0; +let admittedInputFrames = 0; +let pendingInputs = []; +let pendingInputBytes = 0; +let pendingInputTimer; +let inputQueue = Promise.resolve(); +let terminalClosed = false; +let activeGeneration; +const terminal = new WebSocket(runnerPtyUrl, "cfr1-framed-io-v2"); terminal.binaryType = "arraybuffer"; -terminal.onmessage = (event) => { - // Browser input bytes for the active runner. -}; +await new Promise((resolve, reject) => { + terminal.addEventListener("open", resolve, { once: true }); + terminal.addEventListener("error", reject, { once: true }); +}); +const framed = terminal.protocol === "cfr1-framed-io-v2"; +// An empty protocol means an older relay kept this socket in legacy raw mode. + +subscribeSteeringOutput((outputText) => { + terminal.send(framed ? encodeUtf8Output(outputText) : outputText); +}); + +subscribeSteeringExit(() => { + if (terminal.readyState < WebSocket.CLOSING) terminal.close(1000, "pty exited"); +}); + +terminal.addEventListener("message", (event) => { + const input = admitInput(event.data); + if (!input) return; + inputQueue = inputQueue.then(() => { + if (!inputIsActive(input)) { + releaseInputs([input]); + return; + } + return acceptInput(input); + }); +}); + +async function acceptInput(input) { + pendingInputs.push(input); + pendingInputBytes += input.payload.byteLength; + + const payload = new Uint8Array(pendingInputBytes); + let offset = 0; + for (const pending of pendingInputs) { + payload.set(pending.payload, offset); + offset += pending.payload.byteLength; + } + + let text; + try { + text = decodeCompleteUtf8(payload); + } catch { + rejectInputs(takePendingInputs(), 1007, "invalid UTF-8 input"); + return; + } + if (text === null) { + armPendingInputTimer(); + return; + } + const inputs = takePendingInputs(); + try { + await deliverSteeringInput(text); + settleInputs(inputs, true); + } catch { + rejectInputs(inputs, 1011, "steering rejected input"); + } +} + +function admitInput(data) { + if (!framed && typeof data === "string" && data.length > maxAdmittedInputBytes) { + closeRawOverflow(); + return null; + } + const input = framed ? decodeInput(data) : decodeRawInput(data); + if (!input) return null; + if (framed) { + if (activeGeneration === undefined) { + activeGeneration = input.generation; + } else if (input.generation !== activeGeneration) { + sendAck(input, false); + return null; + } + } + const nextBytes = admittedInputBytes + input.payload.byteLength; + const nextFrames = admittedInputFrames + 1; + if (nextBytes > maxAdmittedInputBytes || nextFrames > maxAdmittedInputFrames) { + if (framed) { + sendAck(input, false); + } else { + closeRawOverflow(); + } + return null; + } + admittedInputBytes = nextBytes; + admittedInputFrames = nextFrames; + return input; +} + +function inputIsActive(input) { + return ( + !terminalClosed && + terminal.readyState === WebSocket.OPEN && + (!framed || input.generation === activeGeneration) + ); +} + +function decodeRawInput(data) { + if (typeof data === "string") return { payload: encoder.encode(data) }; + if (data instanceof ArrayBuffer) return { payload: new Uint8Array(data) }; + return null; +} + +function decodeCompleteUtf8(payload) { + const decoder = new TextDecoder("utf-8", { fatal: true, ignoreBOM: true }); + const text = decoder.decode(payload, { stream: true }); + return encoder.encode(text).byteLength === payload.byteLength ? text : null; +} + +function armPendingInputTimer() { + if (pendingInputTimer) return; + pendingInputTimer = setTimeout(() => { + rejectInputs(takePendingInputs(), 1007, "incomplete UTF-8 input"); + }, maxPendingInputAgeMs); +} + +function takePendingInputs() { + if (pendingInputTimer) clearTimeout(pendingInputTimer); + pendingInputTimer = undefined; + const inputs = pendingInputs; + pendingInputs = []; + pendingInputBytes = 0; + return inputs; +} + +function settleInputs(inputs, accepted) { + releaseInputs(inputs); + if (!framed) return; + for (const input of inputs) { + sendAck(input, accepted); + } +} + +function rejectInputs(inputs, rawCloseCode, rawCloseReason) { + settleInputs(inputs, false); + if (!framed && terminal.readyState < WebSocket.CLOSING) { + terminal.close(rawCloseCode, rawCloseReason); + } +} -function writeTerminal(bytes) { - terminal.send(bytes); +function releaseInputs(inputs) { + for (const input of inputs) { + admittedInputBytes -= input.payload.byteLength; + } + admittedInputFrames -= inputs.length; } + +function sendAck(input, accepted) { + if (terminal.readyState !== WebSocket.OPEN) return; + try { + terminal.send(encodeAck(input.inputId, input.generation, accepted)); + } catch { + closeSteering(); + } +} + +function closeRawOverflow() { + if (terminal.readyState < WebSocket.CLOSING) { + terminal.close(1009, "input backlog exceeded"); + } +} + +function decodeInput(data) { + if (!(data instanceof ArrayBuffer)) return null; + const frame = new Uint8Array(data); + if (frame.byteLength < 7 || !magic.every((value, index) => frame[index] === value)) { + return null; + } + if (frame[4] !== 0x05) return null; + const inputIdBytes = frame[5]; + if (!inputIdBytes || inputIdBytes > 80 || 6 + inputIdBytes > frame.byteLength) { + return null; + } + const inputId = inputIdDecoder.decode(frame.subarray(6, 6 + inputIdBytes)); + if (!/^[A-Za-z0-9_-]+$/.test(inputId)) return null; + const generationOffset = 6 + inputIdBytes; + const generationBytes = frame[generationOffset]; + if ( + !generationBytes || + generationBytes > 80 || + generationOffset + 1 + generationBytes > frame.byteLength + ) { + return null; + } + const generation = inputIdDecoder.decode( + frame.subarray(generationOffset + 1, generationOffset + 1 + generationBytes), + ); + if (!/^[A-Za-z0-9_-]+$/.test(generation)) return null; + return { + inputId, + generation, + payload: frame.subarray(generationOffset + 1 + generationBytes), + }; +} + +function encodeAck(inputId, generation, accepted) { + const inputIdBytes = encoder.encode(inputId); + const generationBytes = encoder.encode(generation); + const frame = new Uint8Array(8 + inputIdBytes.byteLength + generationBytes.byteLength); + frame.set(magic); + frame[4] = 0x06; + frame[5] = inputIdBytes.byteLength; + frame.set(inputIdBytes, 6); + const generationOffset = 6 + inputIdBytes.byteLength; + frame[generationOffset] = generationBytes.byteLength; + frame.set(generationBytes, generationOffset + 1); + frame[generationOffset + 1 + generationBytes.byteLength] = accepted ? 1 : 0; + return frame; +} + +function encodeUtf8Output(outputText) { + const payload = encoder.encode(outputText); + const frame = new Uint8Array(6 + payload.byteLength); + frame.set(magic); + frame[4] = 0x04; + frame[5] = 0; + frame.set(payload, 6); + return frame; +} + +function deactivateTerminal() { + if (terminalClosed) return; + terminalClosed = true; + activeGeneration = undefined; + rejectInputs(takePendingInputs(), 1001, "terminal closed"); + closeSteering(); +} + +terminal.addEventListener("close", deactivateTerminal); +terminal.addEventListener("error", deactivateTerminal); ``` +Set `CRABFLEET_RUNNER_PTY_URL` to the `runnerPtyUrl` returned by registration. +Implement `restricted-codex-steering.js` as the integration's narrow +`turn/steer` and `turn/interrupt` adapter. `deliverSteeringInput` must consume +browser input as steering instructions; it must never forward that input to a +shell or subprocess, and the adapter must not expose the GitHub Actions +environment. Await the steering acceptance signal before sending +`encodeAck(..., true)`. Do not acknowledge when the WebSocket merely queues the +input frame. This Node adapter buffers a valid incomplete UTF-8 suffix together +with every affected input ID. It delivers and positively acknowledges those +frames only after a later frame completes the sequence. Invalid UTF-8 rejects +the buffered group without delivering any of it. Every message is decoded and +admitted against the shared 16 KiB and 32-frame limits before it enters the +serialized delivery tail. Those counters retain ownership while input is +pending, queued, or blocked in `deliverSteeringInput`, so a stalled steering +call cannot retain an unbounded sequence of `MessageEvent` payloads. Framed +overflow receives a negative acknowledgement; raw overflow closes the socket +because legacy mode has no acknowledgement channel. An incomplete UTF-8 group +expires after one second. The first framed input pins the relay-owned generation +for that socket. Every admitted input rechecks both that generation and socket +liveness before entering the restricted steering handler; close or error +invalidates the generation and releases buffered or queued input instead of +delivering it through a replacement runner. + +WebSocket subprotocol selection is fixed during the opening handshake. There is +no capability message or mode transition after the socket opens. Older relays +that ignore the offered subprotocol leave `WebSocket.protocol` empty; the +adapter then receives raw input and sends raw string output. The +`runnerProtocol` query remains compatibility-only for already-deployed runners. +New runners must not add it, close, or reconnect solely because +`WebSocket.protocol` is empty: a query cannot confirm that the relay selected +framed I/O, while the existing socket is the required raw fallback. Each `CFR1` +frame occupies one binary WebSocket message. At the wire level, input and output +payloads are opaque terminal bytes. The example is deliberately a UTF-8 text +adapter for the integration's string-based steering surface: it rejects input +that is not complete valid UTF-8 and encodes each output string as UTF-8. +Deployments that require lossless arbitrary terminal bytes must use a +byte-oriented restricted steering adapter instead. + +| Offset | Size | Value | +| ------ | -------- | ------------------------------------------------------------------------------------------------------------------------------------ | +| 0 | 4 | ASCII `CFR1` | +| 4 | 1 | v2 `0x05` input, `0x06` acknowledgement, `0x07` lifecycle event, or shared `0x04` output | +| 5 | 1 | input ID byte length | +| 6 | variable | input ID, then one generation-length byte, the relay generation, and the type-specific payload; output omits the generation envelope | + +Input payloads are raw terminal bytes after the generation envelope. An +acknowledgement payload starts with the generation envelope, then `1` for +accepted or `0` for rejected, followed by optional UTF-8 error text. Lifecycle +events use an empty input ID, the generation envelope, and event code `0x01` +for runner connected, `0x02` for runner disconnected, or `0x03` for runner +waiting. Output uses an empty input ID followed by raw terminal bytes. + +The earlier `cfr1-framed-io-v1` mode remains accepted during rolling upgrades. +Its `0x01`, `0x02`, and `0x03` frames omit generations. The relay translates +between v1 and v2 at each socket boundary. A v2 viewer must send the generation +from its latest lifecycle event; stale-generation input is rejected before it +can reach the replacement runner. +The full wire contract is also specified in +[API](/api/#get-api-agent-interactive-sessions-id-runner-pty). + Properties: -- The URL is directly usable by Node's global `WebSocket`. - Authentication is the session-scoped `agentToken` query value. - Only one runner is current. - A new runner connection replaces the previous runner. - Multiple browser viewers may remain connected. -- Runner output is fanned out to viewers. -- Writable viewer input is sent to the current runner only. -- Runner lifecycle events are visible to viewers even while no runner is - connected. - -The relay transports raw terminal bytes. It does not interpret Codex JSON-RPC. -The runner-side integration decides how terminal input maps to model steering. +- Legacy runners open the returned URL unchanged, receive raw viewer input, and + send raw output. +- Framed runners offer the exact WebSocket subprotocol before connecting, + confirm its selection through `WebSocket.protocol`, and wrap every output + payload in a `0x04` frame. +- Generation-fenced viewers add `viewerProtocol=cfr1-framed-io-v2` before + connecting. They receive `CFR1` output plus relay-generated lifecycle and + acknowledgement frames regardless of the runner's mode. +- Existing v1 runners and viewers remain interoperable through relay-side + frame translation. +- Legacy viewers omit that query. They receive raw terminal output plus JSON + lifecycle and input-acknowledgement messages for compatibility. +- Negotiated input produces `input-accepted` only after the correlated runner + acknowledgement. A definitive negative acknowledgement produces + `input-rejected`. If the terminal hub's acknowledgement deadline expires + while the runner write may still be in flight, it produces + `input-delivery-unknown`, not + `input-rejected`, because that write may still complete. Legacy input reports + acceptance after relay delivery. The unknown-delivery JSON control event + carries `{"type":"input-delivery-unknown","error":"terminal input delivery outcome is unknown; the runner may still complete it"}`. +- Runner replacement also marks unresolved old-generation input as + `input-delivery-unknown`; a write that already entered the old PTY cannot be + proven absent. A runner-side PTY write that exceeds the bounded write deadline + retires that runner socket, so queued frames cannot execute behind a wedged + write. +- Framed viewer lifecycle events remain typed binary frames while no runner is + connected. Legacy viewers receive the JSON fallback. +- When runner and viewer modes differ, the relay wraps or unwraps terminal + output at the viewer boundary. + +The relay does not interpret Codex JSON-RPC. The runner-side integration decides +how accepted terminal input maps to model steering. ## Browser Attach @@ -325,13 +668,15 @@ instead of inventing a local shell. ## Steering Semantics -Crabfleet itself forwards terminal input bytes. In the ClawSweeper integration, -the runner: +Crabfleet forwards negotiated terminal input inside correlated `CFR1` frames. In the +ClawSweeper integration, the runner: -1. Collects printable input until Enter. -2. Echoes `[steer] ` to the terminal. -3. Calls Codex `turn/steer` with the active thread and expected turn ID. -4. Reports rejection or no-active-turn conditions in the terminal. +1. Accepts the framed bytes into its input handler and acknowledges that input + ID. +2. Collects printable input until Enter. +3. Echoes `[steer] ` to the terminal as UTF-8 terminal output. +4. Calls Codex `turn/steer` with the active thread and expected turn ID. +5. Reports rejection or no-active-turn conditions in the terminal. `Ctrl-C` maps to `turn/interrupt`. diff --git a/docs/macos-native-client.md b/docs/macos-native-client.md index 2bea01f6..b3d5b070 100644 --- a/docs/macos-native-client.md +++ b/docs/macos-native-client.md @@ -161,8 +161,10 @@ from the build, or obtain written provenance approval. bound to loopback behind an authenticated SSH tunnel; Share This Mac is the identity-gated tailnet exception. - The hardened prototype negotiates standard VNC password or no-auth security - only. ARD Diffie-Hellman, UltraVNC MS Logon II, Tight security, and TLS remain - disabled until their parsers and cryptography are replaced or fully tested. + only. The bundled ARD Diffie-Hellman path now requires a full-width + probabilistic safe-prime group and nonzero public/shared results, but ARD, + UltraVNC MS Logon II, Tight security, and TLS remain disabled in the app until + their complete interoperability surfaces are enabled and tested. - Password authentication uses a process-global DES key schedule. The fork serializes that path; replace it before concurrent password-auth sessions. - App-owned hosting shares one selected display at a time to a single client. diff --git a/docs/runs.md b/docs/runs.md index 75cede75..1a104345 100644 --- a/docs/runs.md +++ b/docs/runs.md @@ -120,8 +120,9 @@ Terminal contract: GitHub Actions PTY contract: - OpenClaw registers or resumes work through `POST /api/openclaw/action-sessions`. -- The returned `runnerPtyUrl` is a `wss:` URL with a rotated session-scoped query credential, directly usable by Node's global `WebSocket`. -- The Actions process connects outbound and sends raw terminal output bytes. Raw Ghostty input bytes are returned on the same socket. +- The returned `runnerPtyUrl` is a `wss:` URL with a rotated session-scoped query credential. Node's global `WebSocket` can open it without custom headers. +- Legacy runners open the returned URL unchanged and retain raw input/output with relay-level delivery reporting. +- Generation-fenced runners offer `cfr1-framed-io-v2` as a WebSocket subprotocol and use `CFR1` only when the upgrade response selects it. A relay that ignores the offer leaves `WebSocket.protocol` empty, so the runner remains raw-compatible during rolling upgrades. Viewer input and acknowledgements carry the relay-owned runner generation, stale input is rejected before forwarding, and the runner returns the matching generation and correlation ID only after its PTY accepts the write. The relay continues translating v1 framed and raw sockets. - `SessionControlDO` allows one current runner and multiple viewers. A new runner replaces the previous runner; viewers remain connected and receive runner lifecycle events. - Authorized browser viewers attach through the existing `/api/terminal/ws` hub. Service and agent credentials are never included in viewer responses. - The runner updates `state`, `phase`, `summary`, Codex thread/turn IDs, and heartbeat through the agent work-state endpoint. `completed`, `blocked`, `failed`, and `canceled` are terminal. diff --git a/docs/spec.md b/docs/spec.md index f03415f1..238b9f23 100644 --- a/docs/spec.md +++ b/docs/spec.md @@ -178,12 +178,19 @@ Crabfleet owns: - session identity and metadata; - rotating scoped agent token; -- outbound runner relay through `SessionControlDO`; +- outbound runner relay through `SessionControlDO`, preserving raw and v1 framed traffic while an explicitly selected v2 WebSocket subprotocol enables relay-generation-fenced binary `CFR1` input, acknowledgement, and lifecycle frames plus framed output; - browser terminal steering; - work-state heartbeats; - event and transcript finalization. -The Action remains the execution host and mutation authority. Ending the Crabfleet session does not cancel the workflow run. +The Action remains the execution host and mutation authority. A framed +runner acknowledges viewer input only after its PTY accepts the correlated +write; relay queueing is not acceptance. Legacy runners keep raw input/output +and relay-level delivery reporting. New runners offer the v2 WebSocket +subprotocol and switch formats only when the relay selects it in the upgrade +response. Relays that ignore the offer therefore remain raw-compatible, with no +in-band handshake or mode-transition race. Ending the Crabfleet session does +not cancel the workflow run. ## Session Lifecycle diff --git a/internal/fleetapi/client.go b/internal/fleetapi/client.go index 5e53ed5a..bd12cd2d 100644 --- a/internal/fleetapi/client.go +++ b/internal/fleetapi/client.go @@ -10,6 +10,7 @@ import ( "net/http" "net/url" "strings" + "time" "github.com/openclaw/crabfleet/internal/terminalws" ) @@ -18,6 +19,7 @@ type TerminalSize = terminalws.Size const maxResponseBytes = 4 * 1024 * 1024 const maxErrorBytes = 512 +const terminalInputConfirmationTimeout = 15 * time.Second var ErrMissingAuth = errors.New("API mode requires SSH gateway token + fingerprint or agent token + session ID") @@ -193,7 +195,9 @@ func (c *Client) Message( if enter { message += "\n" } - return client.SendInput(ctx, []byte(message)) + confirmationContext, cancel := context.WithTimeout(ctx, terminalInputConfirmationTimeout) + defer cancel() + return client.SendInputConfirmed(confirmationContext, []byte(message)) } func (c *Client) Attach( @@ -222,9 +226,10 @@ func (c *Client) terminal(ctx context.Context, id string, cols uint32, rows uint return nil, err } client, err := terminalws.Dial(ctx, endpoint, id, terminalws.Options{ - Header: headers, - Cols: cols, - Rows: rows, + HTTPClient: c.http, + Header: headers, + Cols: cols, + Rows: rows, }) if err != nil { var statusErr *terminalws.HandshakeStatusError diff --git a/internal/fleetapi/client_test.go b/internal/fleetapi/client_test.go index 7bb54775..95ba5e6b 100644 --- a/internal/fleetapi/client_test.go +++ b/internal/fleetapi/client_test.go @@ -6,6 +6,7 @@ import ( "net/http/httptest" "strings" "testing" + "time" ) func TestClientUsesSSHAuthentication(t *testing.T) { @@ -66,6 +67,12 @@ func TestClientRejectsIncompleteAuthentication(t *testing.T) { } } +func TestTerminalInputConfirmationTimeoutIsBounded(t *testing.T) { + if terminalInputConfirmationTimeout <= 0 || terminalInputConfirmationTimeout > 30*time.Second { + t.Fatalf("terminal input confirmation timeout = %s", terminalInputConfirmationTimeout) + } +} + func TestClientSanitizesStatusErrorBody(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusInternalServerError) diff --git a/internal/terminalws/client.go b/internal/terminalws/client.go index bc84fdc4..afe79210 100644 --- a/internal/terminalws/client.go +++ b/internal/terminalws/client.go @@ -12,6 +12,7 @@ import ( "strings" "sync" "sync/atomic" + "time" "github.com/coder/websocket" ) @@ -22,6 +23,10 @@ const ( maxFrameBytes = 16 * 1024 * 1024 maxErrorBytes = 512 + defaultInputConfirmationTimeout = 5 * time.Second + defaultAttachmentShutdownTimeout = 250 * time.Millisecond + defaultOutputAckTimeout = 5 * time.Second + messageHello = 1 messageWelcome = 2 messageSubscribe = 10 @@ -48,10 +53,15 @@ const ( subscribeOutputAcknowledgements = 1 << 3 ) +var ErrInputDeliveryUnknown = errors.New( + "terminal input delivery outcome is unknown; the runner may still complete it", +) + type Options struct { - Header http.Header - Cols uint32 - Rows uint32 + HTTPClient *http.Client + Header http.Header + Cols uint32 + Rows uint32 } type HandshakeStatusError struct { @@ -73,11 +83,27 @@ type Size struct { } type Client struct { - conn *websocket.Conn - sessionID string - canInput atomic.Bool - lastSize atomic.Uint64 - writeMu sync.Mutex + conn *websocket.Conn + sessionID string + supportsInputAcknowledgement bool + cancel context.CancelFunc + readCancel context.CancelFunc + canInput atomic.Bool + lastSize atomic.Uint64 + writeMu sync.Mutex + confirmOnce sync.Once + confirmGate chan struct{} + confirmationTimeout time.Duration + attachmentShutdownTimeout time.Duration + stateMu sync.Mutex + inputWaiter chan error + attachment *terminalAttachment + attachmentReady chan struct{} + controlGrantGeneration uint64 + handledControlGrant uint64 + terminalErr error + readerDone chan struct{} + readerErr error } type frame struct { @@ -86,6 +112,30 @@ type frame struct { payload []byte } +type terminalAttachment struct { + frames chan attachmentDelivery + done chan struct{} + controlGrantGeneration uint64 +} + +type attachmentDelivery struct { + frame frame + accepted chan bool + controlGrantGeneration uint64 +} + +type inputConfirmationInterruptedError struct { + cause error +} + +func (e *inputConfirmationInterruptedError) Error() string { + return e.cause.Error() +} + +func (e *inputConfirmationInterruptedError) Unwrap() error { + return e.cause +} + type eventPayload struct { Type string `json:"type"` Error string `json:"error"` @@ -93,6 +143,10 @@ type eventPayload struct { CanInput bool `json:"canInput"` } +type welcomePayload struct { + InputAcknowledgements bool `json:"inputAcknowledgements"` +} + type readCanceler interface { CancelRead() error } @@ -122,10 +176,33 @@ func Dial(ctx context.Context, endpoint string, sessionID string, options Option if sessionID == "" { return nil, errors.New("terminal session id is required") } + httpClient := options.HTTPClient + var setupCancel context.CancelFunc + var setupTimer *time.Timer + var setupFinished atomic.Bool + if options.HTTPClient != nil && options.HTTPClient.Timeout > 0 { + cloned := *options.HTTPClient + cloned.Timeout = 0 + httpClient = &cloned + ctx, setupCancel = context.WithCancel(ctx) + setupTimer = time.AfterFunc(options.HTTPClient.Timeout, func() { + if setupFinished.CompareAndSwap(false, true) { + setupCancel() + } + }) + } conn, resp, err := websocket.Dial(ctx, endpoint, &websocket.DialOptions{ + HTTPClient: httpClient, HTTPHeader: options.Header, }) if err != nil { + setupFinished.Store(true) + if setupTimer != nil { + setupTimer.Stop() + } + if setupCancel != nil { + setupCancel() + } if resp != nil { body := "" if resp.Body != nil { @@ -138,9 +215,21 @@ func Dial(ctx context.Context, endpoint string, sessionID string, options Option return nil, err } conn.SetReadLimit(maxFrameBytes) - client := &Client{conn: conn, sessionID: sessionID} + client := &Client{ + conn: conn, + sessionID: sessionID, + cancel: setupCancel, + confirmationTimeout: defaultInputConfirmationTimeout, + } client.rememberSize(Size{Cols: options.Cols, Rows: options.Rows}) closeWithError := func(err error) (*Client, error) { + setupFinished.Store(true) + if setupTimer != nil { + setupTimer.Stop() + } + if setupCancel != nil { + setupCancel() + } _ = conn.Close(websocket.StatusInternalError, "terminal setup failed") return nil, err } @@ -164,6 +253,10 @@ func Dial(ctx context.Context, endpoint string, sessionID string, options Option } switch current.messageType { case messageWelcome: + var welcome welcomePayload + if json.Unmarshal(current.payload, &welcome) == nil { + client.supportsInputAcknowledgement = welcome.InputAcknowledgements + } continue case messageError, messageControlRevoked: return closeWithError(frameError(current, "terminal subscription failed")) @@ -173,7 +266,14 @@ func Dial(ctx context.Context, endpoint string, sessionID string, options Option return closeWithError(fmt.Errorf("decode terminal event: %w", err)) } if event.Type == "subscribed" { + if setupTimer != nil && !setupFinished.CompareAndSwap(false, true) { + return closeWithError(errors.New("terminal subscription timed out")) + } client.canInput.Store(event.CanInput) + if setupTimer != nil { + setupTimer.Stop() + } + client.startReader() return client, nil } if event.Type == "closed" { @@ -184,13 +284,45 @@ func Dial(ctx context.Context, endpoint string, sessionID string, options Option } func (c *Client) Close() error { - return c.conn.Close(websocket.StatusNormalClosure, "") + if c.readCancel != nil { + c.readCancel() + } + err := c.conn.Close(websocket.StatusNormalClosure, "") + if c.cancel != nil { + c.cancel() + } + return err } func (c *Client) SendInput(ctx context.Context, payload []byte) error { if len(payload) == 0 { return nil } + if !c.canInput.Load() { + return errors.New("terminal control has not been granted") + } + if !c.supportsInputAcknowledgement { + return c.writeInput(ctx, payload) + } + if err := c.acquireConfirmation(ctx); err != nil { + return err + } + + waiter := make(chan error, 1) + if err := c.registerInputWaiter(waiter); err != nil { + c.releaseConfirmation() + return err + } + if err := c.writeInput(ctx, payload); err != nil { + c.clearInputWaiter(waiter) + c.releaseConfirmation() + return err + } + go c.drainInputConfirmation(waiter) + return nil +} + +func (c *Client) writeInput(ctx context.Context, payload []byte) error { if !c.canInput.Load() { return errors.New("terminal control has not been granted") } @@ -201,6 +333,110 @@ func (c *Client) SendInput(ctx context.Context, payload []byte) error { }) } +func (c *Client) SendInputConfirmed(ctx context.Context, payload []byte) error { + if len(payload) == 0 { + return nil + } + if !c.supportsInputAcknowledgement { + return c.SendInput(ctx, payload) + } + if err := c.acquireConfirmation(ctx); err != nil { + return err + } + defer c.releaseConfirmation() + + waiter := make(chan error, 1) + if err := c.registerInputWaiter(waiter); err != nil { + return err + } + if err := c.writeInput(ctx, payload); err != nil { + c.clearInputWaiter(waiter) + return err + } + err := c.waitForInputConfirmation(ctx, waiter) + var interrupted *inputConfirmationInterruptedError + if errors.As(err, &interrupted) { + return inputDeliveryUnknownCause(interrupted.cause) + } + return err +} + +func (c *Client) drainInputConfirmation(waiter chan error) { + defer c.releaseConfirmation() + ctx, cancel := context.WithTimeout(context.Background(), c.inputConfirmationTimeout()) + defer cancel() + _ = c.waitForInputConfirmation(ctx, waiter) +} + +func (c *Client) waitForInputConfirmation(ctx context.Context, waiter chan error) error { + select { + case err := <-waiter: + return err + case <-c.readerDone: + select { + case err := <-waiter: + return err + default: + } + if !c.clearInputWaiter(waiter) { + return <-waiter + } + return &inputConfirmationInterruptedError{ + cause: readerUnavailableError(c.readerError()), + } + case <-ctx.Done(): + select { + case err := <-waiter: + return err + default: + } + if !c.clearInputWaiter(waiter) { + return <-waiter + } + c.closeNow() + return &inputConfirmationInterruptedError{cause: ctx.Err()} + } +} + +func (c *Client) closeNow() { + if c.readCancel != nil { + c.readCancel() + } + if c.conn != nil { + _ = c.conn.CloseNow() + } + if c.cancel != nil { + c.cancel() + } +} + +func (c *Client) acquireConfirmation(ctx context.Context) error { + c.confirmOnce.Do(func() { + c.confirmGate = make(chan struct{}, 1) + }) + select { + case c.confirmGate <- struct{}{}: + if err := ctx.Err(); err != nil { + <-c.confirmGate + return err + } + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (c *Client) releaseConfirmation() { + <-c.confirmGate +} + +func (c *Client) inputConfirmationTimeout() time.Duration { + if c.confirmationTimeout > 0 { + return c.confirmationTimeout + } + return defaultInputConfirmationTimeout +} + func (c *Client) Resize(ctx context.Context, size Size) error { if size.Cols == 0 || size.Rows == 0 { return nil @@ -220,6 +456,12 @@ func (c *Client) Attach(ctx context.Context, terminal io.ReadWriter, resizes <-c ctx, cancel := context.WithCancel(ctx) defer cancel() + attachment, err := c.registerAttachment() + if err != nil { + return err + } + defer c.clearAttachment(attachment) + var wg sync.WaitGroup canceler, cancelableRead := terminal.(readCanceler) cancelRead := func() { @@ -230,6 +472,7 @@ func (c *Client) Attach(ctx context.Context, terminal io.ReadWriter, resizes <-c } errCh := make(chan error, 3) + frameConsumerDone := make(chan struct{}) wg.Add(1) go func() { defer wg.Done() @@ -237,7 +480,13 @@ func (c *Client) Attach(ctx context.Context, terminal io.ReadWriter, resizes <-c for { count, err := terminal.Read(buffer) if count > 0 && c.canInput.Load() { - if writeErr := c.SendInput(ctx, buffer[:count]); writeErr != nil { + confirmationCtx, confirmationCancel := context.WithTimeout( + ctx, + c.inputConfirmationTimeout(), + ) + writeErr := c.SendInputConfirmed(confirmationCtx, buffer[:count]) + confirmationCancel() + if writeErr != nil { errCh <- writeErr return } @@ -274,64 +523,99 @@ func (c *Client) Attach(ctx context.Context, terminal io.ReadWriter, resizes <-c wg.Add(1) go func() { defer wg.Done() - for { - current, err := c.read(ctx) - if err != nil { + defer close(frameConsumerDone) + if generation := attachment.controlGrantGeneration; generation > 0 { + if err := c.resendRememberedSize(ctx, generation); err != nil { errCh <- err return } - if current.sessionID != "" && current.sessionID != c.sessionID { - continue - } - switch current.messageType { - case messageOutput: - if _, err := terminal.Write(current.payload); err != nil { - errCh <- err - return - } - if err := c.write(ctx, frame{ - messageType: messageAck, - sessionID: c.sessionID, - payload: ackPayload(uint32(len(current.payload))), - }); err != nil { - errCh <- err + } + for { + select { + case <-ctx.Done(): + return + case <-c.readerDone: + errCh <- c.readerError() + return + case delivery := <-attachment.frames: + if !c.acceptAttachmentDelivery(attachment, delivery) { return } - case messageError: - errCh <- frameError(current, "terminal connection failed") - return - case messageControlRevoked: - c.canInput.Store(false) - case messageControlGranted: - c.canInput.Store(true) - if size := c.rememberedSize(); size.Cols > 0 && size.Rows > 0 { - if err := c.write(ctx, frame{ - messageType: messageResize, + current := delivery.frame + switch current.messageType { + case messageOutput: + if _, err := terminal.Write(current.payload); err != nil { + c.retireConnection(err) + errCh <- err + return + } + ackCtx, ackCancel := context.WithTimeout( + context.Background(), + defaultOutputAckTimeout, + ) + err := c.write(ackCtx, frame{ + messageType: messageAck, sessionID: c.sessionID, - payload: resizePayload(size), - }); err != nil { + payload: ackPayload(uint32(len(current.payload))), + }) + ackCancel() + if err != nil { errCh <- err return } - } - case messageEvent: - var event eventPayload - if json.Unmarshal(current.payload, &event) == nil && event.Type == "closed" { - errCh <- nil + case messageError: + errCh <- frameError(current, "terminal connection failed") return + case messageControlGranted: + if err := c.resendRememberedSize( + ctx, + delivery.controlGrantGeneration, + ); err != nil { + errCh <- err + return + } + case messageEvent: + var event eventPayload + if json.Unmarshal(current.payload, &event) == nil && event.Type == "closed" { + errCh <- nil + return + } } } } }() - err := <-errCh + select { + case err = <-errCh: + case <-ctx.Done(): + err = ctx.Err() + } cancelRead() - if cancelableRead { + // Let completed writes preserve their acknowledgement ordering, but retire the + // attachment if its owner must close the terminal to unblock a write. + shutdownTimer := time.NewTimer(c.frameConsumerShutdownTimeout()) + frameConsumerStopped := false + select { + case <-frameConsumerDone: + frameConsumerStopped = true + if !shutdownTimer.Stop() { + <-shutdownTimer.C + } + case <-shutdownTimer.C: + } + if cancelableRead && frameConsumerStopped { wg.Wait() } return normalizeCloseError(err) } +func (c *Client) frameConsumerShutdownTimeout() time.Duration { + if c.attachmentShutdownTimeout > 0 { + return c.attachmentShutdownTimeout + } + return defaultAttachmentShutdownTimeout +} + func (c *Client) rememberSize(size Size) { c.lastSize.Store(uint64(size.Cols)<<32 | uint64(size.Rows)) } @@ -341,12 +625,339 @@ func (c *Client) rememberedSize() Size { return Size{Cols: uint32(value >> 32), Rows: uint32(value)} } +func (c *Client) resendRememberedSize(ctx context.Context, generation uint64) error { + size := c.rememberedSize() + if size.Cols > 0 && size.Rows > 0 { + if err := c.write(ctx, frame{ + messageType: messageResize, + sessionID: c.sessionID, + payload: resizePayload(size), + }); err != nil { + return err + } + } + c.stateMu.Lock() + if generation > c.handledControlGrant { + c.handledControlGrant = generation + } + c.stateMu.Unlock() + return nil +} + func (c *Client) write(ctx context.Context, current frame) error { c.writeMu.Lock() defer c.writeMu.Unlock() return c.conn.Write(ctx, websocket.MessageBinary, encodeFrame(current)) } +func (c *Client) startReader() { + ctx, cancel := context.WithCancel(context.Background()) + c.readCancel = cancel + c.stateMu.Lock() + c.readerDone = make(chan struct{}) + c.attachmentReady = make(chan struct{}) + c.stateMu.Unlock() + go c.readLoop(ctx) +} + +func (c *Client) readLoop(ctx context.Context) { + for { + current, err := c.read(ctx) + if err != nil { + c.finishReader(err) + return + } + if err := c.handleFrame(ctx, current); err != nil { + c.finishReader(err) + return + } + } +} + +func (c *Client) handleFrame(ctx context.Context, current frame) error { + if current.sessionID != "" && current.sessionID != c.sessionID { + return nil + } + switch current.messageType { + case messageOutput: + return c.deliverOrQueueOutput(ctx, current) + case messageError: + err := frameError(current, "terminal connection failed") + c.canInput.Store(false) + c.completeInput(err) + return err + case messageControlRevoked: + c.canInput.Store(false) + c.stateMu.Lock() + c.handledControlGrant = c.controlGrantGeneration + c.stateMu.Unlock() + c.deliverAttachment(ctx, current) + case messageControlGranted: + c.canInput.Store(true) + c.stateMu.Lock() + c.controlGrantGeneration++ + generation := c.controlGrantGeneration + attachment := c.attachment + c.stateMu.Unlock() + if attachment != nil { + c.deliverControlGranted(ctx, attachment, current, generation) + } + case messageEvent: + var event eventPayload + if err := json.Unmarshal(current.payload, &event); err != nil { + return fmt.Errorf("decode terminal event: %w", err) + } + switch event.Type { + case "subscribed": + c.canInput.Store(event.CanInput) + case "input-accepted": + c.completeInput(nil) + case "input-rejected": + c.completeInput(frameError(current, "terminal input rejected")) + case "input-delivery-unknown": + c.completeInput(inputDeliveryUnknownError(current)) + case "closed": + c.canInput.Store(false) + c.markTerminalClosed(errors.New("terminal closed")) + c.completeInput(errors.New("terminal closed before accepting input")) + c.deliverAttachment(ctx, current) + default: + c.deliverAttachment(ctx, current) + } + default: + c.deliverAttachment(ctx, current) + } + return nil +} + +func (c *Client) registerInputWaiter(waiter chan error) error { + c.stateMu.Lock() + defer c.stateMu.Unlock() + select { + case <-c.readerDone: + return readerUnavailableError(c.readerErr) + default: + } + if c.terminalErr != nil { + return c.terminalErr + } + c.inputWaiter = waiter + if c.attachment == nil && c.attachmentReady != nil { + close(c.attachmentReady) + c.attachmentReady = make(chan struct{}) + } + return nil +} + +func (c *Client) clearInputWaiter(waiter chan error) bool { + c.stateMu.Lock() + defer c.stateMu.Unlock() + if c.inputWaiter == waiter { + c.inputWaiter = nil + return true + } + return false +} + +func (c *Client) completeInput(err error) { + c.stateMu.Lock() + waiter := c.inputWaiter + c.inputWaiter = nil + c.stateMu.Unlock() + if waiter != nil { + waiter <- err + } +} + +func (c *Client) registerAttachment() (*terminalAttachment, error) { + c.stateMu.Lock() + defer c.stateMu.Unlock() + if c.attachment != nil { + return nil, errors.New("terminal client is already attached") + } + if c.terminalErr != nil { + return nil, c.terminalErr + } + select { + case <-c.readerDone: + return nil, readerUnavailableError(c.readerErr) + default: + } + attachment := &terminalAttachment{ + frames: make(chan attachmentDelivery), + done: make(chan struct{}), + } + if c.handledControlGrant < c.controlGrantGeneration { + attachment.controlGrantGeneration = c.controlGrantGeneration + } + c.attachment = attachment + if c.attachmentReady != nil { + close(c.attachmentReady) + c.attachmentReady = nil + } + return attachment, nil +} + +func (c *Client) clearAttachment(attachment *terminalAttachment) { + c.stateMu.Lock() + if c.attachment == attachment { + c.attachment = nil + close(attachment.done) + select { + case <-c.readerDone: + default: + c.attachmentReady = make(chan struct{}) + } + } + c.stateMu.Unlock() +} + +func (c *Client) markTerminalClosed(err error) { + c.stateMu.Lock() + if c.terminalErr == nil { + c.terminalErr = err + } + c.stateMu.Unlock() +} + +func (c *Client) retireConnection(err error) { + c.markTerminalClosed(err) + c.closeNow() +} + +func (c *Client) deliverOrQueueOutput(ctx context.Context, current frame) error { + for { + c.stateMu.Lock() + attachment := c.attachment + ready := c.attachmentReady + discardOutput := c.inputWaiter != nil + c.stateMu.Unlock() + if attachment == nil { + if discardOutput { + return c.write(ctx, frame{ + messageType: messageAck, + sessionID: c.sessionID, + payload: ackPayload(uint32(len(current.payload))), + }) + } + if ready == nil { + return nil + } + select { + case <-ready: + continue + case <-ctx.Done(): + return ctx.Err() + } + } + if c.deliverToAttachment(ctx, attachment, current) { + return nil + } + if err := ctx.Err(); err != nil { + return ctx.Err() + } + } +} + +func (c *Client) deliverAttachment(ctx context.Context, current frame) bool { + c.stateMu.Lock() + attachment := c.attachment + c.stateMu.Unlock() + if attachment == nil { + return false + } + return c.deliverToAttachment(ctx, attachment, current) +} + +func (c *Client) deliverToAttachment( + ctx context.Context, + attachment *terminalAttachment, + current frame, +) bool { + return c.deliverToAttachmentWithControlGrant(ctx, attachment, current, 0) +} + +func (c *Client) deliverControlGranted( + ctx context.Context, + attachment *terminalAttachment, + current frame, + generation uint64, +) bool { + return c.deliverToAttachmentWithControlGrant( + ctx, + attachment, + current, + generation, + ) +} + +func (c *Client) deliverToAttachmentWithControlGrant( + ctx context.Context, + attachment *terminalAttachment, + current frame, + generation uint64, +) bool { + delivery := attachmentDelivery{ + frame: current, + accepted: make(chan bool, 1), + controlGrantGeneration: generation, + } + select { + case attachment.frames <- delivery: + case <-attachment.done: + return false + case <-ctx.Done(): + return false + } + select { + case accepted := <-delivery.accepted: + return accepted + case <-ctx.Done(): + return false + } +} + +func (c *Client) acceptAttachmentDelivery( + attachment *terminalAttachment, + delivery attachmentDelivery, +) bool { + c.stateMu.Lock() + accepted := c.attachment == attachment + c.stateMu.Unlock() + delivery.accepted <- accepted + return accepted +} + +func (c *Client) finishReader(err error) { + c.stateMu.Lock() + c.readerErr = normalizeCloseError(err) + if c.terminalErr != nil { + c.readerErr = c.terminalErr + } + waiter := c.inputWaiter + c.inputWaiter = nil + if waiter != nil { + waiter <- &inputConfirmationInterruptedError{ + cause: readerUnavailableError(c.readerErr), + } + } + close(c.readerDone) + c.stateMu.Unlock() +} + +func (c *Client) readerError() error { + c.stateMu.Lock() + defer c.stateMu.Unlock() + return c.readerErr +} + +func readerUnavailableError(err error) error { + if err != nil { + return err + } + return errors.New("terminal connection closed") +} + func (c *Client) read(ctx context.Context) (frame, error) { messageType, payload, err := c.conn.Read(ctx) if err != nil { @@ -433,6 +1044,21 @@ func frameError(current frame, fallback string) error { return errors.New(fallback) } +func inputDeliveryUnknownError(current frame) error { + detail := frameError(current, ErrInputDeliveryUnknown.Error()) + if detail.Error() == ErrInputDeliveryUnknown.Error() { + return ErrInputDeliveryUnknown + } + return fmt.Errorf("%w: %s", ErrInputDeliveryUnknown, detail) +} + +func inputDeliveryUnknownCause(cause error) error { + if cause == nil { + return ErrInputDeliveryUnknown + } + return fmt.Errorf("%w: %w", ErrInputDeliveryUnknown, cause) +} + func normalizeCloseError(err error) error { if err == nil { return nil diff --git a/internal/terminalws/client_test.go b/internal/terminalws/client_test.go index f1e04cc6..4fc37e94 100644 --- a/internal/terminalws/client_test.go +++ b/internal/terminalws/client_test.go @@ -6,10 +6,13 @@ import ( "encoding/binary" "encoding/hex" "encoding/json" + "errors" + "fmt" "io" "net/http" "net/http/httptest" "os" + "strings" "sync" "testing" "time" @@ -274,6 +277,1581 @@ func TestClientSubscribesSendsInputAndAcknowledgesOutput(t *testing.T) { } } +func TestClientDefersOutputAcknowledgementUntilAttach(t *testing.T) { + acknowledged := make(chan uint32, 1) + outputSent := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: false}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-before-attach", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-before-attach", + payload: []byte("early output\n"), + })); err != nil { + t.Error(err) + return + } + close(outputSent) + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + ack, err := decodeFrame(payload) + if err != nil || ack.messageType != messageAck { + t.Errorf("output acknowledgement = %#v, %v", ack, err) + return + } + acknowledged <- binary.LittleEndian.Uint32(ack.payload) + closed, _ := json.Marshal(eventPayload{Type: "closed"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-before-attach", + payload: closed, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-before-attach", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + <-outputSent + select { + case bytes := <-acknowledged: + t.Fatalf("acknowledged %d bytes before attach", bytes) + default: + } + + terminal := newBlockingTerminal() + if err := client.Attach(context.Background(), terminal, nil); err != nil { + t.Fatal(err) + } + if terminal.String() != "early output\n" { + t.Fatalf("output = %q", terminal.String()) + } + if bytes := <-acknowledged; bytes != uint32(len("early output\n")) { + t.Fatalf("acknowledged = %d", bytes) + } +} + +func TestSendInputConfirmedDoesNotTreatControlRevocationAsInputRejection(t *testing.T) { + revokedSent := make(chan struct{}) + releaseAcceptance := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-revoked", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + revoked, _ := json.Marshal(eventPayload{Error: "terminal control revoked"}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageControlRevoked, + sessionID: "IS-revoked", + payload: revoked, + })); err != nil { + t.Error(err) + return + } + close(revokedSent) + select { + case <-releaseAcceptance: + case <-r.Context().Done(): + return + } + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-revoked", + payload: accepted, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-revoked", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + done := make(chan error, 1) + go func() { + done <- client.SendInputConfirmed(context.Background(), []byte("forwarded\n")) + }() + <-revokedSent + select { + case err := <-done: + t.Fatalf("control revocation completed forwarded input: %v", err) + case <-time.After(50 * time.Millisecond): + } + if client.canInput.Load() { + t.Fatal("control revocation did not remove input capability") + } + close(releaseAcceptance) + if err := <-done; err != nil { + t.Fatalf("later input acceptance = %v", err) + } +} + +func TestSendInputConfirmedReportsUnknownDelivery(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-delivery-unknown", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + unknown, _ := json.Marshal(eventPayload{ + Type: "input-delivery-unknown", + Error: ErrInputDeliveryUnknown.Error(), + }) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-delivery-unknown", + payload: unknown, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-delivery-unknown", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + err = client.SendInputConfirmed(context.Background(), []byte("possibly-delivered\n")) + if !errors.Is(err, ErrInputDeliveryUnknown) { + t.Fatalf("error = %v", err) + } + if err.Error() != ErrInputDeliveryUnknown.Error() { + t.Fatalf("error text = %q", err) + } + if !client.canInput.Load() { + t.Fatal("ambiguous delivery revoked input capability") + } +} + +func TestSendInputConfirmedAcknowledgesOutputWithoutAttachment(t *testing.T) { + acknowledged := make(chan uint32, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-message", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-message", + payload: []byte("prompt\n"), + })); err != nil { + t.Error(err) + return + } + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + ack, err := decodeFrame(payload) + if err != nil || ack.messageType != messageAck { + t.Errorf("output acknowledgement = %#v, %v", ack, err) + return + } + acknowledged <- binary.LittleEndian.Uint32(ack.payload) + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-message", + payload: accepted, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-message", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := client.SendInputConfirmed(ctx, []byte("echo ready\n")); err != nil { + t.Fatal(err) + } + if bytes := <-acknowledged; bytes != uint32(len("prompt\n")) { + t.Fatalf("acknowledged = %d", bytes) + } +} + +func TestSendInputConfirmedWakesOutputWaitingForAttachment(t *testing.T) { + outputSent := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-output-first", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-output-first", + payload: []byte("prompt\n"), + })); err != nil { + t.Error(err) + return + } + close(outputSent) + + seenInput := false + seenAcknowledgement := false + for !seenInput || !seenAcknowledgement { + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + current, err := decodeFrame(payload) + if err != nil { + t.Error(err) + return + } + switch current.messageType { + case messageInput: + seenInput = true + case messageAck: + seenAcknowledgement = true + default: + t.Errorf("message type = %d", current.messageType) + return + } + } + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-output-first", + payload: accepted, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-output-first", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + <-outputSent + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := client.SendInputConfirmed(ctx, []byte("echo ready\n")); err != nil { + t.Fatal(err) + } +} + +func TestSendInputConfirmedFailsWhenConnectionClosesBeforeAcknowledgement(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-close-before-ack", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + _ = conn.Close(websocket.StatusInternalError, "reader lost after input write") + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-close-before-ack", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + err = client.SendInputConfirmed(ctx, []byte("echo ready\n")) + if !errors.Is(err, ErrInputDeliveryUnknown) { + t.Fatalf("error = %v", err) + } + var closeError websocket.CloseError + if !errors.As(err, &closeError) { + t.Fatalf("error does not preserve websocket close: %v", err) + } + if closeError.Code != websocket.StatusInternalError { + t.Fatalf("close code = %v", closeError.Code) + } +} + +func TestSendInputConfirmedKeepsPreWriteReaderLossOrdinary(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-close-before-write", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + _ = conn.Close(websocket.StatusPolicyViolation, "reader lost before input write") + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-close-before-write", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + <-client.readerDone + + err = client.SendInputConfirmed(context.Background(), []byte("never-written\n")) + if errors.Is(err, ErrInputDeliveryUnknown) { + t.Fatalf("pre-write error became ambiguous: %v", err) + } + var closeError websocket.CloseError + if !errors.As(err, &closeError) { + t.Fatalf("error does not preserve websocket close: %v", err) + } + if closeError.Code != websocket.StatusPolicyViolation { + t.Fatalf("close code = %v", closeError.Code) + } +} + +func TestSendInputConfirmedPrefersAcceptanceBeforeReaderClose(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-accepted-before-close", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-accepted-before-close", + payload: accepted, + })); err != nil { + t.Error(err) + return + } + _ = conn.Close(websocket.StatusNormalClosure, "") + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-accepted-before-close", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + if err := client.SendInputConfirmed(context.Background(), []byte("echo ready\n")); err != nil { + t.Fatal(err) + } +} + +func TestAttachRejectsSessionClosedBeforeAttachment(t *testing.T) { + closedSent := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-closed-before-attach", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + closed, _ := json.Marshal(eventPayload{Type: "closed"}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-closed-before-attach", + payload: closed, + })); err != nil { + t.Error(err) + return + } + close(closedSent) + <-r.Context().Done() + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-closed-before-attach", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + <-closedSent + deadline := time.Now().Add(time.Second) + for { + client.stateMu.Lock() + terminalErr := client.terminalErr + client.stateMu.Unlock() + if terminalErr != nil { + break + } + if time.Now().After(deadline) { + t.Fatal("terminal close was not retained") + } + time.Sleep(time.Millisecond) + } + + err = client.Attach(context.Background(), newBlockingTerminal(), nil) + if err == nil || !strings.Contains(err.Error(), "terminal closed") { + t.Fatalf("error = %v", err) + } +} + +func TestRetiredAttachmentCannotAcceptBufferedFrames(t *testing.T) { + attachment := &terminalAttachment{ + frames: make(chan attachmentDelivery), + done: make(chan struct{}), + } + close(attachment.done) + client := &Client{attachment: attachment} + + for range 100 { + if client.deliverAttachment(context.Background(), frame{messageType: messageOutput}) { + t.Fatal("retired attachment accepted output") + } + } +} + +func TestRetiredAttachmentRejectsReceivedDelivery(t *testing.T) { + oldAttachment := &terminalAttachment{ + frames: make(chan attachmentDelivery), + done: make(chan struct{}), + } + replacement := &terminalAttachment{ + frames: make(chan attachmentDelivery), + done: make(chan struct{}), + } + client := &Client{attachment: replacement} + close(oldAttachment.done) + + for range 1_000 { + delivery := attachmentDelivery{ + frame: frame{messageType: messageOutput, payload: []byte("replacement output\n")}, + accepted: make(chan bool, 1), + } + if client.acceptAttachmentDelivery(oldAttachment, delivery) { + t.Fatal("retired attachment accepted a received delivery") + } + if accepted := <-delivery.accepted; accepted { + t.Fatal("retired attachment acknowledged a received delivery") + } + } +} + +func TestStaleAttachmentDeliveryRetriesReplacement(t *testing.T) { + oldAttachment := &terminalAttachment{ + frames: make(chan attachmentDelivery), + done: make(chan struct{}), + } + replacement := &terminalAttachment{ + frames: make(chan attachmentDelivery), + done: make(chan struct{}), + } + client := &Client{ + attachment: oldAttachment, + attachmentReady: make(chan struct{}), + } + + staleCaptured := make(chan struct{}) + releaseStale := make(chan struct{}) + go func() { + delivery := <-oldAttachment.frames + close(staleCaptured) + <-releaseStale + client.acceptAttachmentDelivery(oldAttachment, delivery) + }() + + delivered := make(chan error, 1) + go func() { + delivered <- client.deliverOrQueueOutput(context.Background(), frame{ + messageType: messageOutput, + payload: []byte("replacement output\n"), + }) + }() + <-staleCaptured + + client.stateMu.Lock() + client.attachment = replacement + close(oldAttachment.done) + client.stateMu.Unlock() + + replacementReceived := make(chan frame, 1) + go func() { + delivery := <-replacement.frames + if client.acceptAttachmentDelivery(replacement, delivery) { + replacementReceived <- delivery.frame + } + }() + close(releaseStale) + + select { + case err := <-delivered: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("stale delivery did not retry the replacement attachment") + } + select { + case current := <-replacementReceived: + if got := string(current.payload); got != "replacement output\n" { + t.Fatalf("replacement payload = %q", got) + } + case <-time.After(time.Second): + t.Fatal("replacement attachment did not receive retried output") + } +} + +func TestSendInputConfirmedRejectionDoesNotRevokeControl(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-request-scoped", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + + for index := range 2 { + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + current, err := decodeFrame(payload) + if err != nil { + t.Error(err) + return + } + if current.messageType != messageInput { + t.Errorf("message type = %d", current.messageType) + return + } + event := eventPayload{Type: "input-accepted"} + if index == 0 { + event = eventPayload{ + Type: "input-rejected", + Error: "runner rejected terminal input", + } + } + encoded, _ := json.Marshal(event) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-request-scoped", + payload: encoded, + })); err != nil { + t.Error(err) + return + } + } + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-request-scoped", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + err = client.SendInputConfirmed(context.Background(), []byte("rejected\n")) + if err == nil || !strings.Contains(err.Error(), "runner rejected terminal input") { + t.Fatalf("error = %v", err) + } + if !client.canInput.Load() { + t.Fatal("request-scoped rejection revoked terminal control") + } + if err := client.SendInputConfirmed(context.Background(), []byte("accepted\n")); err != nil { + t.Fatal(err) + } +} + +func TestSendInputConfirmedSharesOneReaderWithAttach(t *testing.T) { + firstInput := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-concurrent", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + for index := range 2 { + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + current, err := decodeFrame(payload) + if err != nil { + t.Error(err) + return + } + if current.messageType != messageInput { + t.Errorf("message type = %d", current.messageType) + return + } + if index == 0 { + close(firstInput) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-concurrent", + payload: []byte("attached\n"), + })); err != nil { + t.Error(err) + return + } + _, acknowledgement, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + ack, err := decodeFrame(acknowledgement) + if err != nil || ack.messageType != messageAck { + t.Errorf("output acknowledgement = %#v, %v", ack, err) + return + } + } + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-concurrent", + payload: accepted, + })); err != nil { + t.Error(err) + return + } + } + closed, _ := json.Marshal(eventPayload{Type: "closed"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-concurrent", + payload: closed, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-concurrent", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + terminal := newBlockingTerminal() + attachDone := make(chan error, 1) + go func() { + attachDone <- client.Attach(context.Background(), terminal, nil) + }() + <-terminal.started + + sendDone := make(chan error, 2) + go func() { + sendDone <- client.SendInputConfirmed(context.Background(), []byte("first\n")) + }() + <-firstInput + go func() { + sendDone <- client.SendInputConfirmed(context.Background(), []byte("second\n")) + }() + for range 2 { + if err := <-sendDone; err != nil { + t.Fatal(err) + } + } + if err := <-attachDone; err != nil { + t.Fatal(err) + } + if terminal.String() != "attached\n" { + t.Fatalf("output = %q", terminal.String()) + } +} + +func TestSendInputDrainsAcknowledgementBeforeConfirmedSend(t *testing.T) { + rawReceived := make(chan struct{}) + confirmedReceived := make(chan struct{}) + releaseRawAcknowledgement := make(chan struct{}) + releaseConfirmedAcknowledgement := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-raw-confirmed", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + + inputs := make(chan frame, 2) + go func() { + for range 2 { + _, payload, readErr := conn.Read(r.Context()) + if readErr != nil { + return + } + current, decodeErr := decodeFrame(payload) + if decodeErr != nil { + t.Error(decodeErr) + return + } + inputs <- current + } + }() + + raw := <-inputs + if raw.messageType != messageInput || string(raw.payload) != "raw\n" { + t.Errorf("raw input = %#v", raw) + return + } + close(rawReceived) + select { + case confirmed := <-inputs: + if confirmed.messageType != messageInput || string(confirmed.payload) != "confirmed\n" { + t.Errorf("confirmed input = %#v", confirmed) + return + } + close(confirmedReceived) + case <-releaseRawAcknowledgement: + } + + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-raw-confirmed", + payload: accepted, + })); err != nil { + t.Error(err) + return + } + select { + case <-confirmedReceived: + default: + confirmed := <-inputs + if confirmed.messageType != messageInput || string(confirmed.payload) != "confirmed\n" { + t.Errorf("confirmed input = %#v", confirmed) + return + } + close(confirmedReceived) + } + <-releaseConfirmedAcknowledgement + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-raw-confirmed", + payload: accepted, + })); err != nil { + t.Error(err) + } + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-raw-confirmed", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + if err := client.SendInput(context.Background(), []byte("raw\n")); err != nil { + t.Fatal(err) + } + <-rawReceived + + confirmedStarted := make(chan struct{}) + confirmedDone := make(chan error, 1) + go func() { + close(confirmedStarted) + confirmedDone <- client.SendInputConfirmed(context.Background(), []byte("confirmed\n")) + }() + <-confirmedStarted + select { + case <-confirmedReceived: + t.Fatal("confirmed input was written before the raw acknowledgement") + case <-time.After(100 * time.Millisecond): + } + + close(releaseRawAcknowledgement) + <-confirmedReceived + select { + case err := <-confirmedDone: + t.Fatalf("confirmed send completed from the raw acknowledgement: %v", err) + case <-time.After(100 * time.Millisecond): + } + + close(releaseConfirmedAcknowledgement) + if err := <-confirmedDone; err != nil { + t.Fatal(err) + } +} + +func TestSendInputConfirmedCanCancelWhileWaitingForPreviousConfirmation(t *testing.T) { + firstInput := make(chan struct{}) + releaseFirst := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-cancel-wait", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + + accepted, _ := json.Marshal(eventPayload{Type: "input-accepted"}) + for index := range 2 { + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + current, err := decodeFrame(payload) + if err != nil { + t.Error(err) + return + } + if current.messageType != messageInput { + t.Errorf("message type = %d", current.messageType) + return + } + if index == 0 { + close(firstInput) + <-releaseFirst + } + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-cancel-wait", + payload: accepted, + })); err != nil { + t.Error(err) + return + } + } + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-cancel-wait", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + firstCtx, firstCancel := context.WithTimeout(context.Background(), time.Second) + defer firstCancel() + firstDone := make(chan error, 1) + go func() { + firstDone <- client.SendInputConfirmed(firstCtx, []byte("first\n")) + }() + <-firstInput + + waitCtx, waitCancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer waitCancel() + started := time.Now() + err = client.SendInputConfirmed(waitCtx, []byte("canceled\n")) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("error = %v", err) + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("waiting cancellation took %s", elapsed) + } + + close(releaseFirst) + if err := <-firstDone; err != nil { + t.Fatal(err) + } + finalCtx, finalCancel := context.WithTimeout(context.Background(), time.Second) + defer finalCancel() + if err := client.SendInputConfirmed(finalCtx, []byte("final\n")); err != nil { + t.Fatal(err) + } +} + +func TestSendInputConfirmedReturnsImmediatelyForEmptyInput(t *testing.T) { + client := &Client{supportsInputAcknowledgement: true} + if err := client.SendInputConfirmed(context.Background(), nil); err != nil { + t.Fatal(err) + } +} + +func TestWaitForInputConfirmationPrefersDetachedAcceptanceOverTimeout(t *testing.T) { + client := &Client{readerDone: make(chan struct{})} + waiter := make(chan error, 1) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + time.AfterFunc(time.Millisecond, func() { + waiter <- nil + }) + + if err := client.waitForInputConfirmation(ctx, waiter); err != nil { + t.Fatalf("confirmation = %v", err) + } +} + +func TestWaitForInputConfirmationPrefersDetachedAcceptanceOverReaderShutdown(t *testing.T) { + readerDone := make(chan struct{}) + close(readerDone) + client := &Client{readerDone: readerDone} + waiter := make(chan error, 1) + time.AfterFunc(time.Millisecond, func() { + waiter <- nil + }) + + if err := client.waitForInputConfirmation(context.Background(), waiter); err != nil { + t.Fatalf("confirmation = %v", err) + } +} + +func TestSendInputConfirmedClosesAfterConfirmationTimeout(t *testing.T) { + inputReceived := make(chan struct{}) + releaseServer := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-confirm-timeout", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + close(inputReceived) + <-releaseServer + })) + defer func() { + close(releaseServer) + server.Close() + }() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-confirm-timeout", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + started := time.Now() + err = client.SendInputConfirmed(ctx, []byte("first\n")) + if !errors.Is(err, ErrInputDeliveryUnknown) { + t.Fatalf("error = %v", err) + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("error = %v", err) + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("confirmation deadline took %s", elapsed) + } + <-inputReceived + if err := client.SendInputConfirmed(context.Background(), []byte("second\n")); err == nil { + t.Fatal("timed-out client accepted another input") + } +} + +func TestAttachBoundsInputConfirmationAndRetiresConnection(t *testing.T) { + inputReceived := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{InputAcknowledgements: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-attach-confirm-timeout", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + close(inputReceived) + _, _, _ = conn.Read(r.Context()) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-attach-confirm-timeout", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.confirmationTimeout = 20 * time.Millisecond + + terminal := &readWriter{reader: strings.NewReader("blocked\n")} + started := time.Now() + err = client.Attach(context.Background(), terminal, nil) + if err == nil { + t.Fatal("attachment succeeded without input confirmation") + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("attachment confirmation timeout took %s", elapsed) + } + <-inputReceived + if err := client.SendInputConfirmed(context.Background(), []byte("second\n")); err == nil { + t.Fatal("timed-out attachment left the connection reusable") + } +} + +func TestSendInputConfirmedFallsBackWithoutServerCapability(t *testing.T) { + receivedInput := make(chan []byte, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-legacy", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + current, err := decodeFrame(payload) + if err != nil { + t.Error(err) + return + } + receivedInput <- append([]byte(nil), current.payload...) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-legacy", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := client.SendInputConfirmed(ctx, []byte("legacy\n")); err != nil { + t.Fatal(err) + } + if input := <-receivedInput; string(input) != "legacy\n" { + t.Fatalf("input = %q", input) + } +} + +func TestDialUsesConfiguredHTTPClientAndTimeout(t *testing.T) { + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + return + } + } + <-r.Context().Done() + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + httpClient := server.Client() + httpClient.Timeout = 25 * time.Millisecond + started := time.Now() + _, err = Dial(context.Background(), endpoint, "IS-timeout", Options{HTTPClient: httpClient}) + if err == nil { + t.Fatal("expected subscription timeout") + } + if elapsed := time.Since(started); elapsed > time.Second { + t.Fatalf("dial timeout took %s", elapsed) + } +} + +func TestDialDoesNotApplyHTTPClientTimeoutToEstablishedConnection(t *testing.T) { + receivedInput := make(chan []byte, 1) + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-established", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + current, err := decodeFrame(payload) + if err != nil { + t.Error(err) + return + } + receivedInput <- append([]byte(nil), current.payload...) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + httpClient := server.Client() + httpClient.Timeout = 25 * time.Millisecond + client, err := Dial(context.Background(), endpoint, "IS-established", Options{ + HTTPClient: httpClient, + }) + if err != nil { + t.Fatal(err) + } + defer client.Close() + time.Sleep(2 * httpClient.Timeout) + if err := client.SendInput(context.Background(), []byte("still-open\n")); err != nil { + t.Fatal(err) + } + if input := <-receivedInput; string(input) != "still-open\n" { + t.Fatalf("input = %q", input) + } +} + func TestAttachClosesCloseableTerminalAfterRemoteClosure(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := websocket.Accept(w, r, nil) @@ -282,7 +1860,179 @@ func TestAttachClosesCloseableTerminalAfterRemoteClosure(t *testing.T) { return } defer conn.Close(websocket.StatusNormalClosure, "") - + + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-1", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + closed, _ := json.Marshal(eventPayload{Type: "closed"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-1", + payload: closed, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-1", Options{Cols: 120, Rows: 34}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + terminal := newBlockingTerminal() + if err := client.Attach(context.Background(), terminal, nil); err != nil { + t.Fatal(err) + } + select { + case <-terminal.closed: + case <-time.After(time.Second): + t.Fatal("terminal was not closed") + } +} + +func TestAttachReturnsWhenContextCancelsAnUncancelableRead(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-cancel", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + <-r.Context().Done() + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-cancel", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + terminal := newUncancelableTerminal() + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- client.Attach(ctx, terminal, nil) + }() + <-terminal.started + cancel() + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("error = %v", err) + } + case <-time.After(time.Second): + t.Fatal("Attach did not return after context cancellation") + } + close(terminal.release) +} + +func TestAttachBoundsBlockedFrameConsumerShutdown(t *testing.T) { + client := &Client{ + readerDone: make(chan struct{}), + attachmentReady: make(chan struct{}), + attachmentShutdownTimeout: 10 * time.Millisecond, + } + terminal := newUncancelableReadBlockingWriteTerminal() + ctx, cancel := context.WithCancel(context.Background()) + attachDone := make(chan error, 1) + go func() { + attachDone <- client.Attach(ctx, terminal, nil) + }() + + <-terminal.readStarted + client.stateMu.Lock() + attachment := client.attachment + client.stateMu.Unlock() + if attachment == nil { + t.Fatal("attachment was not registered") + } + delivered := make(chan bool, 1) + go func() { + delivered <- client.deliverAttachment(context.Background(), frame{ + messageType: messageOutput, + payload: []byte("old output\n"), + }) + }() + <-terminal.writeStarted + if !<-delivered { + t.Fatal("output was not delivered to the attachment") + } + + cancel() + select { + case err := <-attachDone: + if !errors.Is(err, context.Canceled) { + t.Fatalf("error = %v", err) + } + case <-time.After(time.Second): + t.Fatal("Attach did not retire the blocked frame consumer") + } + + replacement := newBlockingTerminal() + replacementCtx, replacementCancel := context.WithCancel(context.Background()) + replacementDone := make(chan error, 1) + go func() { + replacementDone <- client.Attach(replacementCtx, replacement, nil) + }() + <-replacement.started + replacementCancel() + if err := <-replacementDone; !errors.Is(err, context.Canceled) { + t.Fatalf("replacement error = %v", err) + } + + close(terminal.releaseWrite) + select { + case <-terminal.writeDone: + case <-time.After(time.Second): + t.Fatal("blocked frame consumer did not exit after terminal shutdown") + } + close(terminal.releaseRead) +} + +func TestLateRetiredAttachmentWriteFailureClosesConnection(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") for range 2 { if _, _, err := conn.Read(r.Context()); err != nil { t.Error(err) @@ -292,18 +2042,21 @@ func TestAttachClosesCloseableTerminalAfterRemoteClosure(t *testing.T) { subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ messageType: messageEvent, - sessionID: "IS-1", + sessionID: "IS-late-write-failure", payload: subscribed, })); err != nil { t.Error(err) return } - closed, _ := json.Marshal(eventPayload{Type: "closed"}) - _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ - messageType: messageEvent, - sessionID: "IS-1", - payload: closed, - })) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-late-write-failure", + payload: []byte("blocked output\n"), + })); err != nil { + t.Error(err) + return + } + <-r.Context().Done() })) defer server.Close() @@ -311,21 +2064,169 @@ func TestAttachClosesCloseableTerminalAfterRemoteClosure(t *testing.T) { if err != nil { t.Fatal(err) } - client, err := Dial(context.Background(), endpoint, "IS-1", Options{Cols: 120, Rows: 34}) + client, err := Dial(context.Background(), endpoint, "IS-late-write-failure", Options{}) if err != nil { t.Fatal(err) } defer client.Close() + client.attachmentShutdownTimeout = 10 * time.Millisecond - terminal := newBlockingTerminal() - if err := client.Attach(context.Background(), terminal, nil); err != nil { + oldTerminal := newUncancelableReadBlockingWriteTerminal() + oldCtx, oldCancel := context.WithCancel(context.Background()) + oldDone := make(chan error, 1) + go func() { + oldDone <- client.Attach(oldCtx, oldTerminal, nil) + }() + <-oldTerminal.readStarted + <-oldTerminal.writeStarted + oldCancel() + if err := <-oldDone; !errors.Is(err, context.Canceled) { + t.Fatalf("old attachment error = %v", err) + } + + replacement := newBlockingTerminal() + replacementDone := make(chan error, 1) + go func() { + replacementDone <- client.Attach(context.Background(), replacement, nil) + }() + <-replacement.started + + close(oldTerminal.releaseWrite) + select { + case <-oldTerminal.writeDone: + case <-time.After(time.Second): + t.Fatal("retired attachment write did not finish") + } + select { + case err := <-replacementDone: + if !errors.Is(err, errBlockedTerminalWrite) { + t.Fatalf("replacement attachment error = %v", err) + } + case <-time.After(time.Second): + t.Fatal("late write failure did not retire the replacement attachment") + } + select { + case <-client.readerDone: + case <-time.After(time.Second): + t.Fatal("late write failure did not close the connection") + } + close(oldTerminal.releaseRead) +} + +func TestRetiredBlockedAttachmentAcknowledgementKeepsReplacementConnectionOpen(t *testing.T) { + firstAcknowledged := make(chan struct{}) + secondAcknowledged := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + welcome, _ := json.Marshal(welcomePayload{}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageWelcome, + payload: welcome, + })); err != nil { + t.Error(err) + return + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: true}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-stale-ack", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-stale-ack", + payload: []byte("old output\n"), + })); err != nil { + t.Error(err) + return + } + if err := readOutputAcknowledgement(r.Context(), conn, len("old output\n")); err != nil { + t.Error(err) + return + } + close(firstAcknowledged) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageOutput, + sessionID: "IS-stale-ack", + payload: []byte("replacement output\n"), + })); err != nil { + t.Error(err) + return + } + if err := readOutputAcknowledgement(r.Context(), conn, len("replacement output\n")); err != nil { + t.Error(err) + return + } + close(secondAcknowledged) + <-r.Context().Done() + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-stale-ack", Options{}) + if err != nil { t.Fatal(err) } + defer client.Close() + client.attachmentShutdownTimeout = 10 * time.Millisecond + + oldTerminal := newUncancelableReadBlockingSuccessfulWriteTerminal() + oldCtx, oldCancel := context.WithCancel(context.Background()) + oldDone := make(chan error, 1) + go func() { + oldDone <- client.Attach(oldCtx, oldTerminal, nil) + }() + <-oldTerminal.readStarted + <-oldTerminal.writeStarted + oldCancel() + if err := <-oldDone; !errors.Is(err, context.Canceled) { + t.Fatalf("old attachment error = %v", err) + } + + replacement := newBlockingTerminal() + replacementCtx, replacementCancel := context.WithCancel(context.Background()) + replacementDone := make(chan error, 1) + go func() { + replacementDone <- client.Attach(replacementCtx, replacement, nil) + }() + <-replacement.started + + close(oldTerminal.releaseWrite) + select { + case <-firstAcknowledged: + case <-time.After(time.Second): + t.Fatal("retired attachment did not acknowledge completed output") + } select { - case <-terminal.closed: + case <-secondAcknowledged: case <-time.After(time.Second): - t.Fatal("terminal was not closed") + t.Fatal("replacement connection did not acknowledge subsequent output") } + replacementCancel() + if err := <-replacementDone; !errors.Is(err, context.Canceled) { + t.Fatalf("replacement error = %v", err) + } + if got := replacement.String(); got != "replacement output\n" { + t.Fatalf("replacement output = %q", got) + } + close(oldTerminal.releaseRead) } func TestClientSubscribesReadOnlyAndSuppressesInput(t *testing.T) { @@ -578,9 +2479,131 @@ func TestClientContinuesReadOnlyAndResumesControl(t *testing.T) { if size := <-resumedSize; size != (Size{Cols: 132, Rows: 43}) { t.Fatalf("resumed size = %#v", size) } - if !client.canInput.Load() { - t.Fatal("client did not resume input control") + if client.canInput.Load() { + t.Fatal("closed terminal retained input control") + } +} + +func TestClientReplaysControlGrantedToNextAttachment(t *testing.T) { + grantControl := make(chan struct{}) + resumedSize := make(chan Size, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + t.Error(err) + return + } + defer conn.Close(websocket.StatusNormalClosure, "") + for range 2 { + if _, _, err := conn.Read(r.Context()); err != nil { + t.Error(err) + return + } + } + subscribed, _ := json.Marshal(eventPayload{Type: "subscribed", CanInput: false}) + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-between-attachments", + payload: subscribed, + })); err != nil { + t.Error(err) + return + } + <-grantControl + if err := conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageControlGranted, + sessionID: "IS-between-attachments", + })); err != nil { + t.Error(err) + return + } + _, payload, err := conn.Read(r.Context()) + if err != nil { + t.Error(err) + return + } + resize, err := decodeFrame(payload) + if err != nil || resize.messageType != messageResize { + t.Errorf("resumed resize = %#v, %v", resize, err) + return + } + resumedSize <- Size{ + Cols: binary.LittleEndian.Uint32(resize.payload[0:4]), + Rows: binary.LittleEndian.Uint32(resize.payload[4:8]), + } + closed, _ := json.Marshal(eventPayload{Type: "closed"}) + _ = conn.Write(r.Context(), websocket.MessageBinary, encodeFrame(frame{ + messageType: messageEvent, + sessionID: "IS-between-attachments", + payload: closed, + })) + })) + defer server.Close() + + endpoint, err := Endpoint(server.URL) + if err != nil { + t.Fatal(err) + } + client, err := Dial(context.Background(), endpoint, "IS-between-attachments", Options{}) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + if err := client.Attach( + context.Background(), + &readWriter{reader: bytes.NewReader(nil)}, + nil, + ); err != nil { + t.Fatal(err) + } + expectedSize := Size{Cols: 132, Rows: 43} + if err := client.Resize(context.Background(), expectedSize); err != nil { + t.Fatal(err) + } + close(grantControl) + deadline := time.Now().Add(time.Second) + for !client.canInput.Load() { + if time.Now().After(deadline) { + t.Fatal("control grant was not received") + } + time.Sleep(time.Millisecond) + } + + terminal := newBlockingTerminal() + attachCtx, attachCancel := context.WithTimeout(context.Background(), time.Second) + defer attachCancel() + if err := client.Attach(attachCtx, terminal, nil); err != nil { + t.Fatal(err) + } + if size := <-resumedSize; size != expectedSize { + t.Fatalf("resumed size = %#v", size) + } +} + +func readOutputAcknowledgement( + ctx context.Context, + conn *websocket.Conn, + expectedBytes int, +) error { + _, payload, err := conn.Read(ctx) + if err != nil { + return err } + current, err := decodeFrame(payload) + if err != nil { + return err + } + if current.messageType != messageAck { + return fmt.Errorf("acknowledgement message type = %d", current.messageType) + } + if len(current.payload) != 4 { + return fmt.Errorf("acknowledgement payload length = %d", len(current.payload)) + } + if acknowledged := binary.LittleEndian.Uint32(current.payload); acknowledged != uint32(expectedBytes) { + return fmt.Errorf("acknowledged bytes = %d, want %d", acknowledged, expectedBytes) + } + return nil } type readWriter struct { @@ -601,16 +2624,96 @@ func (rw *readWriter) CancelRead() error { } type blockingTerminal struct { - closed chan struct{} - once sync.Once + closed chan struct{} + once sync.Once + started chan struct{} + startOnce sync.Once + bytes.Buffer +} + +type uncancelableTerminal struct { + started chan struct{} + startOnce sync.Once + release chan struct{} bytes.Buffer } +var errBlockedTerminalWrite = errors.New("blocked terminal write") + +type uncancelableReadBlockingWriteTerminal struct { + readStarted chan struct{} + readOnce sync.Once + releaseRead chan struct{} + writeStarted chan struct{} + writeOnce sync.Once + releaseWrite chan struct{} + writeDone chan struct{} + writeErr error +} + +func newUncancelableReadBlockingWriteTerminal() *uncancelableReadBlockingWriteTerminal { + return &uncancelableReadBlockingWriteTerminal{ + readStarted: make(chan struct{}), + releaseRead: make(chan struct{}), + writeStarted: make(chan struct{}), + releaseWrite: make(chan struct{}), + writeDone: make(chan struct{}), + writeErr: errBlockedTerminalWrite, + } +} + +func newUncancelableReadBlockingSuccessfulWriteTerminal() *uncancelableReadBlockingWriteTerminal { + terminal := newUncancelableReadBlockingWriteTerminal() + terminal.writeErr = nil + return terminal +} + +func (terminal *uncancelableReadBlockingWriteTerminal) Read(_ []byte) (int, error) { + terminal.readOnce.Do(func() { + close(terminal.readStarted) + }) + <-terminal.releaseRead + return 0, io.EOF +} + +func (terminal *uncancelableReadBlockingWriteTerminal) Write(payload []byte) (int, error) { + defer close(terminal.writeDone) + terminal.writeOnce.Do(func() { + close(terminal.writeStarted) + }) + <-terminal.releaseWrite + if terminal.writeErr != nil { + return 0, terminal.writeErr + } + return len(payload), nil +} + +func newUncancelableTerminal() *uncancelableTerminal { + return &uncancelableTerminal{ + started: make(chan struct{}), + release: make(chan struct{}), + } +} + +func (terminal *uncancelableTerminal) Read(_ []byte) (int, error) { + terminal.startOnce.Do(func() { + close(terminal.started) + }) + <-terminal.release + return 0, io.EOF +} + func newBlockingTerminal() *blockingTerminal { - return &blockingTerminal{closed: make(chan struct{})} + return &blockingTerminal{ + closed: make(chan struct{}), + started: make(chan struct{}), + } } func (terminal *blockingTerminal) Read(_ []byte) (int, error) { + terminal.startOnce.Do(func() { + close(terminal.started) + }) <-terminal.closed return 0, io.ErrClosedPipe } diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/CrabboxVNCBridge.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/CrabboxVNCBridge.swift index 082046f5..a29bf205 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/CrabboxVNCBridge.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/CrabboxVNCBridge.swift @@ -36,6 +36,30 @@ private struct CrabboxVNCHandoff: Decodable { } final class CrabboxVNCBridge: @unchecked Sendable { + private static let configEnvironmentKeys = [ + "CRABBOX_CONFIG", + "XDG_CONFIG_HOME", + "XDG_STATE_HOME", + ] + + private static let networkEnvironmentKeys = [ + "ALL_PROXY", + "HTTP_PROXY", + "HTTPS_PROXY", + "NO_PROXY", + "all_proxy", + "http_proxy", + "https_proxy", + "no_proxy", + "AWS_CA_BUNDLE", + "CURL_CA_BUNDLE", + "GIT_SSL_CAINFO", + "NODE_EXTRA_CA_CERTS", + "REQUESTS_CA_BUNDLE", + "SSL_CERT_DIR", + "SSL_CERT_FILE", + ] + let request: VNCConnectionRequest private let process: Process @@ -129,6 +153,7 @@ final class CrabboxVNCBridge: @unchecked Sendable { process.standardInput = stdin process.standardOutput = stdout process.standardError = stderr + process.environment = commandEnvironment(from: ProcessInfo.processInfo.environment) do { try process.run() @@ -240,6 +265,15 @@ final class CrabboxVNCBridge: @unchecked Sendable { } } + static func commandEnvironment(from source: [String: String]) -> [String: String] { + SubprocessEnvironment.minimal( + from: source, + includeSSHAgent: true, + additionalInheritedKeys: networkEnvironmentKeys, + additionalInheritedPathKeys: configEnvironmentKeys + ) + } + private static func drain(_ pipe: Pipe) { pipe.fileHandleForReading.readabilityHandler = { handle in _ = handle.availableData @@ -253,11 +287,12 @@ final class CrabboxVNCBridge: @unchecked Sendable { value.unicodeScalars.allSatisfy { !CharacterSet.controlCharacters.contains($0) } } - private static func validGrant(_ grant: NativeVNCGrant) -> Bool { + static func validGrant(_ grant: NativeVNCGrant) -> Bool { let ticketPrefix = "native_vnc_" let ticketSuffix = grant.ticket.dropFirst(ticketPrefix.count) - let secureBroker = grant.brokerURL.scheme == "https" - || (grant.brokerURL.scheme == "http" + let scheme = grant.brokerURL.scheme?.lowercased() + let secureBroker = (scheme == "https" && grant.brokerURL.host?.isEmpty == false) + || (scheme == "http" && ["localhost", "127.0.0.1", "::1"].contains(grant.brokerURL.host ?? "")) return secureBroker && grant.brokerURL.user == nil diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetDesktopRegistration.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetDesktopRegistration.swift index f0347263..6fec9aaa 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetDesktopRegistration.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetDesktopRegistration.swift @@ -1,23 +1,152 @@ import Foundation +struct DesktopHostRegistrationRecoveryScope: Equatable, Hashable, Sendable { + let apiOrigin: String + let ownerSubject: String +} + +protocol DesktopHostRegistrationRecoveryScoping: Sendable { + func recoveryScope() async throws -> DesktopHostRegistrationRecoveryScope +} + protocol DesktopHostRegistering: Sendable { - func register(identity: TailnetIdentity, port: UInt16) async throws + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws +} + +struct DesktopHostRegistrationResultUncertainError: LocalizedError, Equatable, Sendable { + let message: String + + var errorDescription: String? { message } +} + +struct DesktopHostRegistrationSupersededError: LocalizedError, Equatable, Sendable { + var errorDescription: String? { + "The previous desktop publication is no longer current." + } +} + +actor DesktopHostRegistrationCoordinator { + private let registration: any DesktopHostRegistering + private var pendingOperation: Task? + + init(registration: any DesktopHostRegistering) { + self.registration = registration + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + let registration = self.registration + let operation = enqueue { + try await registration.register( + identity: identity, + port: port, + publicationID: publicationID + ) + } + return try await operation.value + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + let registration = self.registration + let operation = enqueue { + try await registration.recover(identity: identity, publicationID: publicationID) + } + return try await operation.value + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + let registration = self.registration + let operation = enqueue { + try await registration.unregister(identity: identity, ownershipToken: ownershipToken) + } + try await operation.value + } + + private func enqueue( + _ operation: @escaping @Sendable () async throws -> Value + ) -> Task { + let previous = pendingOperation + let task = Task { + await previous?.value + return try await operation() + } + pendingOperation = Task { + _ = try? await task.value + } + return task + } } -struct CrabfleetDesktopRegistration: DesktopHostRegistering, Sendable { +struct CrabfleetDesktopRegistration: + DesktopHostRegistering, + DesktopHostRegistrationRecoveryScoping, + @unchecked Sendable +{ + private struct RegistrationResponse: Decodable { + private enum CodingKeys: String, CodingKey { + case host + case ownershipToken + } + + struct Host: Decodable { + let id: String + } + + let host: Host + let ownershipToken: String? + + init(from decoder: any Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) + host = try container.decode(Host.self, forKey: .host) + ownershipToken = + container.contains(.ownershipToken) + ? try container.decode(String.self, forKey: .ownershipToken) + : nil + } + } + private struct RegistrationBody: Encodable { let name: String let address: String let port: UInt16 } + private struct RecoveryBody: Encodable { + let publicationID: String + } + + private struct RecoveryResponse: Decodable { + let ownershipToken: String? + } + + private struct NativeSessionResponse: Decodable { + struct User: Decodable { + let subject: String + } + + let user: User + } + private let baseURL: URL + private let apiOrigin: String private let sessionCookie: String - private let session: URLSession + private let transport: any HTTPDataTransport + static let ownershipModeHeader = "X-Crabfleet-Ownership-Mode" + static let publicationIDHeader = "X-Crabfleet-Publication-ID" + static let tokenOwnershipMode = "token-v1" init?( environment: [String: String] = ProcessInfo.processInfo.environment, - session: URLSession = .shared + transport: any HTTPDataTransport = RejectingRedirectURLSessionTransport() ) { guard let rawURL = environment["CRABFLEET_API_URL"], @@ -35,23 +164,145 @@ struct CrabfleetDesktopRegistration: DesktopHostRegistering, Sendable { normalizedURL.deleteLastPathComponent() } guard normalizedURL.path.isEmpty || normalizedURL.path == "/" else { return nil } + guard let apiOrigin = Self.normalizedAPIOrigin(normalizedURL) else { return nil } self.baseURL = normalizedURL + self.apiOrigin = apiOrigin self.sessionCookie = cookie - self.session = session + self.transport = transport } - func register(identity: TailnetIdentity, port: UInt16) async throws { - let request = try registrationRequest(identity: identity, port: port) - let (_, response) = try await session.data(for: request) - guard let http = response as? HTTPURLResponse else { + func recoveryScope() async throws -> DesktopHostRegistrationRecoveryScope { + let request = nativeSessionRequest() + let data: Data + let http: HTTPURLResponse + (data, http) = try await transport.data(for: request) + try validate(response: http, for: request, acceptingNotFound: false) + guard + let response = try? JSONDecoder().decode(NativeSessionResponse.self, from: data), + Self.isValidOwnerSubject(response.user.subject) + else { throw DesktopHostRegistrationError.invalidResponse } - guard (200..<300).contains(http.statusCode) else { - throw DesktopHostRegistrationError.httpStatus(http.statusCode) + return DesktopHostRegistrationRecoveryScope( + apiOrigin: apiOrigin, + ownerSubject: response.user.subject + ) + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + let request = try registrationRequest( + identity: identity, + port: port, + publicationID: publicationID + ) + let data: Data + let http: HTTPURLResponse + do { + (data, http) = try await transport.data(for: request) + } catch { + throw DesktopHostRegistrationResultUncertainError(message: error.localizedDescription) + } + do { + try validate(response: http, for: request, acceptingNotFound: false) + } catch let error as DesktopHostRegistrationError { + if case .httpStatus(let status) = error, status >= 500 { + throw DesktopHostRegistrationResultUncertainError( + message: error.localizedDescription + ) + } + throw error } + guard + let response = try? JSONDecoder().decode(RegistrationResponse.self, from: data), + response.host.id == Self.hostID(identity: identity) + else { + throw DesktopHostRegistrationResultUncertainError( + message: DesktopHostRegistrationError.invalidResponse.localizedDescription + ) + } + if let ownershipToken = response.ownershipToken, + !Self.isValidOwnershipToken(ownershipToken) + { + throw DesktopHostRegistrationResultUncertainError( + message: DesktopHostRegistrationError.invalidResponse.localizedDescription + ) + } + return response.ownershipToken } - func registrationRequest(identity: TailnetIdentity, port: UInt16) throws -> URLRequest { + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + let request = try recoveryRequest(identity: identity, publicationID: publicationID) + let data: Data + let http: HTTPURLResponse + do { + (data, http) = try await transport.data(for: request) + } catch { + throw DesktopHostRegistrationResultUncertainError(message: error.localizedDescription) + } + if http.statusCode == 404 { + try validate(response: http, for: request, acceptingNotFound: true) + throw DesktopHostRegistrationResultUncertainError( + message: "Desktop publication recovery is unavailable on this server." + ) + } + do { + try validate(response: http, for: request, acceptingNotFound: false) + } catch let error as DesktopHostRegistrationError { + if case .httpStatus(let status) = error, status >= 500 { + throw DesktopHostRegistrationResultUncertainError( + message: error.localizedDescription + ) + } + throw error + } + guard let response = try? JSONDecoder().decode(RecoveryResponse.self, from: data) else { + throw DesktopHostRegistrationResultUncertainError( + message: DesktopHostRegistrationError.invalidResponse.localizedDescription + ) + } + if let ownershipToken = response.ownershipToken, + !Self.isValidOwnershipToken(ownershipToken) + { + throw DesktopHostRegistrationResultUncertainError( + message: DesktopHostRegistrationError.invalidResponse.localizedDescription + ) + } + return response.ownershipToken + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + let request = try removalRequest(identity: identity, ownershipToken: ownershipToken) + let (_, http) = try await transport.data(for: request) + try validate(response: http, for: request, acceptingNotFound: true) + } + + private func validate( + response: HTTPURLResponse, + for request: URLRequest, + acceptingNotFound: Bool + ) throws { + guard response.url == request.url else { + throw DesktopHostRegistrationError.redirectRejected + } + guard (200..<300).contains(response.statusCode) + || (acceptingNotFound && response.statusCode == 404) + else { + throw DesktopHostRegistrationError.httpStatus(response.statusCode) + } + } + + func registrationRequest( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) throws -> URLRequest { + guard Self.isValidOwnershipToken(publicationID) else { + throw DesktopHostRegistrationError.invalidResponse + } let hostID = Self.hostID(identity: identity) let url = baseURL @@ -64,6 +315,8 @@ struct CrabfleetDesktopRegistration: DesktopHostRegistering, Sendable { request.setValue("application/json", forHTTPHeaderField: "Content-Type") request.setValue("application/json", forHTTPHeaderField: "Accept") request.setValue(sessionCookie, forHTTPHeaderField: "Cookie") + request.setValue(Self.tokenOwnershipMode, forHTTPHeaderField: Self.ownershipModeHeader) + request.setValue(publicationID, forHTTPHeaderField: Self.publicationIDHeader) request.httpBody = try JSONEncoder().encode( RegistrationBody( name: identity.hostName.isEmpty ? identity.dnsName : identity.hostName, @@ -73,6 +326,67 @@ struct CrabfleetDesktopRegistration: DesktopHostRegistering, Sendable { return request } + func recoveryRequest(identity: TailnetIdentity, publicationID: String) throws -> URLRequest { + guard Self.isValidOwnershipToken(publicationID) else { + throw DesktopHostRegistrationError.invalidResponse + } + var components = URLComponents( + url: + baseURL + .appending(path: "api") + .appending(path: "desktop-hosts") + .appending(path: Self.hostID(identity: identity)), + resolvingAgainstBaseURL: false + ) + components?.queryItems = [URLQueryItem(name: "recover", value: "1")] + guard let url = components?.url else { + throw DesktopHostRegistrationError.invalidResponse + } + var request = URLRequest(url: url) + request.httpMethod = "POST" + request.timeoutInterval = 15 + request.setValue("application/json", forHTTPHeaderField: "Content-Type") + request.setValue("application/json", forHTTPHeaderField: "Accept") + request.setValue(sessionCookie, forHTTPHeaderField: "Cookie") + request.httpBody = try JSONEncoder().encode(RecoveryBody(publicationID: publicationID)) + return request + } + + func removalRequest(identity: TailnetIdentity, ownershipToken: String?) throws -> URLRequest { + if let ownershipToken, !Self.isValidOwnershipToken(ownershipToken) { + throw DesktopHostRegistrationError.invalidResponse + } + let url = + baseURL + .appending(path: "api") + .appending(path: "desktop-hosts") + .appending(path: Self.hostID(identity: identity)) + var request = URLRequest(url: url) + request.httpMethod = "DELETE" + request.timeoutInterval = 15 + request.setValue("application/json", forHTTPHeaderField: "Accept") + request.setValue(sessionCookie, forHTTPHeaderField: "Cookie") + if let ownershipToken { + request.setValue(ownershipToken, forHTTPHeaderField: "X-Crabfleet-Ownership-Token") + } + return request + } + + private func nativeSessionRequest() -> URLRequest { + let url = + baseURL + .appending(path: "api") + .appending(path: "native") + .appending(path: "v1") + .appending(path: "session") + var request = URLRequest(url: url) + request.httpMethod = "GET" + request.timeoutInterval = 15 + request.setValue("application/json", forHTTPHeaderField: "Accept") + request.setValue(sessionCookie, forHTTPHeaderField: "Cookie") + return request + } + static func hostID(identity: TailnetIdentity) -> String { let dnsLabel = identity.dnsName.split(separator: ".").first.map(String.init) ?? "" let normalized = dnsLabel.lowercased().filter { @@ -95,16 +409,43 @@ struct CrabfleetDesktopRegistration: DesktopHostRegistering, Sendable { if scheme == "https" { return true } return scheme == "http" && (host == "127.0.0.1" || host == "::1") } + + private static func normalizedAPIOrigin(_ url: URL) -> String? { + let scheme = url.scheme?.lowercased() + var components = URLComponents() + components.scheme = scheme + components.host = url.host?.lowercased() + if !((scheme == "https" && url.port == 443) || (scheme == "http" && url.port == 80)) { + components.port = url.port + } + return components.url?.absoluteString + } + + private static func isValidOwnerSubject(_ value: String) -> Bool { + !value.isEmpty && value.utf8.count <= 512 + && !value.unicodeScalars.contains(where: CharacterSet.controlCharacters.contains) + } + + private static func isValidOwnershipToken(_ value: String) -> Bool { + !value.isEmpty && value.utf8.count <= 200 + && !value.unicodeScalars.contains { + CharacterSet.whitespacesAndNewlines.contains($0) + || CharacterSet.controlCharacters.contains($0) + } + } } -enum DesktopHostRegistrationError: LocalizedError { +enum DesktopHostRegistrationError: LocalizedError, Equatable { case invalidResponse + case redirectRejected case httpStatus(Int) var errorDescription: String? { switch self { case .invalidResponse: "Crabfleet returned an invalid registration response." + case .redirectRejected: + "Crabfleet redirected the desktop registration request." case .httpStatus(let status): "Crabfleet registration returned HTTP \(status)." } diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetMacApp.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetMacApp.swift index 25cf1939..4c4b6b2e 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetMacApp.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/CrabfleetMacApp.swift @@ -47,19 +47,65 @@ enum VNCConnectionLaunchMode { @MainActor final class CrabfleetApplicationDelegate: NSObject, NSApplicationDelegate { - private var shareController: PrivateMacShareController? + let shareController: PrivateMacShareController + private let replyToTerminationRequest: @MainActor (Bool) -> Void + private let isAutoShareRequested: @MainActor () -> Bool + private let autoShareDelay: Duration private var autoShareTask: Task? + private var terminationTask: Task? + + override init() { + shareController = PrivateMacShareController() + replyToTerminationRequest = { shouldTerminate in + NSApp.reply(toApplicationShouldTerminate: shouldTerminate) + } + isAutoShareRequested = { PrivateMacShareLaunchMode.isRequested() } + autoShareDelay = .milliseconds(500) + super.init() + } + + init( + shareController: PrivateMacShareController, + replyToTerminationRequest: @escaping @MainActor (Bool) -> Void = { + NSApp.reply(toApplicationShouldTerminate: $0) + }, + isAutoShareRequested: @escaping @MainActor () -> Bool = { + PrivateMacShareLaunchMode.isRequested() + }, + autoShareDelay: Duration = .milliseconds(500) + ) { + self.shareController = shareController + self.replyToTerminationRequest = replyToTerminationRequest + self.isAutoShareRequested = isAutoShareRequested + self.autoShareDelay = autoShareDelay + super.init() + } func applicationDidFinishLaunching(_ notification: Notification) { - guard PrivateMacShareLaunchMode.isRequested() else { return } + guard isAutoShareRequested() else { return } NSApp.activate(ignoringOtherApps: true) - let controller = PrivateMacShareController() - shareController = controller autoShareTask = Task { [weak self] in guard let self else { return } - try? await Task.sleep(for: .milliseconds(500)) - await self.startPrivateShare(controller) + do { + try await Task.sleep(for: autoShareDelay) + } catch { + return + } + guard !Task.isCancelled else { return } + await self.startPrivateShare(shareController) + } + } + + func applicationShouldTerminate(_ sender: NSApplication) -> NSApplication.TerminateReply { + autoShareTask?.cancel() + guard terminationTask == nil else { return .terminateLater } + terminationTask = Task { [weak self] in + guard let self else { return } + let cleanupCanRecover = await shareController.stopAndWaitForCleanup() + terminationTask = nil + replyToTerminationRequest(cleanupCanRecover) } + return .terminateLater } func applicationWillTerminate(_ notification: Notification) { @@ -68,6 +114,7 @@ final class CrabfleetApplicationDelegate: NSObject, NSApplicationDelegate { private func startPrivateShare(_ controller: PrivateMacShareController) async { await controller.refresh() + guard !Task.isCancelled else { return } report( "private share prerequisites: tailnet \(controller.identity == nil ? "unavailable" : "ready"), " + "Screen Recording \(controller.screenRecordingGranted ? "allowed" : "denied")" @@ -75,23 +122,34 @@ final class CrabfleetApplicationDelegate: NSObject, NSApplicationDelegate { ) if !controller.screenRecordingGranted { await controller.requestScreenRecordingPermission() + guard !Task.isCancelled else { return } } let clock = ContinuousClock() let deadline = clock.now.advanced(by: .seconds(300)) while !Task.isCancelled, clock.now < deadline { await controller.refresh() + guard !Task.isCancelled else { return } if controller.canStart { await controller.start() + guard !Task.isCancelled else { return } for _ in 0..<50 where controller.phase == .starting { - try? await Task.sleep(for: .milliseconds(100)) + do { + try await Task.sleep(for: .milliseconds(100)) + } catch { + return + } } let address = controller.connectionAddress.map { " at \($0)" } ?? "" let notice = controller.notice.map { ": \($0)" } ?? "" report("private share \(controller.phase.title.lowercased())\(address)\(notice)") return } - try? await Task.sleep(for: .seconds(2)) + do { + try await Task.sleep(for: .seconds(2)) + } catch { + return + } } let missing = [ @@ -119,6 +177,7 @@ struct CrabfleetMacApp: App { fleetStore: fleetStore, connectionLibrary: connectionLibrary, sessionPool: sessionPool, + privateShare: appDelegate.shareController, launchConnection: launchConnection ) } @@ -142,6 +201,7 @@ private struct CrabfleetAppRoot: View { @ObservedObject var fleetStore: FleetStore @ObservedObject var connectionLibrary: ConnectionLibrary @ObservedObject var sessionPool: VNCSessionPool + @ObservedObject var privateShare: PrivateMacShareController let launchConnection: VNCAddress? @Environment(\.scenePhase) private var scenePhase @@ -151,11 +211,13 @@ private struct CrabfleetAppRoot: View { fleetStore: FleetStore, connectionLibrary: ConnectionLibrary, sessionPool: VNCSessionPool, + privateShare: PrivateMacShareController, launchConnection: VNCAddress? ) { self.fleetStore = fleetStore self.connectionLibrary = connectionLibrary self.sessionPool = sessionPool + self.privateShare = privateShare self.launchConnection = launchConnection _localOnly = State( initialValue: ProcessInfo.processInfo.environment["CRABFLEET_LOCAL_ONLY"] == "1" @@ -170,6 +232,7 @@ private struct CrabfleetAppRoot: View { store: fleetStore, connections: connectionLibrary, sessions: sessionPool, + privateShare: privateShare, launchConnection: launchConnection, deploymentLabel: fleetStore.isConnected ? fleetStore.deploymentLabel : "Local VNC", accountLabel: fleetStore.isConnected ? fleetStore.accountLabel : NSUserName(), diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/FleetRootView.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/FleetRootView.swift index 96171fe3..1194db77 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/FleetRootView.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/FleetRootView.swift @@ -5,13 +5,12 @@ struct FleetRootView: View { @ObservedObject var store: FleetStore @ObservedObject var connections: ConnectionLibrary @ObservedObject var sessions: VNCSessionPool + @ObservedObject var privateShare: PrivateMacShareController let launchConnection: VNCAddress? let deploymentLabel: String let accountLabel: String let disconnectLabel: String let disconnectDeployment: () -> Void - @StateObject private var privateShare = PrivateMacShareController() - @Namespace private var desktopTransition @Environment(\.accessibilityReduceMotion) private var reduceMotion diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/FleetStore.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/FleetStore.swift index 9ad5fdc9..c3a50b5b 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/FleetStore.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/FleetStore.swift @@ -498,7 +498,12 @@ final class FleetStore: ObservableObject { throw CancellationError() } do { - return (try await operation(token), token) + let value = try await operation(token) + try Task.checkCancellation() + guard isCurrent(generation), client === api, connectedOrigin == origin else { + throw CancellationError() + } + return (value, token) } catch NativeAPIError.unauthorized { try Task.checkCancellation() guard isCurrent(generation), client === api, connectedOrigin == origin else { @@ -518,7 +523,12 @@ final class FleetStore: ObservableObject { credential = .init(origin: origin, token: refreshed) adopted = true } - return (try await operation(refreshed), refreshed) + let value = try await operation(refreshed) + try Task.checkCancellation() + guard isCurrent(generation), client === api, connectedOrigin == origin else { + throw CancellationError() + } + return (value, refreshed) } catch { if !adopted, preserveRotatedCredentialOnTransientFailure, diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/HostClipboard.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/HostClipboard.swift index 0ace3b80..b5ac18c9 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/HostClipboard.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/HostClipboard.swift @@ -27,7 +27,6 @@ final class HostClipboardBridge: HostClipboardSyncing, @unchecked Sendable { private var lastObservedChangeCount: Int? private var suppressedChangeCount: Int? private var lastKnownText: String? - private var lastAppliedClientText: String? init( pasteboard: NSPasteboard = .general, @@ -67,7 +66,6 @@ final class HostClipboardBridge: HostClipboardSyncing, @unchecked Sendable { func detach() { withLock { pusher = nil - lastAppliedClientText = nil suppressedChangeCount = nil } DispatchQueue.main.async { [weak self] in @@ -80,9 +78,14 @@ final class HostClipboardBridge: HostClipboardSyncing, @unchecked Sendable { guard text.utf8.count <= RFBWire.maximumClipboardBytes else { return } DispatchQueue.main.async { [weak self] in guard let self else { return } - let alreadyCurrent = self.withLock { self.lastKnownText == text } - if alreadyCurrent { - self.withLock { self.lastAppliedClientText = text } + let changeCount = self.pasteboard.changeCount + let currentText = self.pasteboard.string(forType: .string) + guard self.pasteboard.changeCount == changeCount else { return } + if currentText == text { + self.withLock { + self.lastObservedChangeCount = changeCount + self.lastKnownText = text + } return } self.pasteboard.clearContents() @@ -90,7 +93,6 @@ final class HostClipboardBridge: HostClipboardSyncing, @unchecked Sendable { self.withLock { self.suppressedChangeCount = self.pasteboard.changeCount self.lastObservedChangeCount = self.pasteboard.changeCount - self.lastAppliedClientText = text self.lastKnownText = text } } @@ -106,6 +108,7 @@ final class HostClipboardBridge: HostClipboardSyncing, @unchecked Sendable { let previous = withLock { lastObservedChangeCount } guard changeCount != previous else { return } + let types = pasteboard.types ?? [] let text = pasteboard.string(forType: .string) guard pasteboard.changeCount == changeCount else { return } @@ -113,19 +116,19 @@ final class HostClipboardBridge: HostClipboardSyncing, @unchecked Sendable { var pushHandler: (@Sendable (String) -> Void)? withLock { lastObservedChangeCount = changeCount - lastKnownText = text if suppressedChangeCount == changeCount { suppressedChangeCount = nil return } - guard let text, !text.isEmpty, - text != lastAppliedClientText, - text.utf8.count <= RFBWire.maximumClipboardBytes - else { + lastKnownText = text + guard let outboundText = text ?? (types.isEmpty ? "" : nil) else { + return + } + guard outboundText.utf8.count <= RFBWire.maximumClipboardBytes else { return } - textToPush = text + textToPush = outboundText pushHandler = pusher } if let textToPush, let pushHandler { diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/MacRemoteInput.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/MacRemoteInput.swift index 3244940e..60c00cc7 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/MacRemoteInput.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/MacRemoteInput.swift @@ -14,21 +14,98 @@ extension RemoteInputForwarding { func releaseAllInput() {} } +final class RemoteInputSessionGate: @unchecked Sendable { + private let input: any RemoteInputForwarding + private let lock = NSLock() + private var acceptingInput = true + private var viewOnly: Bool + + init(input: any RemoteInputForwarding, viewOnly: Bool) { + self.input = input + self.viewOnly = viewOnly + } + + func setViewOnly(_ enabled: Bool) { + withLock { + guard acceptingInput else { return } + let shouldReleaseInput = enabled && !viewOnly + viewOnly = enabled + if shouldReleaseInput { + input.releaseAllInput() + } + } + } + + func keyEvent(down: Bool, keysym: UInt32) { + withLock { + guard acceptingInput, !viewOnly else { return } + input.keyEvent(down: down, keysym: keysym) + } + } + + func pointerEvent(buttonMask: UInt8, x: UInt16, y: UInt16) { + withLock { + guard acceptingInput, !viewOnly else { return } + input.pointerEvent(buttonMask: buttonMask, x: x, y: y) + } + } + + func finish() { + withLock { + guard acceptingInput else { return } + acceptingInput = false + input.releaseAllInput() + } + } + + private func withLock(_ body: () -> T) -> T { + lock.lock() + defer { lock.unlock() } + return body() + } +} + final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable { + private static let releaseRetryDelay: DispatchTimeInterval = .milliseconds(250) + private static let releaseRetryLimit = 120 + private let descriptor: CapturedDisplayDescriptor private let eventQueue = DispatchQueue( label: "org.openclaw.crabfleet.remote-input", qos: .userInteractive ) + private let accessibilityGranted: @Sendable () -> Bool + private let pendingReleaseRetryDelay: DispatchTimeInterval + private let pendingReleaseRetryLimit: Int + private let keyEventPoster: (@Sendable (Bool, UInt32) -> Void)? + private let mouseEventPoster: + (@Sendable (CGEventType, CGPoint, CGMouseButton) -> Void)? private let frameSizeLock = NSLock() private var frameWidth: Int private var frameHeight: Int private var previousButtonMask: UInt8 = 0 private var previousPointerLocation: CGPoint private var pressedKeysyms: Set = [] + private var hasPendingRelease = false + private var pendingReleaseRetryScheduled = false + private var pendingReleaseRetriesRemaining = 0 - init(descriptor: CapturedDisplayDescriptor) { + init( + descriptor: CapturedDisplayDescriptor, + accessibilityGranted: @escaping @Sendable () -> Bool = { + MacRemoteInputController.isAccessibilityGranted + }, + pendingReleaseRetryDelay: DispatchTimeInterval = MacRemoteInputController.releaseRetryDelay, + pendingReleaseRetryLimit: Int = MacRemoteInputController.releaseRetryLimit, + keyEventPoster: (@Sendable (Bool, UInt32) -> Void)? = nil, + mouseEventPoster: (@Sendable (CGEventType, CGPoint, CGMouseButton) -> Void)? = nil + ) { self.descriptor = descriptor + self.accessibilityGranted = accessibilityGranted + self.pendingReleaseRetryDelay = pendingReleaseRetryDelay + self.pendingReleaseRetryLimit = max(pendingReleaseRetryLimit, 0) + self.keyEventPoster = keyEventPoster + self.mouseEventPoster = mouseEventPoster frameWidth = descriptor.frameWidth frameHeight = descriptor.frameHeight previousPointerLocation = descriptor.displayBounds.origin @@ -58,7 +135,8 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable func keyEvent(down: Bool, keysym: UInt32) { eventQueue.async { [self] in - guard Self.isAccessibilityGranted else { return } + guard accessibilityGranted() else { return } + flushPendingRelease() postKeyEvent(down: down, keysym: keysym) if down { pressedKeysyms.insert(keysym) @@ -70,7 +148,8 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable func pointerEvent(buttonMask: UInt8, x: UInt16, y: UInt16) { eventQueue.async { [self] in - guard Self.isAccessibilityGranted else { return } + guard accessibilityGranted() else { return } + flushPendingRelease() let location = mappedLocation(x: x, y: y) previousPointerLocation = location let changedButtons = previousButtonMask ^ buttonMask @@ -79,12 +158,7 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable for button in Self.mouseButtons where changedButtons & button.mask != 0 { let isDown = buttonMask & button.mask != 0 let type = isDown ? button.downType : button.upType - CGEvent( - mouseEventSource: eventSource(), - mouseType: type, - mouseCursorPosition: location, - mouseButton: button.button - )?.post(tap: .cghidEventTap) + postMouseEvent(type: type, location: location, button: button.button) postedButtonChange = true } @@ -99,12 +173,7 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable } else { moveType = .mouseMoved } - CGEvent( - mouseEventSource: eventSource(), - mouseType: moveType, - mouseCursorPosition: location, - mouseButton: .left - )?.post(tap: .cghidEventTap) + postMouseEvent(type: moveType, location: location, button: .left) } let newWheelBits = buttonMask & ~previousButtonMask @@ -118,29 +187,22 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable func releaseAllInput() { eventQueue.async { [self] in - guard Self.isAccessibilityGranted else { return } - for keysym in pressedKeysyms { - postKeyEvent(down: false, keysym: keysym) - } - pressedKeysyms.removeAll() - - for button in Self.mouseButtons where previousButtonMask & button.mask != 0 { - CGEvent( - mouseEventSource: eventSource(), - mouseType: button.upType, - mouseCursorPosition: previousPointerLocation, - mouseButton: button.button - )?.post(tap: .cghidEventTap) - } - previousButtonMask = 0 + guard !pressedKeysyms.isEmpty || previousButtonMask != 0 else { return } + hasPendingRelease = true + pendingReleaseRetriesRemaining = pendingReleaseRetryLimit + flushPendingRelease() } } private func postKeyEvent(down: Bool, keysym: UInt32) { + if let keyEventPoster { + keyEventPoster(down, keysym) + return + } let event: CGEvent? if let keyCode = Self.keyCode(for: keysym) { event = CGEvent(keyboardEventSource: eventSource(), virtualKey: keyCode, keyDown: down) - } else if let scalar = UnicodeScalar(keysym) { + } else if let scalar = Self.unicodeScalar(for: keysym) { let candidate = CGEvent(keyboardEventSource: eventSource(), virtualKey: 0, keyDown: down) var codeUnits = Array(String(scalar).utf16) candidate?.keyboardSetUnicodeString( @@ -154,6 +216,56 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable event?.post(tap: .cghidEventTap) } + private func postMouseEvent( + type: CGEventType, + location: CGPoint, + button: CGMouseButton + ) { + if let mouseEventPoster { + mouseEventPoster(type, location, button) + return + } + CGEvent( + mouseEventSource: eventSource(), + mouseType: type, + mouseCursorPosition: location, + mouseButton: button + )?.post(tap: .cghidEventTap) + } + + private func flushPendingRelease() { + guard hasPendingRelease else { return } + guard accessibilityGranted() else { + schedulePendingReleaseRetry() + return + } + + for keysym in pressedKeysyms { + postKeyEvent(down: false, keysym: keysym) + } + for button in Self.mouseButtons where previousButtonMask & button.mask != 0 { + postMouseEvent( + type: button.upType, + location: previousPointerLocation, + button: button.button + ) + } + pressedKeysyms.removeAll() + previousButtonMask = 0 + hasPendingRelease = false + pendingReleaseRetriesRemaining = 0 + } + + private func schedulePendingReleaseRetry() { + guard !pendingReleaseRetryScheduled, pendingReleaseRetriesRemaining > 0 else { return } + pendingReleaseRetriesRemaining -= 1 + pendingReleaseRetryScheduled = true + eventQueue.asyncAfter(deadline: .now() + pendingReleaseRetryDelay) { [self] in + self.pendingReleaseRetryScheduled = false + self.flushPendingRelease() + } + } + private func eventSource() -> CGEventSource? { CGEventSource(stateID: .hidSystemState) } @@ -225,6 +337,14 @@ final class MacRemoteInputController: RemoteInputForwarding, @unchecked Sendable } } + static func unicodeScalar(for keysym: UInt32) -> UnicodeScalar? { + let value = + keysym & 0xFF00_0000 == 0x0100_0000 + ? keysym & 0x00FF_FFFF + : keysym + return UnicodeScalar(value) + } + private static let asciiKeyCodes: [Character: CGKeyCode] = [ "a": CGKeyCode(kVK_ANSI_A), "b": CGKeyCode(kVK_ANSI_B), "c": CGKeyCode(kVK_ANSI_C), "d": CGKeyCode(kVK_ANSI_D), diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/NativeAPIClient.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/NativeAPIClient.swift index 7712e5dc..7501a240 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/NativeAPIClient.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/NativeAPIClient.swift @@ -647,8 +647,10 @@ final class NativeAPIClient: NativeAPIClientProtocol { guard url.user == nil, url.password == nil, url.query == nil, url.fragment == nil else { return false } - if url.scheme == "https" { return true } - return url.scheme == "http" && ["localhost", "127.0.0.1", "::1"].contains(url.host ?? "") + let scheme = url.scheme?.lowercased() + let host = url.host ?? "" + if scheme == "https" { return !host.isEmpty } + return scheme == "http" && ["localhost", "127.0.0.1", "::1"].contains(host) } private func validOpaqueNativeVNCValue(_ value: String, maximumBytes: Int) -> Bool { diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/PrivateMacShareController.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/PrivateMacShareController.swift index 3decba85..1347087e 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/PrivateMacShareController.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/PrivateMacShareController.swift @@ -1,5 +1,6 @@ import AppKit import CoreGraphics +import CryptoKit import Foundation import ServiceManagement @@ -12,10 +13,452 @@ enum PrivateMacSharePermissionPolicy { } } +@MainActor +final class PrivateMacShareStopCoordinator { + private var isPerforming = false + private var waiters: [CheckedContinuation] = [] + + func perform(_ body: @escaping @MainActor () async -> Void) async { + if isPerforming { + await withCheckedContinuation { continuation in + self.waiters.append(continuation) + } + return + } + + isPerforming = true + await body() + isPerforming = false + let waiters = waiters + self.waiters.removeAll() + for waiter in waiters { + waiter.resume() + } + } +} + +@MainActor +protocol DesktopHostRegistrationStateStoring: AnyObject { + func containsState() -> Bool + func load(scope: DesktopHostRegistrationRecoveryScope) throws -> Data? + func save(_ data: Data?, scope: DesktopHostRegistrationRecoveryScope) throws +} + +extension DesktopHostRegistrationStateStoring { + func containsState() -> Bool { true } +} + +enum DesktopHostRegistrationPersistenceError: LocalizedError { + case missingScope + case unreadableState + case writeFailed + + var errorDescription: String? { + switch self { + case .missingScope: + "The desktop publication recovery scope is unavailable." + case .unreadableState: + "The saved desktop publication recovery state is unreadable." + case .writeFailed: + "The desktop publication recovery state could not be saved." + } + } +} + +@MainActor +final class UserDefaultsDesktopHostRegistrationStateStore: + DesktopHostRegistrationStateStoring +{ + nonisolated static let defaultKey = "org.openclaw.crabfleet.share.desktop-publications" + + private let defaults: UserDefaults + private let key: String + + init(defaults: UserDefaults = .standard, key: String = defaultKey) { + self.defaults = defaults + self.key = key + } + + func containsState() -> Bool { + let prefix = "\(key).v2." + return defaults.dictionaryRepresentation().keys.contains { $0.hasPrefix(prefix) } + } + + func load(scope: DesktopHostRegistrationRecoveryScope) throws -> Data? { + defaults.data(forKey: scopedKey(scope)) + } + + func save(_ data: Data?, scope: DesktopHostRegistrationRecoveryScope) throws { + let scopedKey = scopedKey(scope) + if let data { + defaults.set(data, forKey: scopedKey) + } else { + defaults.removeObject(forKey: scopedKey) + } + guard defaults.synchronize() else { + throw DesktopHostRegistrationPersistenceError.writeFailed + } + } + + private func scopedKey(_ scope: DesktopHostRegistrationRecoveryScope) -> String { + let scopeData = Data("\(scope.apiOrigin)\u{0}\(scope.ownerSubject)".utf8) + let digest = SHA256.hash(data: scopeData).map { String(format: "%02x", $0) }.joined() + return "\(key).v2.\(digest)" + } +} + +@MainActor +final class DesktopHostRegistrationLifecycle { + private struct PersistedIdentity: Codable, Equatable { + let tailnetName: String + let loginName: String + let dnsName: String + let hostName: String + let ipv4Address: String + let userID: Int64 + + init(_ identity: TailnetIdentity) { + tailnetName = identity.tailnetName + loginName = identity.loginName + dnsName = identity.dnsName + hostName = identity.hostName + ipv4Address = identity.ipv4Address + userID = identity.userID + } + + var identity: TailnetIdentity { + TailnetIdentity( + tailnetName: tailnetName, + loginName: loginName, + dnsName: dnsName, + hostName: hostName, + ipv4Address: ipv4Address, + userID: userID + ) + } + } + + private struct RegistrationTarget: Codable, Equatable { + private let persistedIdentity: PersistedIdentity + let hostID: String + let port: UInt16 + let publicationID: String + + init(identity: TailnetIdentity, hostID: String, port: UInt16, publicationID: String) { + persistedIdentity = PersistedIdentity(identity) + self.hostID = hostID + self.port = port + self.publicationID = publicationID + } + + var identity: TailnetIdentity { persistedIdentity.identity } + } + + private struct PublishedRegistration: Equatable { + let identity: TailnetIdentity + let hostID: String + let publicationID: String + let ownershipToken: String? + let usesLegacyCleanup: Bool + } + + private struct PersistedPublishedRegistration: Codable { + private let persistedIdentity: PersistedIdentity + let hostID: String + let publicationID: String + let usesLegacyCleanup: Bool + + init(_ registration: PublishedRegistration) { + persistedIdentity = PersistedIdentity(registration.identity) + hostID = registration.hostID + publicationID = registration.publicationID + usesLegacyCleanup = registration.usesLegacyCleanup + } + + var registration: PublishedRegistration { + PublishedRegistration( + identity: persistedIdentity.identity, + hostID: hostID, + publicationID: publicationID, + ownershipToken: nil, + usesLegacyCleanup: usesLegacyCleanup + ) + } + } + + private struct PersistedState: Codable { + var uncertainRegistrations: [RegistrationTarget] + var publishedRegistration: PersistedPublishedRegistration? + var pendingRemovals: [PersistedPublishedRegistration] + } + + private let coordinator: DesktopHostRegistrationCoordinator + private let createPublicationID: () -> String + private let stateStore: (any DesktopHostRegistrationStateStoring)? + private let recoveryScopeProvider: (() async throws -> DesktopHostRegistrationRecoveryScope)? + private var recoveryScope: DesktopHostRegistrationRecoveryScope? + private var stateLoaded = false + private var publishedRegistration: PublishedRegistration? + private var uncertainRegistrations: [RegistrationTarget] = [] + private var pendingRemovals: [PublishedRegistration] = [] + private var stateLoadError: Error? + private var lastPersistenceError: Error? + + init( + registration: any DesktopHostRegistering, + createPublicationID: @escaping () -> String = { UUID().uuidString }, + stateStore: (any DesktopHostRegistrationStateStoring)? = nil, + recoveryScopeProvider: (() async throws -> DesktopHostRegistrationRecoveryScope)? = nil + ) { + coordinator = DesktopHostRegistrationCoordinator(registration: registration) + self.createPublicationID = createPublicationID + self.stateStore = stateStore + self.recoveryScopeProvider = recoveryScopeProvider + } + + var hasDurableRecoveryState: Bool { + stateStore != nil && stateLoaded && stateLoadError == nil && lastPersistenceError == nil + && (!uncertainRegistrations.isEmpty || publishedRegistration != nil + || !pendingRemovals.isEmpty) + } + + var canTerminateAfterCleanupFailure: Bool { + let hasActiveState = + !uncertainRegistrations.isEmpty || publishedRegistration != nil || !pendingRemovals.isEmpty + return !hasActiveState || hasDurableRecoveryState + } + + func publish(identity: TailnetIdentity, port: UInt16) async throws { + try await loadStateIfNeeded() + let hostID = CrabfleetDesktopRegistration.hostID(identity: identity) + let existingTarget = uncertainRegistrations.first { + $0.hostID == hostID && $0.identity == identity && $0.port == port + } + let target = + existingTarget + ?? RegistrationTarget( + identity: identity, + hostID: hostID, + port: port, + publicationID: createPublicationID() + ) + if existingTarget == nil { + uncertainRegistrations.append(target) + try persistState() + } + let ownershipToken: String? + if existingTarget != nil { + ownershipToken = try await coordinator.recover( + identity: identity, + publicationID: target.publicationID + ) + guard ownershipToken != nil else { + uncertainRegistrations.removeAll { $0 == target } + try persistState() + throw DesktopHostRegistrationSupersededError() + } + } else { + do { + ownershipToken = try await coordinator.register( + identity: identity, + port: port, + publicationID: target.publicationID + ) + } catch { + if !(error is DesktopHostRegistrationResultUncertainError) { + uncertainRegistrations.removeAll { $0 == target } + try persistState() + } + throw error + } + } + uncertainRegistrations.removeAll { $0 == target } + if let publishedRegistration, publishedRegistration.hostID != hostID, + !pendingRemovals.contains(publishedRegistration) + { + pendingRemovals.append(publishedRegistration) + } + pendingRemovals.removeAll { $0.hostID == hostID } + publishedRegistration = PublishedRegistration( + identity: identity, + hostID: hostID, + publicationID: target.publicationID, + ownershipToken: ownershipToken, + usesLegacyCleanup: ownershipToken == nil + ) + try persistState() + } + + func removePublishedIdentities() async throws { + try await loadStateIfNeeded(requireRecoveryScope: false) + var firstError: Error? + let uncertainRegistrations = uncertainRegistrations + for target in uncertainRegistrations { + do { + let ownershipToken = try await coordinator.recover( + identity: target.identity, + publicationID: target.publicationID + ) + self.uncertainRegistrations.removeAll { $0 == target } + guard let ownershipToken else { + try persistState() + continue + } + let recovered = PublishedRegistration( + identity: target.identity, + hostID: target.hostID, + publicationID: target.publicationID, + ownershipToken: ownershipToken, + usesLegacyCleanup: false + ) + if !pendingRemovals.contains(recovered) { + pendingRemovals.append(recovered) + } + try persistState() + } catch { + firstError = firstError ?? error + } + } + + if let publishedRegistration { + if !pendingRemovals.contains(publishedRegistration) { + pendingRemovals.append(publishedRegistration) + } + self.publishedRegistration = nil + do { + try persistState() + } catch { + firstError = firstError ?? error + } + } + + let removals = pendingRemovals + for removal in removals { + var ownershipToken = removal.ownershipToken + if ownershipToken == nil, !removal.usesLegacyCleanup { + do { + guard + let recoveredToken = try await coordinator.recover( + identity: removal.identity, + publicationID: removal.publicationID + ) + else { + pendingRemovals.removeAll { $0 == removal } + try persistState() + continue + } + ownershipToken = recoveredToken + } catch { + firstError = firstError ?? error + continue + } + } + do { + try await coordinator.unregister( + identity: removal.identity, + ownershipToken: ownershipToken + ) + pendingRemovals.removeAll { $0 == removal } + try persistState() + } catch { + firstError = firstError ?? error + if removal.usesLegacyCleanup { + // A tokenless legacy DELETE cannot identify its publication. Try it + // only while stopping the process that published it; never retain it + // for a delayed retry that could delete a replacement publisher. + pendingRemovals.removeAll { $0 == removal } + do { + try persistState() + } catch { + firstError = firstError ?? error + } + } + } + } + if let firstError { throw firstError } + } + + private func loadStateIfNeeded(requireRecoveryScope: Bool = true) async throws { + guard let stateStore else { return } + if stateLoaded { + if let stateLoadError { throw stateLoadError } + if requireRecoveryScope, recoveryScope == nil { + try await loadRecoveryScope() + } + return + } + if !requireRecoveryScope, !stateStore.containsState() { + stateLoaded = true + return + } + try await loadRecoveryScope() + guard let recoveryScope else { + throw DesktopHostRegistrationPersistenceError.missingScope + } + var filteredLegacyState = false + do { + if let data = try stateStore.load(scope: recoveryScope) { + let state = try JSONDecoder().decode(PersistedState.self, from: data) + filteredLegacyState = + state.publishedRegistration?.usesLegacyCleanup == true + || state.pendingRemovals.contains { $0.usesLegacyCleanup } + uncertainRegistrations = state.uncertainRegistrations + publishedRegistration = state.publishedRegistration + .map(\.registration) + .flatMap { $0.usesLegacyCleanup ? nil : $0 } + pendingRemovals = state.pendingRemovals.map(\.registration) + .filter { !$0.usesLegacyCleanup } + } + stateLoaded = true + } catch { + stateLoadError = DesktopHostRegistrationPersistenceError.unreadableState + stateLoaded = true + throw DesktopHostRegistrationPersistenceError.unreadableState + } + if filteredLegacyState { + try persistState() + } + } + + private func loadRecoveryScope() async throws { + guard let recoveryScopeProvider else { + throw DesktopHostRegistrationPersistenceError.missingScope + } + recoveryScope = try await recoveryScopeProvider() + } + + private func persistState() throws { + guard let stateStore else { return } + guard stateLoaded, let recoveryScope else { + throw DesktopHostRegistrationPersistenceError.missingScope + } + let state = PersistedState( + uncertainRegistrations: uncertainRegistrations, + publishedRegistration: publishedRegistration + .flatMap { $0.usesLegacyCleanup ? nil : PersistedPublishedRegistration($0) }, + pendingRemovals: pendingRemovals + .filter { !$0.usesLegacyCleanup } + .map(PersistedPublishedRegistration.init) + ) + let hasState = + !state.uncertainRegistrations.isEmpty || state.publishedRegistration != nil + || !state.pendingRemovals.isEmpty + do { + let data = hasState ? try JSONEncoder().encode(state) : nil + try stateStore.save(data, scope: recoveryScope) + lastPersistenceError = nil + } catch { + lastPersistenceError = error + throw error + } + } +} + @MainActor final class PrivateMacShareController: ObservableObject { enum RegistryPhase: Equatable { case notConfigured + case notPublished case registering case registered case failed(String) @@ -23,6 +466,7 @@ final class PrivateMacShareController: ObservableObject { var detail: String { switch self { case .notConfigured: "Not configured" + case .notPublished: "Not published" case .registering: "Registering" case .registered: "Published" case .failed(let message): message @@ -98,22 +542,44 @@ final class PrivateMacShareController: ObservableObject { private let runner: (any TailscaleCommandRunning)? private let desktopRegistration: (any DesktopHostRegistering)? + private let desktopRegistrationLifecycle: DesktopHostRegistrationLifecycle? private let runnerInitializationError: Error? private let defaults: UserDefaults private var capture: MacScreenCapture? private var server: TailnetRFBServer? private var clipboardBridge: HostClipboardBridge? - private var serverGeneration: UUID? + private var activeIdentity: TailnetIdentity? + private var lifecycleGeneration: UInt64 = 0 + private var serverGeneration: UInt64? private var registrationTask: Task? + private var registryOperationGeneration: UInt64 = 0 + private var publishingServerGeneration: UInt64? + private var refreshWaiters: [CheckedContinuation] = [] + private let stopCoordinator = PrivateMacShareStopCoordinator() init( runner: (any TailscaleCommandRunning)? = nil, desktopRegistration: (any DesktopHostRegistering)? = CrabfleetDesktopRegistration(), + registrationLifecycle: DesktopHostRegistrationLifecycle? = nil, defaults: UserDefaults = .standard ) { self.desktopRegistration = desktopRegistration + let registrationStateStore = UserDefaultsDesktopHostRegistrationStateStore(defaults: defaults) + let recoveryScopeProvider = (desktopRegistration as? any DesktopHostRegistrationRecoveryScoping) + .map { registration in + { try await registration.recoveryScope() } + } + desktopRegistrationLifecycle = + registrationLifecycle + ?? desktopRegistration.map { + DesktopHostRegistrationLifecycle( + registration: $0, + stateStore: registrationStateStore, + recoveryScopeProvider: recoveryScopeProvider + ) + } self.defaults = defaults - registryPhase = desktopRegistration == nil ? .notConfigured : .registering + registryPhase = desktopRegistration == nil ? .notConfigured : .notPublished let savedDisplayID = defaults.object(forKey: Self.selectedDisplayDefaultsKey) as? Int selectedDisplayID = savedDisplayID.map(CGDirectDisplayID.init) ?? CGMainDisplayID() clipboardSyncEnabled = @@ -139,7 +605,7 @@ final class PrivateMacShareController: ObservableObject { } var canStart: Bool { - phase == .idle + phase == .idle && !isRefreshing && PrivateMacSharePermissionPolicy.canStart( identityAvailable: identity != nil, screenRecordingGranted: screenRecordingGranted @@ -147,14 +613,23 @@ final class PrivateMacShareController: ObservableObject { } func refresh() async { - guard !isRefreshing, phase != .starting, phase != .stopping else { return } + if isRefreshing { + await waitForRefreshCompletion() + return + } + guard phase != .starting, phase != .stopping else { return } isRefreshing = true + defer { finishRefresh() } notice = nil - await loadIdentity() + do { + identity = try await fetchIdentity() + } catch { + identity = nil + notice = error.localizedDescription + } refreshPermissions() await refreshDisplays() launchAtLoginEnabled = SMAppService.mainApp.status == .enabled - isRefreshing = false } /// Registers or removes the login item and remembers to auto-start the @@ -208,12 +683,25 @@ final class PrivateMacShareController: ObservableObject { func start() async { guard phase == .idle else { return } + let generation = beginLifecycleTransition() phase = .starting notice = nil connectedPeer = nil streamStats = nil registryPhase = desktopRegistration == nil ? .notConfigured : .registering - await loadIdentity() + await waitForRefreshCompletion() + guard canContinueStarting(generation) else { return } + do { + let loadedIdentity = try await fetchIdentity() + guard canContinueStarting(generation) else { return } + identity = loadedIdentity + } catch { + guard canContinueStarting(generation) else { return } + identity = nil + phase = .failed + notice = error.localizedDescription + return + } refreshPermissions() guard let identity else { @@ -223,6 +711,7 @@ final class PrivateMacShareController: ObservableObject { } guard screenRecordingGranted else { phase = .idle + registryPhase = desktopRegistration == nil ? .notConfigured : .notPublished notice = PrivateMacShareError.screenRecordingDenied.localizedDescription return } @@ -237,9 +726,12 @@ final class PrivateMacShareController: ObservableObject { let capture = MacScreenCapture() do { let descriptor = try await capture.start(displayID: selectedDisplayID) + guard canContinueStarting(generation) else { + await capture.stop() + return + } let input = MacRemoteInputController(descriptor: descriptor) let bridge = clipboardSyncEnabled ? HostClipboardBridge() : nil - let generation = UUID() serverGeneration = generation let server = TailnetRFBServer( identity: identity, @@ -258,7 +750,12 @@ final class PrivateMacShareController: ObservableObject { self.capture = capture self.server = server self.clipboardBridge = bridge + activeIdentity = identity } catch { + guard canContinueStarting(generation) else { + await capture.stop() + return + } serverGeneration = nil await capture.stop() phase = .failed @@ -267,12 +764,19 @@ final class PrivateMacShareController: ObservableObject { } func stop() async { + await stopCoordinator.perform { [weak self] in + await self?.performStop() + } + } + + private func performStop() async { guard phase.isRunning || phase == .failed else { return } + let generation = beginLifecycleTransition() phase = .stopping connectedPeer = nil streamStats = nil - registrationTask?.cancel() - registrationTask = nil + let registrationTask = self.registrationTask + self.registrationTask = nil serverGeneration = nil server?.stop() server = nil @@ -281,9 +785,28 @@ final class PrivateMacShareController: ObservableObject { let capture = capture self.capture = nil await capture?.stop() + activeIdentity = nil + removeDesktopHost(after: registrationTask) + guard isCurrent(generation) else { return } phase = .idle } + func stopAndWaitForCleanup() async -> Bool { + await stop() + let cleanupTask = registrationTask + await cleanupTask?.value + guard let desktopRegistrationLifecycle else { return true } + do { + try await desktopRegistrationLifecycle.removePublishedIdentities() + registryPhase = .notPublished + return true + } catch { + registryPhase = .failed(error.localizedDescription) + notice = error.localizedDescription + return desktopRegistrationLifecycle.canTerminateAfterCleanupFailure + } + } + func openPrivacySettings(_ pane: PrivacyPane) { let value: String switch pane { @@ -305,25 +828,16 @@ final class PrivateMacShareController: ObservableObject { case accessibility } - private func loadIdentity() async { + private func fetchIdentity() async throws -> TailnetIdentity { guard let runner else { - identity = nil - notice = - (runnerInitializationError ?? PrivateMacShareError.tailscaleNotInstalled) - .localizedDescription - return - } - do { - let result = try await runner.run(arguments: ["status", "--json"]) - let document = try JSONDecoder().decode( - TailscaleStatusDocument.self, - from: Data(result.standardOutput.utf8) - ) - identity = try TailnetIdentityPolicy.identity(from: document) - } catch { - identity = nil - notice = error.localizedDescription + throw runnerInitializationError ?? PrivateMacShareError.tailscaleNotInstalled } + let result = try await runner.run(arguments: ["status", "--json"]) + let document = try JSONDecoder().decode( + TailscaleStatusDocument.self, + from: Data(result.standardOutput.utf8) + ) + return try TailnetIdentityPolicy.identity(from: document) } private func refreshPermissions() { @@ -344,7 +858,23 @@ final class PrivateMacShareController: ObservableObject { } } - private func handle(_ event: TailnetRFBServerEvent, generation: UUID) { + private func waitForRefreshCompletion() async { + guard isRefreshing else { return } + await withCheckedContinuation { continuation in + refreshWaiters.append(continuation) + } + } + + private func finishRefresh() { + isRefreshing = false + let waiters = refreshWaiters + refreshWaiters.removeAll() + for waiter in waiters { + waiter.resume() + } + } + + private func handle(_ event: TailnetRFBServerEvent, generation: UInt64) { guard serverGeneration == generation else { return } switch event { case .listening: @@ -368,8 +898,7 @@ final class PrivateMacShareController: ObservableObject { connectedPeer = nil streamStats = nil case .listenerFailed(let message): - registrationTask?.cancel() - registrationTask = nil + let pendingRegistration = registrationTask phase = .failed connectedPeer = nil streamStats = nil @@ -382,6 +911,7 @@ final class PrivateMacShareController: ObservableObject { let failedCapture = capture capture = nil Task { await failedCapture?.stop() } + removeDesktopHost(after: pendingRegistration) case .sessionFailed(let message): phase = .sharing connectedPeer = nil @@ -390,24 +920,94 @@ final class PrivateMacShareController: ObservableObject { } } - private func registerDesktopHost(generation: UUID) { - registrationTask?.cancel() - guard let desktopRegistration, let identity else { - registryPhase = .notConfigured + private func registerDesktopHost(generation: UInt64) { + guard publishingServerGeneration != generation else { return } + guard let desktopRegistrationLifecycle, let identity = activeIdentity else { + registryPhase = desktopRegistration == nil ? .notConfigured : .notPublished return } + let pendingOperation = registrationTask + let operationGeneration = beginRegistryOperation() + publishingServerGeneration = generation registryPhase = .registering registrationTask = Task { [weak self] in do { - try await desktopRegistration.register(identity: identity, port: Self.port) - guard !Task.isCancelled, self?.serverGeneration == generation else { return } + await pendingOperation?.value + try await desktopRegistrationLifecycle.publish(identity: identity, port: Self.port) + guard + self?.isCurrentRegistryOperation(operationGeneration) == true, + self?.serverGeneration == generation + else { return } self?.registryPhase = .registered - } catch is CancellationError { - return } catch { - guard self?.serverGeneration == generation else { return } + guard + self?.isCurrentRegistryOperation(operationGeneration) == true, + self?.serverGeneration == generation + else { return } + self?.registryPhase = .failed(error.localizedDescription) + } + if self?.publishingServerGeneration == generation { + self?.publishingServerGeneration = nil + } + self?.finishRegistryOperation(operationGeneration) + } + } + + private func removeDesktopHost(after pendingRegistration: Task?) { + guard let desktopRegistrationLifecycle else { + registrationTask = nil + registryPhase = desktopRegistration == nil ? .notConfigured : .notPublished + return + } + publishingServerGeneration = nil + let operationGeneration = beginRegistryOperation() + registrationTask = Task { [weak self] in + await pendingRegistration?.value + do { + try await desktopRegistrationLifecycle.removePublishedIdentities() + guard self?.isCurrentRegistryOperation(operationGeneration) == true else { return } + self?.registryPhase = .notPublished + } catch { + guard self?.isCurrentRegistryOperation(operationGeneration) == true else { return } self?.registryPhase = .failed(error.localizedDescription) + self?.notice = error.localizedDescription } + self?.finishRegistryOperation(operationGeneration) + } + } + + @discardableResult + private func beginLifecycleTransition() -> UInt64 { + lifecycleGeneration &+= 1 + return lifecycleGeneration + } + + private func isCurrent(_ generation: UInt64) -> Bool { + lifecycleGeneration == generation + } + + private func canContinueStarting(_ generation: UInt64) -> Bool { + guard isCurrent(generation), phase == .starting else { return false } + guard !Task.isCancelled else { + phase = .idle + registryPhase = desktopRegistration == nil ? .notConfigured : .notPublished + return false } + return true + } + + @discardableResult + private func beginRegistryOperation() -> UInt64 { + registryOperationGeneration &+= 1 + return registryOperationGeneration + } + + private func isCurrentRegistryOperation(_ generation: UInt64) -> Bool { + registryOperationGeneration == generation + } + + private func finishRegistryOperation(_ generation: UInt64) { + guard isCurrentRegistryOperation(generation) else { return } + registrationTask = nil } } diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/SubprocessEnvironment.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/SubprocessEnvironment.swift new file mode 100644 index 00000000..609d1604 --- /dev/null +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/SubprocessEnvironment.swift @@ -0,0 +1,59 @@ +import Darwin +import Foundation + +enum SubprocessEnvironment { + private static let inheritedKeys = [ + "HOME", + "LANG", + "LC_ALL", + "LC_CTYPE", + "TMPDIR", + ] + + static let safePath = "/usr/bin:/bin:/usr/sbin:/sbin" + + static func minimal( + from source: [String: String], + includeSSHAgent: Bool = false, + additionalInheritedKeys: [String] = [], + additionalInheritedPathKeys: [String] = [], + overrides: [String: String] = [:] + ) -> [String: String] { + var environment = Dictionary( + uniqueKeysWithValues: inheritedKeys.compactMap { key in + source[key].map { (key, $0) } + } + ) + for key in additionalInheritedKeys { + environment[key] = source[key] + } + for key in additionalInheritedPathKeys { + environment.removeValue(forKey: key) + if let value = source[key], isSafeAbsolutePath(value) { + environment[key] = value + } + } + environment["PATH"] = safePath + if includeSSHAgent, let socket = source["SSH_AUTH_SOCK"], !socket.isEmpty { + environment["SSH_AUTH_SOCK"] = socket + } + for (key, value) in overrides { + environment[key] = value + } + return environment + } + + private static func isSafeAbsolutePath(_ value: String) -> Bool { + guard + value.hasPrefix("/"), + value.utf8.count < Int(PATH_MAX), + value.unicodeScalars.allSatisfy({ + !CharacterSet.controlCharacters.contains($0) + }) + else { + return false + } + return !value.split(separator: "/", omittingEmptySubsequences: false) + .contains { $0 == "." || $0 == ".." } + } +} diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetIdentity.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetIdentity.swift index 77103a02..82f162ce 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetIdentity.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetIdentity.swift @@ -11,13 +11,18 @@ protocol TailscaleCommandRunning: Sendable { struct SystemTailscaleCommandRunner: TailscaleCommandRunning { private static let maximumOutputBytes = 4 * 1_024 * 1_024 + private static let defaultTimeout: TimeInterval = 15 static let executableCandidates = [ "/Applications/Tailscale.app/Contents/MacOS/Tailscale", ] let executableURL: URL + let timeout: TimeInterval - init(fileManager: FileManager = .default) throws { + init( + fileManager: FileManager = .default, + timeout: TimeInterval = Self.defaultTimeout + ) throws { guard let path = Self.executableCandidates.first(where: { Self.isTrustedExecutable(atPath: $0, fileManager: fileManager) @@ -26,89 +31,40 @@ struct SystemTailscaleCommandRunner: TailscaleCommandRunning { throw PrivateMacShareError.tailscaleNotInstalled } executableURL = URL(fileURLWithPath: path) + self.timeout = timeout } - init(executableURL: URL) { + init(executableURL: URL, timeout: TimeInterval = Self.defaultTimeout) { self.executableURL = executableURL + self.timeout = timeout } func run(arguments: [String]) async throws -> TailscaleCommandResult { - try await withCheckedThrowingContinuation { continuation in - DispatchQueue.global(qos: .userInitiated).async { - let process = Process() - let outputPipe = Pipe() - let errorPipe = Pipe() - let readGroup = DispatchGroup() - let capture = CommandCapture() - - process.executableURL = executableURL - process.arguments = arguments - process.standardOutput = outputPipe - process.standardError = errorPipe - process.qualityOfService = .userInitiated - - process.environment = Self.commandEnvironment( - from: ProcessInfo.processInfo.environment - ) - - readGroup.enter() - DispatchQueue.global(qos: .userInitiated).async { - capture.setStandardOutput(outputPipe.fileHandleForReading.readDataToEndOfFile()) - readGroup.leave() - } - - readGroup.enter() + let execution = TailscaleCommandExecution( + executableURL: executableURL, + arguments: arguments, + environment: Self.commandEnvironment(from: ProcessInfo.processInfo.environment), + timeout: timeout, + maximumOutputBytes: Self.maximumOutputBytes + ) + return try await withTaskCancellationHandler { + let result = try await withCheckedThrowingContinuation { continuation in DispatchQueue.global(qos: .userInitiated).async { - capture.setStandardError(errorPipe.fileHandleForReading.readDataToEndOfFile()) - readGroup.leave() - } - - do { - try process.run() - process.waitUntilExit() - readGroup.wait() - - let (outputData, errorData) = capture.values() - guard - outputData.count <= Self.maximumOutputBytes, - errorData.count <= Self.maximumOutputBytes - else { - continuation.resume(throwing: PrivateMacShareError.commandOutputTooLarge) - return - } - - let result = TailscaleCommandResult( - standardOutput: String(decoding: outputData, as: UTF8.self), - standardError: String(decoding: errorData, as: UTF8.self) - ) - guard process.terminationStatus == 0 else { - let message = result.standardError.trimmingCharacters(in: .whitespacesAndNewlines) - continuation.resume( - throwing: PrivateMacShareError.commandFailed( - status: process.terminationStatus, - message: String(message.prefix(500)) - )) - return - } - continuation.resume(returning: result) - } catch { - outputPipe.fileHandleForWriting.closeFile() - errorPipe.fileHandleForWriting.closeFile() - readGroup.wait() - continuation.resume(throwing: error) + continuation.resume(with: Result { try execution.run() }) } } + try Task.checkCancellation() + return result + } onCancel: { + execution.cancel() } } static func commandEnvironment(from source: [String: String]) -> [String: String] { - var environment = source - for key in environment.keys - where key.hasPrefix("TS_") || key.hasPrefix("TAILSCALE_") { - environment.removeValue(forKey: key) - } - environment["TAILSCALE_BE_CLI"] = "1" - return environment + SubprocessEnvironment.minimal( + from: source, + overrides: ["TAILSCALE_BE_CLI": "1"] + ) } static func isTrustedExecutable( @@ -131,27 +87,294 @@ struct SystemTailscaleCommandRunner: TailscaleCommandRunning { } } -private final class CommandCapture: @unchecked Sendable { +private final class TailscaleCommandExecution: @unchecked Sendable { + private static let processDrainTimeout: DispatchTimeInterval = .milliseconds(250) + private static let terminationGracePeriod: TimeInterval = 0.5 + + private enum StopReason { + case cancelled + case outputTooLarge + case timedOut + } + private let lock = NSLock() + private let outputPipe = Pipe() + private let errorPipe = Pipe() + private let readGroup = DispatchGroup() + private let executableURL: URL + private let arguments: [String] + private let environment: [String: String] + private let timeout: TimeInterval + private let maximumOutputBytes: Int + private var processID: pid_t? + private var stopReason: StopReason? + private var captureShouldStop = false private var standardOutput = Data() private var standardError = Data() - func setStandardOutput(_ data: Data) { + init( + executableURL: URL, + arguments: [String], + environment: [String: String], + timeout: TimeInterval, + maximumOutputBytes: Int + ) { + self.executableURL = executableURL + self.arguments = arguments + self.environment = environment + self.timeout = max(0.1, timeout) + self.maximumOutputBytes = maximumOutputBytes + } + + func run() throws -> TailscaleCommandResult { + if currentStopReason() != nil { throw CancellationError() } + startCapture(pipe: outputPipe, isStandardOutput: true) + startCapture(pipe: errorPipe, isStandardOutput: false) + + let pid: pid_t + do { + pid = try spawn() + } catch { + outputPipe.fileHandleForWriting.closeFile() + errorPipe.fileHandleForWriting.closeFile() + readGroup.wait() + throw error + } + + setProcessID(pid) + outputPipe.fileHandleForWriting.closeFile() + errorPipe.fileHandleForWriting.closeFile() + if currentStopReason() != nil { signalProcessGroup(pid, signal: SIGTERM) } + + let clock = ContinuousClock() + let deadline = clock.now.advanced(by: .seconds(timeout)) + var terminationDeadline: ContinuousClock.Instant? + var waitStatus: Int32 = 0 + while true { + let waitResult = Darwin.waitpid(pid, &waitStatus, WNOHANG) + if waitResult == pid { break } + if waitResult == -1 { + if errno == EINTR { continue } + throw POSIXError(POSIXErrorCode(rawValue: errno) ?? .ECHILD) + } + + if currentStopReason() != nil { + if terminationDeadline == nil { + signalProcessGroup(pid, signal: SIGTERM) + terminationDeadline = clock.now.advanced(by: .seconds(Self.terminationGracePeriod)) + } else if clock.now >= terminationDeadline! { + signalProcessGroup(pid, signal: SIGKILL) + } + } else if clock.now >= deadline { + stop(.timedOut) + } + Thread.sleep(forTimeInterval: 0.01) + } + + if currentStopReason() != nil { + signalProcessGroup(pid, signal: SIGKILL) + } + clearProcessID(pid) + finishCapture() + + switch currentStopReason() { + case .cancelled: + throw CancellationError() + case .outputTooLarge: + throw PrivateMacShareError.commandOutputTooLarge + case .timedOut: + throw PrivateMacShareError.commandTimedOut + case nil: + break + } + + let result = values() + let terminationStatus = Self.terminationStatus(from: waitStatus) + guard terminationStatus == 0 else { + let message = result.standardError.trimmingCharacters(in: .whitespacesAndNewlines) + throw PrivateMacShareError.commandFailed( + status: terminationStatus, + message: String(message.prefix(500)) + ) + } + return result + } + + func cancel() { + stop(.cancelled) + } + + private func startCapture(pipe: Pipe, isStandardOutput: Bool) { + readGroup.enter() + DispatchQueue.global(qos: .userInitiated).async { [self] in + defer { readGroup.leave() } + let fileDescriptor = pipe.fileHandleForReading.fileDescriptor + var data = Data() + var buffer = [UInt8](repeating: 0, count: 64 * 1_024) + while !shouldStopCapture() { + var descriptor = pollfd(fd: fileDescriptor, events: Int16(POLLIN | POLLHUP), revents: 0) + let pollResult = Darwin.poll(&descriptor, 1, 50) + if pollResult == 0 { continue } + if pollResult < 0 { + if errno == EINTR { continue } + break + } + + let bytesRead = buffer.withUnsafeMutableBytes { + Darwin.read(fileDescriptor, $0.baseAddress, $0.count) + } + if bytesRead == 0 { break } + if bytesRead < 0 { + if errno == EINTR || errno == EAGAIN { continue } + break + } + guard bytesRead <= maximumOutputBytes - data.count else { + stop(.outputTooLarge) + break + } + data.append(buffer, count: bytesRead) + } + setCaptured(data, isStandardOutput: isStandardOutput) + } + } + + private func setCaptured(_ data: Data, isStandardOutput: Bool) { lock.lock() - standardOutput = data + if isStandardOutput { + standardOutput = data + } else { + standardError = data + } lock.unlock() } - func setStandardError(_ data: Data) { + private func values() -> TailscaleCommandResult { lock.lock() - standardError = data - lock.unlock() + defer { lock.unlock() } + return TailscaleCommandResult( + standardOutput: String(decoding: standardOutput, as: UTF8.self), + standardError: String(decoding: standardError, as: UTF8.self) + ) + } + + private func currentStopReason() -> StopReason? { + lock.lock() + defer { lock.unlock() } + return stopReason } - func values() -> (Data, Data) { + private func shouldStopCapture() -> Bool { lock.lock() defer { lock.unlock() } - return (standardOutput, standardError) + return captureShouldStop + } + + private func stopCapture() { + lock.lock() + captureShouldStop = true + lock.unlock() + } + + private func stop(_ reason: StopReason) { + lock.lock() + if stopReason == nil { stopReason = reason } + let pid = processID + lock.unlock() + if let pid { + signalProcessGroup(pid, signal: SIGTERM) + } + } + + private func spawn() throws -> pid_t { + var fileActions: posix_spawn_file_actions_t? + var attributes: posix_spawnattr_t? + guard posix_spawn_file_actions_init(&fileActions) == 0 else { + throw POSIXError(.ENOMEM) + } + defer { posix_spawn_file_actions_destroy(&fileActions) } + guard posix_spawnattr_init(&attributes) == 0 else { + throw POSIXError(.ENOMEM) + } + defer { posix_spawnattr_destroy(&attributes) } + + let outputRead = outputPipe.fileHandleForReading.fileDescriptor + let outputWrite = outputPipe.fileHandleForWriting.fileDescriptor + let errorRead = errorPipe.fileHandleForReading.fileDescriptor + let errorWrite = errorPipe.fileHandleForWriting.fileDescriptor + posix_spawn_file_actions_addclose(&fileActions, outputRead) + posix_spawn_file_actions_addclose(&fileActions, errorRead) + posix_spawn_file_actions_adddup2(&fileActions, outputWrite, STDOUT_FILENO) + posix_spawn_file_actions_adddup2(&fileActions, errorWrite, STDERR_FILENO) + posix_spawn_file_actions_addclose(&fileActions, outputWrite) + posix_spawn_file_actions_addclose(&fileActions, errorWrite) + + let flags = Int16(POSIX_SPAWN_SETPGROUP) + posix_spawnattr_setflags(&attributes, flags) + posix_spawnattr_setpgroup(&attributes, 0) + + let argv = [executableURL.path] + arguments + let env = environment.map { "\($0.key)=\($0.value)" } + return try withCStringArray(argv) { argumentPointers in + try withCStringArray(env) { environmentPointers in + var pid: pid_t = 0 + let result = posix_spawn( + &pid, + executableURL.path, + &fileActions, + &attributes, + argumentPointers, + environmentPointers + ) + guard result == 0 else { + throw POSIXError(POSIXErrorCode(rawValue: result) ?? .EINVAL) + } + return pid + } + } + } + + private func withCStringArray( + _ strings: [String], + body: ([UnsafeMutablePointer?]) throws -> Result + ) rethrows -> Result { + let pointers = strings.map { strdup($0) } + defer { pointers.forEach { free($0) } } + return try body(pointers + [nil]) + } + + private func setProcessID(_ pid: pid_t) { + lock.lock() + processID = pid + lock.unlock() + } + + private func clearProcessID(_ pid: pid_t) { + lock.lock() + if processID == pid { processID = nil } + lock.unlock() + } + + private func signalProcessGroup(_ pid: pid_t, signal: Int32) { + guard pid > 0 else { return } + if Darwin.kill(-pid, signal) != 0, errno != ESRCH { + return + } + } + + private static func terminationStatus(from waitStatus: Int32) -> Int32 { + let signal = waitStatus & 0x7f + if signal == 0 { + return (waitStatus >> 8) & 0xff + } + return 128 + signal + } + + private func finishCapture() { + guard readGroup.wait(timeout: .now() + Self.processDrainTimeout) == .timedOut else { + return + } + stopCapture() + _ = readGroup.wait(timeout: .now() + Self.processDrainTimeout) } } @@ -317,7 +540,10 @@ struct TailnetPeerAuthorizer: TailnetPeerAuthorizing, Sendable { let expectedIdentity: TailnetIdentity func authorize(remoteAddress: String) async -> Bool { - guard TailnetIdentityPolicy.isTailscaleIPv4(remoteAddress) else { return false } + guard + TailnetIdentityPolicy.isTailscaleIPv4(remoteAddress), + remoteAddress != expectedIdentity.ipv4Address + else { return false } do { let result = try await runner.run(arguments: ["whois", "--json", remoteAddress]) let document = try JSONDecoder().decode( @@ -347,6 +573,7 @@ enum PrivateMacShareError: LocalizedError, Equatable { case accessibilityDenied case commandFailed(status: Int32, message: String) case commandOutputTooLarge + case commandTimedOut case captureUnavailable case listenerFailed(String) case protocolError(String) @@ -375,6 +602,8 @@ enum PrivateMacShareError: LocalizedError, Equatable { : "Tailscale exited with status \(status): \(message)" case .commandOutputTooLarge: "Tailscale returned more status data than Crabfleet will accept." + case .commandTimedOut: + "Tailscale did not respond before the command deadline." case .captureUnavailable: "Crabfleet could not capture the main display." case .listenerFailed(let message): diff --git a/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetRFBServer.swift b/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetRFBServer.swift index e3b00699..3480cfb7 100644 --- a/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetRFBServer.swift +++ b/macos/CrabfleetMac/Sources/CrabfleetMac/TailnetRFBServer.swift @@ -31,6 +31,7 @@ final class TailnetRFBServer: @unchecked Sendable { private let clipboard: (any HostClipboardSyncing)? private let peerAuthorizer: (any TailnetPeerAuthorizing)? private let port: UInt16 + private let handshakeTimeout: Duration private let queue = DispatchQueue(label: "org.openclaw.crabfleet.rfb-listener") private let lock = NSLock() private let eventHandler: EventHandler @@ -47,6 +48,7 @@ final class TailnetRFBServer: @unchecked Sendable { clipboard: (any HostClipboardSyncing)? = nil, peerAuthorizer: (any TailnetPeerAuthorizing)? = nil, port: UInt16, + handshakeTimeout: Duration = .seconds(10), eventHandler: @escaping EventHandler ) { self.identity = identity @@ -57,6 +59,7 @@ final class TailnetRFBServer: @unchecked Sendable { self.clipboard = clipboard self.peerAuthorizer = peerAuthorizer self.port = port + self.handshakeTimeout = handshakeTimeout self.eventHandler = eventHandler } @@ -134,6 +137,7 @@ final class TailnetRFBServer: @unchecked Sendable { clipboard: clipboard, requiredLocalAddress: identity.ipv4Address, desktopName: "Crabfleet — \(identity.hostName)", + handshakeTimeout: handshakeTimeout, viewOnly: false, didAuthorize: { [weak capture] in capture?.setConsumerActive(true) }, eventHandler: eventHandler, @@ -170,9 +174,11 @@ private final class RFBHostSession: @unchecked Sendable { private let capture: MacScreenCapture private let descriptor: CapturedDisplayDescriptor private let input: any RemoteInputForwarding + private let inputGate: RemoteInputSessionGate private let clipboard: (any HostClipboardSyncing)? private let requiredLocalAddress: String private let desktopName: String + private let handshakeTimeout: Duration private let didAuthorize: @Sendable () -> Void private let queue = DispatchQueue(label: "org.openclaw.crabfleet.rfb-session") private let eventHandler: TailnetRFBServer.EventHandler @@ -180,6 +186,8 @@ private final class RFBHostSession: @unchecked Sendable { private let lock = NSLock() private var started = false private var finished = false + private var handshakeFinished = false + private var handshakeTimedOut = false private var task: Task? private var pushIO: RFBConnectionIO? @@ -193,7 +201,6 @@ private final class RFBHostSession: @unchecked Sendable { private var clientClipboardCaps: VNCExtendedClipboardCaps? private var currentWidth: Int private var currentHeight: Int - private var viewOnly: Bool private var videoEncoder: MacVideoEncoder? private var videoFrameMailbox: VideoMailbox? private var videoPixelMailbox: VideoMailbox? @@ -213,6 +220,7 @@ private final class RFBHostSession: @unchecked Sendable { clipboard: (any HostClipboardSyncing)?, requiredLocalAddress: String, desktopName: String, + handshakeTimeout: Duration, viewOnly: Bool, didAuthorize: @escaping @Sendable () -> Void, eventHandler: @escaping TailnetRFBServer.EventHandler, @@ -223,10 +231,11 @@ private final class RFBHostSession: @unchecked Sendable { self.capture = capture self.descriptor = descriptor self.input = input + inputGate = RemoteInputSessionGate(input: input, viewOnly: viewOnly) self.clipboard = clipboard self.requiredLocalAddress = requiredLocalAddress self.desktopName = desktopName - self.viewOnly = viewOnly + self.handshakeTimeout = handshakeTimeout self.didAuthorize = didAuthorize self.eventHandler = eventHandler self.didFinish = didFinish @@ -258,11 +267,7 @@ private final class RFBHostSession: @unchecked Sendable { } func setViewOnly(_ enabled: Bool) { - let shouldReleaseInput = withLock { () -> Bool in - defer { viewOnly = enabled } - return enabled && !viewOnly - } - if shouldReleaseInput { input.releaseAllInput() } + inputGate.setViewOnly(enabled) } private func beginProtocolIfNeeded() { @@ -310,7 +315,7 @@ private final class RFBHostSession: @unchecked Sendable { didAuthorize() let io = RFBConnectionIO(connection: connection) - try await handshake(io: io) + try await handshakeBeforeDeadline(io: io) withLock { pushIO = io } attachClipboard() eventHandler(.connected(remoteAddress)) @@ -351,6 +356,45 @@ private final class RFBHostSession: @unchecked Sendable { )) } + private func handshakeBeforeDeadline(io: RFBConnectionIO) async throws { + let deadlineTask = Task { [weak self, handshakeTimeout] in + do { + try await Task.sleep(for: handshakeTimeout) + } catch { + return + } + self?.expireHandshake() + } + do { + try await handshake(io: io) + deadlineTask.cancel() + let timedOut = withLock { () -> Bool in + handshakeFinished = true + return handshakeTimedOut + } + guard !timedOut else { + throw PrivateMacShareError.protocolError("RFB handshake timed out") + } + } catch { + deadlineTask.cancel() + if withLock({ handshakeTimedOut }) { + throw PrivateMacShareError.protocolError("RFB handshake timed out") + } + throw error + } + } + + private func expireHandshake() { + let shouldCancel = withLock { () -> Bool in + guard !finished, !handshakeFinished else { return false } + handshakeTimedOut = true + return true + } + if shouldCancel { + finish(event: .sessionFailed("RFB handshake timed out")) + } + } + private func messageLoop(io: RFBConnectionIO) async throws { var hasSentJPEGFrame = false var lastSentJPEGSequence: UInt64 = 0 @@ -476,23 +520,15 @@ private final class RFBHostSession: @unchecked Sendable { case 4: // KeyEvent let payload = try await io.readExactly(7) - withLock { - if !viewOnly { - input.keyEvent(down: payload[0] != 0, keysym: payload.readUInt32(at: 3)) - } - } + inputGate.keyEvent(down: payload[0] != 0, keysym: payload.readUInt32(at: 3)) case 5: // PointerEvent let payload = try await io.readExactly(5) - withLock { - if !viewOnly { - input.pointerEvent( - buttonMask: payload[0], - x: payload.readUInt16(at: 1), - y: payload.readUInt16(at: 3) - ) - } - } + inputGate.pointerEvent( + buttonMask: payload[0], + x: payload.readUInt16(at: 1), + y: payload.readUInt16(at: 3) + ) case 6: // ClientCutText try await receiveClientCutText(io: io) @@ -588,27 +624,11 @@ private final class RFBHostSession: @unchecked Sendable { } guard let io else { return } - let payload: Data? - if extendedNegotiated, let caps, caps.supportsText { - let wireByteCount = VNCExtendedClipboard.wireTextByteCount(text) - switch VNCExtendedClipboard.textRoute(wireByteCount: wireByteCount, caps: caps) { - case .provide: - payload = (try? VNCExtendedClipboard.encodeProvide(text: text)).map { - VNCExtendedClipboard.frame(messageType: 3, body: $0) - } - case .notify: - payload = VNCExtendedClipboard.frame( - messageType: 3, - body: VNCExtendedClipboard.encodeNotify(hasText: true) - ) - case .legacy: - payload = RFBWire.legacyServerCutText(text: text) - } - } else { - // Legacy path: silently skip text that cannot survive Latin-1. - payload = RFBWire.legacyServerCutText(text: text) - } - + let payload = RFBWire.hostClipboardPayload( + text: text, + extendedNegotiated: extendedNegotiated, + caps: caps + ) guard let payload else { return } Task { try? await io.send(payload) @@ -1005,6 +1025,7 @@ private final class RFBHostSession: @unchecked Sendable { finishPixelMailbox() let encoder = replaceVideoEncoder(with: nil) encoder?.invalidate() + inputGate.finish() clipboard?.detach() connection.cancel() guard encoder != nil else { @@ -1035,6 +1056,34 @@ private final class RFBHostSession: @unchecked Sendable { } } +extension RFBWire { + static func hostClipboardPayload( + text: String, + extendedNegotiated: Bool, + caps: VNCExtendedClipboardCaps? + ) -> Data? { + guard extendedNegotiated, let caps, caps.supportsText else { + // Legacy path: silently skip text that cannot survive Latin-1. + return legacyServerCutText(text: text) + } + + let wireByteCount = VNCExtendedClipboard.wireTextByteCount(text) + switch VNCExtendedClipboard.textRoute(wireByteCount: wireByteCount, caps: caps) { + case .provide: + return (try? VNCExtendedClipboard.encodeProvide(text: text)).map { + VNCExtendedClipboard.frame(messageType: 3, body: $0) + } + case .notify: + return VNCExtendedClipboard.frame( + messageType: 3, + body: VNCExtendedClipboard.encodeNotify(hasText: true) + ) + case .legacy: + return legacyServerCutText(text: text) + } + } +} + private struct RFBConnectionIO: Sendable { let connection: NWConnection diff --git a/macos/CrabfleetMac/Tests/CrabfleetMacTests/FleetModelsTests.swift b/macos/CrabfleetMac/Tests/CrabfleetMacTests/FleetModelsTests.swift index a2bad09c..bd75521d 100644 --- a/macos/CrabfleetMac/Tests/CrabfleetMacTests/FleetModelsTests.swift +++ b/macos/CrabfleetMac/Tests/CrabfleetMacTests/FleetModelsTests.swift @@ -16,6 +16,54 @@ private func nativeVNCGrant(leaseID: String = "cbx_native123") -> NativeVNCGrant } struct FleetModelsTests { + @Test + func crabboxReceivesOnlyItsMinimalSubprocessEnvironment() { + let environment = CrabboxVNCBridge.commandEnvironment( + from: [ + "HOME": "/Users/tester", + "PATH": "/tmp/untrusted", + "SSH_AUTH_SOCK": "/tmp/agent.sock", + "HTTPS_PROXY": "http://proxy.example.test:8443", + "NO_PROXY": "localhost,.example.test", + "SSL_CERT_FILE": "/etc/ssl/custom-ca.pem", + "SSL_CERT_DIR": "/etc/ssl/custom-certs", + "CRABBOX_CONFIG": "/Users/tester/.config/crabbox/config.yaml", + "XDG_CONFIG_HOME": "/Users/tester/.config", + "XDG_STATE_HOME": "/Users/tester/.local/state", + "CRABFLEET_SESSION_COOKIE": "secret", + "NODE_TLS_REJECT_UNAUTHORIZED": "0", + ] + ) + + #expect(environment["HOME"] == "/Users/tester") + #expect(environment["PATH"] == SubprocessEnvironment.safePath) + #expect(environment["SSH_AUTH_SOCK"] == "/tmp/agent.sock") + #expect(environment["HTTPS_PROXY"] == "http://proxy.example.test:8443") + #expect(environment["NO_PROXY"] == "localhost,.example.test") + #expect(environment["SSL_CERT_FILE"] == "/etc/ssl/custom-ca.pem") + #expect(environment["SSL_CERT_DIR"] == "/etc/ssl/custom-certs") + #expect(environment["CRABBOX_CONFIG"] == "/Users/tester/.config/crabbox/config.yaml") + #expect(environment["XDG_CONFIG_HOME"] == "/Users/tester/.config") + #expect(environment["XDG_STATE_HOME"] == "/Users/tester/.local/state") + #expect(environment["CRABFLEET_SESSION_COOKIE"] == nil) + #expect(environment["NODE_TLS_REJECT_UNAUTHORIZED"] == nil) + } + + @Test + func crabboxRejectsUnsafeConfigEnvironmentPaths() { + let environment = CrabboxVNCBridge.commandEnvironment( + from: [ + "CRABBOX_CONFIG": "relative/config.yaml", + "XDG_CONFIG_HOME": "/Users/tester/../other-config", + "XDG_STATE_HOME": "/" + String(repeating: "a", count: Int(PATH_MAX)), + ] + ) + + #expect(environment["CRABBOX_CONFIG"] == nil) + #expect(environment["XDG_CONFIG_HOME"] == nil) + #expect(environment["XDG_STATE_HOME"] == nil) + } + @Test func sizesRemoteDesktopToEvenViewportPixelsWithinPerformanceCap() { #expect( @@ -159,6 +207,32 @@ struct FleetModelsTests { } } + @Test + func rejectsHTTPSNativeGrantWithoutAHost() async { + let grant = NativeVNCGrant( + brokerURL: URL(string: "https:///native-vnc")!, + leaseID: "cbx_native123", + ticket: nativeVNCTicket, + expiresAt: Date().addingTimeInterval(60) + ) + + await #expect(throws: CrabboxVNCBridgeError.invalidHandoff) { + _ = try await CrabboxVNCBridge.start(grant: grant) + } + } + + @Test + func acceptsCaseInsensitiveNativeGrantSchemes() { + let grant = NativeVNCGrant( + brokerURL: URL(string: "HTTPS://crabbox.example.test/native-vnc")!, + leaseID: "cbx_native123", + ticket: nativeVNCTicket, + expiresAt: Date().addingTimeInterval(60) + ) + + #expect(CrabboxVNCBridge.validGrant(grant)) + } + @Test func parsesGenericVNCAddresses() throws { let direct = try VNCAddress.parse("workstation.example:5907") diff --git a/macos/CrabfleetMac/Tests/CrabfleetMacTests/HostShareProtocolTests.swift b/macos/CrabfleetMac/Tests/CrabfleetMacTests/HostShareProtocolTests.swift index 943d8a77..f5eee6cd 100644 --- a/macos/CrabfleetMac/Tests/CrabfleetMacTests/HostShareProtocolTests.swift +++ b/macos/CrabfleetMac/Tests/CrabfleetMacTests/HostShareProtocolTests.swift @@ -1,5 +1,6 @@ import AppKit import Foundation +import RoyalVNCKit import Testing @testable import CrabfleetMac @@ -15,20 +16,21 @@ struct HostShareWireTests { ) #expect( - update == Data([ - 0, 0, // FramebufferUpdate + padding - 0, 1, // one rectangle - 0, 1, // x = reason (client-requested) - 0, 0, // y = status (no error) - 0x05, 0x00, // width 1280 - 0x02, 0xD0, // height 720 - 0xFF, 0xFF, 0xFE, 0xCC, // ExtendedDesktopSize (-308) - 1, 0, 0, 0, // one screen + padding - 0, 0, 0, 1, // screen id - 0, 0, 0, 0, // position - 0x05, 0x00, 0x02, 0xD0, // screen size - 0, 0, 0, 0, // flags - ]) + update + == Data([ + 0, 0, // FramebufferUpdate + padding + 0, 1, // one rectangle + 0, 1, // x = reason (client-requested) + 0, 0, // y = status (no error) + 0x05, 0x00, // width 1280 + 0x02, 0xD0, // height 720 + 0xFF, 0xFF, 0xFE, 0xCC, // ExtendedDesktopSize (-308) + 1, 0, 0, 0, // one screen + padding + 0, 0, 0, 1, // screen id + 0, 0, 0, 0, // position + 0x05, 0x00, 0x02, 0xD0, // screen size + 0, 0, 0, 0, // flags + ]) ) } @@ -52,6 +54,21 @@ struct HostShareWireTests { #expect(RFBWire.legacyServerCutText(text: "emoji 🦀") == nil) } + @Test + func emptyExtendedClipboardTextRemainsRequestable() throws { + let caps = VNCExtendedClipboardCaps( + supportsText: true, + maximumUnsolicitedTextBytes: 0, + actions: VNCExtendedClipboard.notifyAction + ) + let packet = try #require( + RFBWire.hostClipboardPayload(text: "", extendedNegotiated: true, caps: caps) + ) + + #expect(packet[0] == 3) + #expect(try VNCExtendedClipboard.decode(body: packet.subdata(in: 8.. Bool diff --git a/macos/CrabfleetMac/Tests/CrabfleetMacTests/NativeConnectionTests.swift b/macos/CrabfleetMac/Tests/CrabfleetMacTests/NativeConnectionTests.swift index b6d9ce1d..6d2bf33e 100644 --- a/macos/CrabfleetMac/Tests/CrabfleetMacTests/NativeConnectionTests.swift +++ b/macos/CrabfleetMac/Tests/CrabfleetMacTests/NativeConnectionTests.swift @@ -123,6 +123,26 @@ struct NativeConnectionTests { #expect(grant.leaseID == "cbx_native123") #expect(grant.ticket == "native_vnc_0123456789abcdef0123456789abcdef") + let missingHostTransport = RecordingHTTPTransport { request in + let body = Data( + """ + { + "grant": { + "brokerUrl": "https:///native-vnc", + "leaseId": "cbx_native123", + "ticket": "native_vnc_0123456789abcdef0123456789abcdef", + "expiresAt": "\(expiryFormatter.string(from: expiresAt))" + } + } + """.utf8 + ) + return (body, httpResponse(url: request.url!, status: 200)) + } + await #expect(throws: NativeAPIError.invalidResponse) { + try await NativeAPIClient(origin: origin, transport: missingHostTransport) + .nativeVNCGrant(sessionID: "IS-257", accessToken: "access-token") + } + await #expect(throws: NativeAPIError.invalidResponse) { try await NativeAPIClient(origin: origin, transport: transport) .nativeVNCGrant(sessionID: "IS-0", accessToken: "access-token") @@ -693,6 +713,44 @@ struct NativeConnectionTests { #expect(api.fleetTokens == ["saved-token"]) } + @Test + func disconnectDiscardsAnInFlightNativeVNCGrant() async throws { + let origin = try DeploymentOrigin("https://fleet.example.test") + let origins = MemoryOriginStore(value: origin.displayValue) + let tokens = MemoryTokenStore(values: [origin.displayValue: "saved-token"]) + let api = StubNativeAPIClient(origin: origin) + api.sessionResult = .success(testSession()) + api.fleetResult = .success(testFleet()) + var grantStarted = false + var grantContinuation: CheckedContinuation? + api.nativeVNCGrantHandler = { _, _ in + grantStarted = true + return try await withCheckedThrowingContinuation { continuation in + grantContinuation = continuation + } + } + let store = FleetStore( + environment: [:], + originStore: origins, + tokenStore: tokens, + clientFactory: { _ in api }, + openURL: { _ in false } + ) + await store.restore() + + let grantTask = Task { + try await store.nativeVNCGrant(sessionID: "IS-257") + } + try await waitUntil { grantStarted } + store.disconnect() + let continuation = try #require(grantContinuation) + continuation.resume(returning: testNativeVNCGrant()) + + await #expect(throws: CancellationError.self) { + try await grantTask.value + } + } + @Test func unauthorizedOAuthReadRotatesAndPersistsCredentialBeforeRetry() async throws { let origin = try DeploymentOrigin("https://fleet.example.test") @@ -1725,6 +1783,7 @@ private final class StubNativeAPIClient: NativeAPIClientProtocol { var fleetResult: Result = .failure(NativeAPIError.invalidResponse) var nativeVNCGrantResult: Result = .failure( NativeAPIError.invalidResponse) + var nativeVNCGrantHandler: ((String, String) async throws -> NativeVNCGrant)? var fleetHandler: (() async throws -> NativeAPIFleet)? var sessionTokens: [String] = [] var fleetTokens: [String] = [] @@ -1773,7 +1832,10 @@ private final class StubNativeAPIClient: NativeAPIClientProtocol { } func nativeVNCGrant(sessionID: String, accessToken: String) async throws -> NativeVNCGrant { - try nativeVNCGrantResult.get() + if let nativeVNCGrantHandler { + return try await nativeVNCGrantHandler(sessionID, accessToken) + } + return try nativeVNCGrantResult.get() } func refreshCredential(accessToken: String) async throws -> String? { @@ -1890,6 +1952,15 @@ private func testSession() -> NativeAPISession { ) } +private func testNativeVNCGrant() -> NativeVNCGrant { + .init( + brokerURL: URL(string: "https://crabbox.example.test")!, + leaseID: "cbx_native123", + ticket: "native_vnc_0123456789abcdef0123456789abcdef", + expiresAt: Date().addingTimeInterval(60) + ) +} + private func testLease(id: String = "IS-live") -> CrabboxLease { .init( id: id, diff --git a/macos/CrabfleetMac/Tests/CrabfleetMacTests/PrivateMacShareTests.swift b/macos/CrabfleetMac/Tests/CrabfleetMacTests/PrivateMacShareTests.swift index f33a57b5..9061351d 100644 --- a/macos/CrabfleetMac/Tests/CrabfleetMacTests/PrivateMacShareTests.swift +++ b/macos/CrabfleetMac/Tests/CrabfleetMacTests/PrivateMacShareTests.swift @@ -1,5 +1,6 @@ import AppKit import Foundation +import Network import Testing @testable import CrabfleetMac @@ -22,12 +23,16 @@ struct PrivateMacShareTests { let environment = SystemTailscaleCommandRunner.commandEnvironment( from: [ - "PATH": "/usr/bin:/bin", + "HOME": "/Users/tester", + "PATH": "/tmp/untrusted", + "SECRET_TOKEN": "test-token-placeholder", "TS_DEBUG": "unsafe", "TAILSCALE_SOCKET": "/tmp/unsafe.sock", ] ) - #expect(environment["PATH"] == "/usr/bin:/bin") + #expect(environment["HOME"] == "/Users/tester") + #expect(environment["PATH"] == SubprocessEnvironment.safePath) + #expect(environment["SECRET_TOKEN"] == nil) #expect(environment["TS_DEBUG"] == nil) #expect(environment["TAILSCALE_SOCKET"] == nil) #expect(environment["TAILSCALE_BE_CLI"] == "1") @@ -56,476 +61,2198 @@ struct PrivateMacShareTests { } @Test - func privateShareCanStartViewOnlyWithoutAccessibility() { - #expect( - PrivateMacSharePermissionPolicy.canStart( - identityAvailable: true, - screenRecordingGranted: true - )) - #expect( - !PrivateMacSharePermissionPolicy.canStart( - identityAvailable: false, - screenRecordingGranted: true - )) - #expect( - !PrivateMacSharePermissionPolicy.canStart( - identityAvailable: true, - screenRecordingGranted: false - )) + func tailscaleCommandTimesOutAndRespondsToCancellation() async throws { + let directory = FileManager.default.temporaryDirectory + .appendingPathComponent("CrabfleetMacTests.\(UUID().uuidString)") + try FileManager.default.createDirectory(at: directory, withIntermediateDirectories: true) + defer { try? FileManager.default.removeItem(at: directory) } + let executable = directory.appendingPathComponent("tailscale") + let pidFile = directory.appendingPathComponent("pid") + try Data( + """ + #!/bin/sh + printf '%s' "$$" > '\(pidFile.path)' + exec sleep 30 + """.utf8 + ).write(to: executable) + try FileManager.default.setAttributes([.posixPermissions: 0o700], ofItemAtPath: executable.path) + + let timedRunner = SystemTailscaleCommandRunner(executableURL: executable, timeout: 0.1) + await #expect(throws: PrivateMacShareError.commandTimedOut) { + _ = try await timedRunner.run(arguments: ["status"]) + } + + try? FileManager.default.removeItem(at: pidFile) + let cancellableRunner = SystemTailscaleCommandRunner(executableURL: executable, timeout: 30) + let task = Task { + try await cancellableRunner.run(arguments: ["status"]) + } + let launched = await waitUntilAsync { + FileManager.default.fileExists(atPath: pidFile.path) + } + #expect(launched) + let cancelledPID = try #require(Int(String(contentsOf: pidFile, encoding: .utf8))) + task.cancel() + await #expect(throws: CancellationError.self) { + try await task.value + } + #expect(await waitUntilAsync { Darwin.kill(Int32(cancelledPID), 0) != 0 }) } @Test - func recognizesExplicitPrivateShareLaunchMode() { - #expect( - PrivateMacShareLaunchMode.isEnabled( - arguments: ["CrabfleetMac", "--share-this-mac"], environment: [:])) - #expect( - PrivateMacShareLaunchMode.isEnabled( - arguments: ["CrabfleetMac"], environment: ["CRABFLEET_AUTO_SHARE": "1"])) - #expect( - !PrivateMacShareLaunchMode.isEnabled( - arguments: ["CrabfleetMac"], environment: ["CRABFLEET_AUTO_SHARE": "true"])) - #expect(!PrivateMacShareLaunchMode.isEnabled(arguments: ["CrabfleetMac"], environment: [:])) + func tailscaleCommandTimeoutTerminatesDescendantProcessGroup() async throws { + let directory = FileManager.default.temporaryDirectory + .appendingPathComponent("CrabfleetMacTests.\(UUID().uuidString)") + try FileManager.default.createDirectory(at: directory, withIntermediateDirectories: true) + defer { try? FileManager.default.removeItem(at: directory) } + let executable = directory.appendingPathComponent("tailscale") + let descendantPIDFile = directory.appendingPathComponent("descendant-pid") + try Data( + """ + #!/bin/sh + ( + trap '' HUP TERM + exec sleep 30 + ) & + printf '%s' "$!" > '\(descendantPIDFile.path)' + exec sleep 30 + """.utf8 + ).write(to: executable) + try FileManager.default.setAttributes([.posixPermissions: 0o700], ofItemAtPath: executable.path) + + // Give the helper time to spawn under a concurrently loaded Swift test runner. + let runner = SystemTailscaleCommandRunner(executableURL: executable, timeout: 2) + let clock = ContinuousClock() + let startedAt = clock.now + await #expect(throws: PrivateMacShareError.commandTimedOut) { + _ = try await runner.run(arguments: ["status"]) + } + let elapsed = startedAt.duration(to: clock.now) + + let descendantPID = try #require( + Int32(String(contentsOf: descendantPIDFile, encoding: .utf8)) + ) + #expect(await waitUntilAsync { Darwin.kill(descendantPID, 0) != 0 }) + #expect(elapsed < .seconds(4)) } @Test - func parsesExplicitVNCConnectionLaunchMode() throws { - let explicitAddress = try VNCConnectionLaunchMode.address( - arguments: ["CrabfleetMac", "--connect", "vnc://100.64.0.8:5901"], - environment: ["CRABFLEET_AUTO_CONNECT": "vnc://ignored.example:5900"] - ) - let explicit = try #require(explicitAddress) - #expect(explicit.host == "100.64.0.8") - #expect(explicit.port == 5_901) + func successfulTailscaleCommandDoesNotWaitForDescendantPipeEOF() async throws { + let directory = FileManager.default.temporaryDirectory + .appendingPathComponent("CrabfleetMacTests.\(UUID().uuidString)") + try FileManager.default.createDirectory(at: directory, withIntermediateDirectories: true) + defer { try? FileManager.default.removeItem(at: directory) } + let executable = directory.appendingPathComponent("tailscale") + let descendantPIDFile = directory.appendingPathComponent("descendant-pid") + try Data( + """ + #!/bin/sh + ( + trap '' HUP TERM + exec sleep 30 + ) & + printf '%s' "$!" > '\(descendantPIDFile.path)' + printf 'status complete' + exit 0 + """.utf8 + ).write(to: executable) + try FileManager.default.setAttributes([.posixPermissions: 0o700], ofItemAtPath: executable.path) - let environmentAddress = try VNCConnectionLaunchMode.address( - arguments: ["CrabfleetMac"], - environment: ["CRABFLEET_AUTO_CONNECT": "viewer.example:5999"] + let runner = SystemTailscaleCommandRunner(executableURL: executable, timeout: 5) + let clock = ContinuousClock() + let startedAt = clock.now + let result = try await runner.run(arguments: ["status"]) + let elapsed = startedAt.duration(to: clock.now) + + let descendantPID = try #require( + Int32(String(contentsOf: descendantPIDFile, encoding: .utf8)) ) - let environment = try #require(environmentAddress) - #expect(environment.host == "viewer.example") - #expect(environment.port == 5_999) - #expect( - try VNCConnectionLaunchMode.address( - arguments: ["CrabfleetMac"], environment: [:]) == nil) + defer { + _ = Darwin.kill(descendantPID, SIGKILL) + } + #expect(result.standardOutput == "status complete") + #expect(Darwin.kill(descendantPID, 0) == 0) + #expect(elapsed < .seconds(2)) } @Test - func rejectsMissingOrCredentialedVNCConnectionLaunchAddress() { - #expect(throws: VNCAddressError.missingHost) { - try VNCConnectionLaunchMode.address( - arguments: ["CrabfleetMac", "--connect"], environment: [:]) + func tailscaleCommandCancellationTerminatesDescendantProcessGroup() async throws { + let directory = FileManager.default.temporaryDirectory + .appendingPathComponent("CrabfleetMacTests.\(UUID().uuidString)") + try FileManager.default.createDirectory(at: directory, withIntermediateDirectories: true) + defer { try? FileManager.default.removeItem(at: directory) } + let executable = directory.appendingPathComponent("tailscale") + let descendantPIDFile = directory.appendingPathComponent("descendant-pid") + try Data( + """ + #!/bin/sh + ( + trap '' HUP TERM + exec sleep 30 + ) & + printf '%s' "$!" > '\(descendantPIDFile.path)' + exec sleep 30 + """.utf8 + ).write(to: executable) + try FileManager.default.setAttributes([.posixPermissions: 0o700], ofItemAtPath: executable.path) + + let runner = SystemTailscaleCommandRunner(executableURL: executable, timeout: 30) + let task = Task { + try await runner.run(arguments: ["status"]) } - #expect(throws: VNCAddressError.embeddedPassword) { - try VNCConnectionLaunchMode.address( - arguments: ["CrabfleetMac", "--connect", "vnc://user:secret@example.test"], - environment: [:] - ) + #expect(await waitUntilAsync { + FileManager.default.fileExists(atPath: descendantPIDFile.path) + }) + let descendantPID = try #require( + Int32(String(contentsOf: descendantPIDFile, encoding: .utf8)) + ) + task.cancel() + await #expect(throws: CancellationError.self) { + try await task.value } + #expect(await waitUntilAsync { Darwin.kill(descendantPID, 0) != 0 }) } - @Test - func acceptsOnlineUserOnActiveTailnet() throws { - let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + @Test @MainActor + func stopInvalidatesAnInFlightPrivateShareStart() async throws { + let runner = SuspendedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: nil, + defaults: defaults + ) + let startTask = Task { await controller.start() } + let started = await waitUntilAsync { await runner.hasStarted } + #expect(started) - #expect(identity.tailnetName == "example.com") - #expect(identity.loginName == "operator@example.com") - #expect(identity.ipv4Address == "100.64.12.34") - #expect(identity.vncAddress(port: 5901) == "vnc://100.64.12.34:5901") + await controller.stop() + await runner.resume( + .success(.init(standardOutput: statusJSON(), standardError: "")) + ) + await startTask.value + + #expect(controller.phase == .idle) + #expect(controller.identity == nil) } - @Test - func derivesGenericStableDesktopHostID() { - let identity = TailnetIdentity( - tailnetName: "example.com", - loginName: "operator@example.com", - dnsName: "workstation-1.example.ts.net", - hostName: "Workstation", - ipv4Address: "100.64.12.34", - userID: 42 + @Test @MainActor + func startWaitsForAnInFlightRefresh() async throws { + let runner = SequencedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: nil, + defaults: defaults ) - #expect(CrabfleetDesktopRegistration.hostID(identity: identity) == "workstation-1") - let fallback = TailnetIdentity( - tailnetName: identity.tailnetName, - loginName: identity.loginName, - dnsName: "", - hostName: identity.hostName, - ipv4Address: identity.ipv4Address, - userID: identity.userID + let refreshTask = Task { await controller.refresh() } + #expect(await waitUntilAsync { await runner.callCount == 1 }) + #expect(controller.isRefreshing) + + let startState = AsyncInvocationState() + let startTask = Task { + await startState.markStarted() + await controller.start() + await startState.markFinished() + } + #expect(await waitUntilAsync { await startState.started }) + try await Task.sleep(for: .milliseconds(20)) + #expect(!(await startState.finished)) + + await runner.resumeNext( + .success(.init(standardOutput: statusJSON(), standardError: "")) ) - #expect(CrabfleetDesktopRegistration.hostID(identity: fallback) == "mac-100-64-12-34") - } + await refreshTask.value + #expect(await waitUntilAsync { await runner.callCount == 2 }) + await runner.resumeNext(.failure(PrivateMacShareError.tailscaleOffline)) + await startTask.value - @Test - func acceptsOnlySecureCrabfleetAPIURLs() throws { - #expect( - CrabfleetDesktopRegistration.isSecureAPIURL( - try #require(URL(string: "https://fleet.example/api/fleet")))) - #expect( - CrabfleetDesktopRegistration.isSecureAPIURL( - try #require(URL(string: "http://127.0.0.1:8787")))) - #expect( - !CrabfleetDesktopRegistration.isSecureAPIURL( - try #require(URL(string: "http://fleet.example")))) - #expect( - !CrabfleetDesktopRegistration.isSecureAPIURL( - try #require(URL(string: "https://user@fleet.example")))) - #expect( - !CrabfleetDesktopRegistration.isSecureAPIURL( - try #require(URL(string: "https://fleet.example?token=value")))) + #expect(controller.phase == .failed) + #expect(await runner.callCount == 2) } - @Test - func buildsAuthenticatedDesktopRegistrationRequest() throws { - let registration = try #require( - CrabfleetDesktopRegistration(environment: [ - "CRABFLEET_API_URL": "https://fleet.example/api/fleet", - "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", - ])) - let identity = TailnetIdentity( - tailnetName: "example.com", - loginName: "operator@example.com", - dnsName: "workstation-1.example.ts.net", - hostName: "Workstation", - ipv4Address: "100.64.12.34", - userID: 42 + @Test @MainActor + func stopInvalidatesAStartWaitingForRefreshCompletion() async throws { + let runner = SequencedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: nil, + defaults: defaults ) - let request = try registration.registrationRequest(identity: identity, port: 5901) - #expect(request.url?.absoluteString == "https://fleet.example/api/desktop-hosts/workstation-1") - #expect(request.httpMethod == "PUT") - #expect(request.value(forHTTPHeaderField: "Cookie") == "crabbox_session=secret") - let body = try #require(request.httpBody) - let json = try #require(JSONSerialization.jsonObject(with: body) as? [String: Any]) - #expect(json["name"] as? String == "Workstation") - #expect(json["address"] as? String == "100.64.12.34") - #expect(json["port"] as? Int == 5901) - } + let refreshTask = Task { await controller.refresh() } + #expect(await waitUntilAsync { await runner.callCount == 1 }) - @Test - func rejectsInvalidTailnetAndIdentityFields() throws { - var value = statusJSON() - value = value.replacingOccurrences( - of: #""Name": "example.com""#, with: #""Name": """#) - let missingTailnet = try JSONDecoder().decode( - TailscaleStatusDocument.self, - from: Data(value.utf8) + let startTask = Task { await controller.start() } + #expect(await waitUntilAsync { controller.phase == .starting }) + + await controller.stop() + #expect(controller.phase == .idle) + + await runner.resumeNext( + .success(.init(standardOutput: statusJSON(), standardError: "")) ) - #expect(throws: PrivateMacShareError.invalidTailnetIdentity) { - try TailnetIdentityPolicy.identity(from: missingTailnet) - } + await refreshTask.value + await startTask.value - value = statusJSON().replacingOccurrences( - of: "operator@example.com", - with: "" + #expect(await runner.callCount == 1) + #expect(controller.phase == .idle) + } + + @Test @MainActor + func cancellationRestoresIdleWhileStartWaitsForRefresh() async throws { + let runner = SequencedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") ) - let missingUser = try JSONDecoder().decode( - TailscaleStatusDocument.self, - from: Data(value.utf8) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: nil, + defaults: defaults ) - #expect(throws: PrivateMacShareError.invalidTailnetUser) { - try TailnetIdentityPolicy.identity(from: missingUser) - } - } - @Test - func recognizesOnlyTailscaleIPv4Range() { - #expect(TailnetIdentityPolicy.isTailscaleIPv4("100.64.0.1")) - #expect(TailnetIdentityPolicy.isTailscaleIPv4("100.127.255.254")) - #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.63.255.255")) - #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.128.0.1")) - #expect(!TailnetIdentityPolicy.isTailscaleIPv4("10.0.0.1")) - #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.64.invalid.1.2")) - #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.64..1")) - } + let refreshTask = Task { await controller.refresh() } + #expect(await waitUntilAsync { await runner.callCount == 1 }) - @Test - func validatesBoundedTailnetIdentityFields() { - #expect(TailnetIdentityPolicy.isValidTailnetName("example.com")) - #expect(TailnetIdentityPolicy.isValidTailnetName("example.github")) - #expect(!TailnetIdentityPolicy.isValidTailnetName("")) - #expect(!TailnetIdentityPolicy.isValidTailnetName(" example.com")) - #expect(!TailnetIdentityPolicy.isValidTailnetName("bad\nname")) - #expect(!TailnetIdentityPolicy.isValidTailnetName(String(repeating: "a", count: 254))) + let startTask = Task { await controller.start() } + #expect(await waitUntilAsync { controller.phase == .starting }) + startTask.cancel() - #expect(TailnetIdentityPolicy.isValidLogin("operator@example.com")) - #expect(TailnetIdentityPolicy.isValidLogin("github-user")) - #expect(!TailnetIdentityPolicy.isValidLogin("")) - #expect(!TailnetIdentityPolicy.isValidLogin("github-user ")) - #expect(!TailnetIdentityPolicy.isValidLogin("bad\u{0}login")) - #expect(!TailnetIdentityPolicy.isValidLogin(String(repeating: "a", count: 321))) + await runner.resumeNext( + .success(.init(standardOutput: statusJSON(), standardError: "")) + ) + await refreshTask.value + await startTask.value + + #expect(await runner.callCount == 1) + #expect(controller.phase == .idle) } @Test - func authorizesOnlySameTailnetUserAndExactPeerAddress() async throws { + func desktopRemovalWaitsForACommittedRegistrationAfterCancellation() async throws { + let registration = SuspendedDesktopRegistration() + let coordinator = DesktopHostRegistrationCoordinator(registration: registration) let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) - let accepted = TailnetPeerAuthorizer( - runner: StaticTailscaleRunner(output: whoisJSON(login: identity.loginName)), - expectedIdentity: identity - ) - #expect(await accepted.authorize(remoteAddress: "100.100.10.20")) - let otherUser = TailnetPeerAuthorizer( - runner: StaticTailscaleRunner(output: whoisJSON(login: "other@example.com")), - expectedIdentity: identity - ) - let otherUserID = TailnetPeerAuthorizer( - runner: StaticTailscaleRunner( - output: whoisJSON(login: identity.loginName, userID: 43)), - expectedIdentity: identity + let publish = Task { + _ = try await coordinator.register( + identity: identity, + port: 5_901, + publicationID: "publication-id" + ) + } + #expect(await waitUntilAsync { await registration.hasStartedRegistration }) + publish.cancel() + + let remove = Task { + try await coordinator.unregister(identity: identity, ownershipToken: "registration-token") + } + try await Task.sleep(for: .milliseconds(20)) + #expect(await registration.events == [.registerStarted]) + + await registration.finishRegistration() + try await publish.value + try await remove.value + + #expect( + await registration.events + == [.registerStarted, .registerFinished, .unregisterStarted] ) - let otherAddress = TailnetPeerAuthorizer( - runner: StaticTailscaleRunner( - output: whoisJSON(login: identity.loginName, addresses: ["100.100.10.21/32"])), - expectedIdentity: identity + } + + @Test @MainActor + func stopReturnsBeforeSlowDesktopRegistryCleanup() async throws { + let registration = SuspendedDesktopCleanupRegistration() + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + let identity = desktopIdentity(name: "slow-cleanup", address: "100.64.12.43") + try await lifecycle.publish(identity: identity, port: 5_901) + + let runner = SuspendedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") ) - let unauthorizedNode = TailnetPeerAuthorizer( - runner: StaticTailscaleRunner( - output: whoisJSON(login: identity.loginName, machineAuthorized: false)), - expectedIdentity: identity + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: registration, + registrationLifecycle: lifecycle, + defaults: defaults ) - #expect(!(await otherUser.authorize(remoteAddress: "100.100.10.20"))) - #expect(!(await otherUserID.authorize(remoteAddress: "100.100.10.20"))) - #expect(!(await otherAddress.authorize(remoteAddress: "100.100.10.20"))) - #expect(!(await unauthorizedNode.authorize(remoteAddress: "100.100.10.20"))) - #expect(!(await accepted.authorize(remoteAddress: "192.168.1.4"))) - } + let startTask = Task { await controller.start() } + #expect(await waitUntilAsync { await runner.hasStarted }) - @Test - func keepsNewestCapturedFrameWhenUpdatesArriveOutOfOrder() async throws { - let store = CapturedDesktopFrameStore() - await store.update(.init(jpegData: Data([2]), sequence: 2, width: 2, height: 2)) - await store.update(.init(jpegData: Data([1]), sequence: 1, width: 2, height: 2)) + let stopState = AsyncInvocationState() + let stopTask = Task { + await controller.stop() + await stopState.markFinished() + } + #expect(await waitUntilAsync { await stopState.finished }) + #expect(controller.phase == .idle) + #expect(await waitUntilAsync { await registration.hasStartedUnregistration }) - #expect(await store.latest()?.sequence == 2) + await runner.resume( + .success(.init(standardOutput: statusJSON(), standardError: "")) + ) + await startTask.value + await registration.finishUnregistration() + await stopTask.value + #expect(await waitUntilAsync { controller.registryPhase == .notPublished }) } - @Test - func buildsTightJPEGFramebufferUpdate() throws { - let jpeg = Data([0xFF, 0xD8, 0xFF, 0xD9]) - let frame = CapturedDesktopFrame(jpegData: jpeg, sequence: 7, width: 1_600, height: 900) - let packet = try RFBWire.tightJPEGUpdate(frame: frame) + @Test @MainActor + func concurrentStopsAwaitTheSameInFlightOperation() async throws { + let coordinator = PrivateMacShareStopCoordinator() + let operation = SuspendedAsyncOperation() + let firstState = AsyncInvocationState() + let secondState = AsyncInvocationState() - #expect(packet[0] == 0) - #expect(packet.readUInt16(at: 2) == 1) - #expect(packet.readUInt16(at: 8) == 1_600) - #expect(packet.readUInt16(at: 10) == 900) - #expect(packet.readInt32(at: 12) == RFBWire.tightEncoding) - #expect(packet[16] == 0x90) - #expect(packet[17] == 4) - #expect(packet.suffix(4) == jpeg) + let first = Task { + await coordinator.perform { + await operation.run() + } + await firstState.markFinished() + } + #expect(await waitUntilAsync { await operation.invocationCount == 1 }) + + let second = Task { + await coordinator.perform { + await operation.run() + } + await secondState.markFinished() + } + try await Task.sleep(for: .milliseconds(20)) + + #expect(await operation.invocationCount == 1) + #expect(!(await firstState.finished)) + #expect(!(await secondState.finished)) + + await operation.finish() + await first.value + await second.value + + #expect(await firstState.finished) + #expect(await secondState.finished) } - @Test - func encodesTightCompactLengths() { - #expect(RFBWire.tightCompactLength(0) == Data([0x00])) - #expect(RFBWire.tightCompactLength(127) == Data([0x7F])) - #expect(RFBWire.tightCompactLength(128) == Data([0x80, 0x01])) - #expect(RFBWire.tightCompactLength(16_383) == Data([0xFF, 0x7F])) - #expect(RFBWire.tightCompactLength(16_384) == Data([0x80, 0x80, 0x01])) + @Test @MainActor + func completedStopDoesNotCoalesceWithTheNextOperation() async { + let coordinator = PrivateMacShareStopCoordinator() + let operation = SuspendedAsyncOperation() + + let first = Task { + await coordinator.perform { + await operation.run() + } + } + #expect(await waitUntilAsync { await operation.invocationCount == 1 }) + await operation.finish() + await first.value + + let second = Task { + await coordinator.perform { + await operation.run() + } + } + #expect(await waitUntilAsync { await operation.invocationCount == 2 }) + await operation.finish() + await second.value } - @Test - func scalesCaptureWithinBoundedEvenDimensions() { - let retina = MacScreenCapture.captureDimensions(sourceWidth: 5_120, sourceHeight: 2_880) - #expect(retina.width == 2_560) - #expect(retina.height == 1_440) - #expect(retina.width.isMultiple(of: 2)) - #expect(retina.height.isMultiple(of: 2)) + @Test @MainActor + func failedDesktopPublicationIsNotUnregistered() async throws { + let identity = desktopIdentity(name: "failed-publish", address: "100.64.12.40") + let registration = RecordingDesktopRegistration(registerFailures: [identity.dnsName: 1]) + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) - let small = MacScreenCapture.captureDimensions(sourceWidth: 1_280, sourceHeight: 800) - #expect(small.width == 1_280) - #expect(small.height == 800) + await #expect(throws: DesktopRegistrationTestError.failed) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + try await lifecycle.removePublishedIdentities() + + #expect(await registration.events == [.register(identity.dnsName)]) } - @Test - func mapsRFBKeysymsToMacKeys() { - #expect(MacRemoteInputController.keyCode(for: 0x61) != nil) - #expect(MacRemoteInputController.keyCode(for: 0xFF51) != nil) - #expect(MacRemoteInputController.keyCode(for: 0xFFE7) != nil) - #expect(MacRemoteInputController.keyCode(for: 0x1F980) == nil) + @Test @MainActor + func ambiguousDesktopPublicationIsRecoveredBeforeCleanup() async throws { + let identity = desktopIdentity(name: "ambiguous-publish", address: "100.64.12.46") + let registration = AmbiguousDesktopRegistration() + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + try await lifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register(identity.dnsName), + .recover(identity.dnsName), + .unregister(identity.dnsName, "recovered-token"), + ] + ) } @Test @MainActor - func servesRoyalVNCKitOverTheCurrentTailnet() async throws { - guard ProcessInfo.processInfo.environment["CRABFLEET_TAILNET_RFB_SMOKE"] == "1" else { - return + func negativeRecoveryDoesNotDeleteANewerTokenlessPublisher() async throws { + let identity = desktopIdentity(name: "negative-recovery", address: "100.64.12.62") + let registration = NegativeRecoveryDesktopRegistration() + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) } + await registration.publishNewerEndpoint() + try await lifecycle.removePublishedIdentities() - let runner = try SystemTailscaleCommandRunner() - let status = try await runner.run(arguments: ["status", "--json"]) - let document = try JSONDecoder().decode( - TailscaleStatusDocument.self, - from: Data(status.standardOutput.utf8) + #expect(await registration.activeEndpoint == "newer-publisher") + #expect(await registration.events == [.register, .recover]) + } + + @Test @MainActor + func ambiguousDesktopPublicationRetryRecoversOnlyTheExactIdentity() async throws { + let identity = desktopIdentity(name: "retry-publish", address: "100.64.12.50") + let registration = IdentityAwareAmbiguousDesktopRegistration( + uncertainPublicationIDs: ["publication-a"] ) - let identity = try TailnetIdentityPolicy.identity(from: document) - let capture = MacScreenCapture() - let jpeg = try #require(testJPEG()) - await capture.frameStore.update( - .init(jpegData: jpeg, sequence: 1, width: 64, height: 64) + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "publication-a" } ) - let port: UInt16 = 5_909 - let server = TailnetRFBServer( - identity: identity, - runner: runner, - capture: capture, - descriptor: .init( - displayID: 0, - displayBounds: CGRect(x: 0, y: 0, width: 64, height: 64), - frameWidth: 64, - frameHeight: 64, - sourcePixelWidth: 64, - sourcePixelHeight: 64 - ), - input: NoopRemoteInput(), - port: port, - eventHandler: { _ in } + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + try await lifecycle.publish(identity: identity, port: 5_901) + try await lifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register(identity.ipv4Address, 5_901, "publication-a"), + .recover(identity.ipv4Address, "publication-a"), + .unregister(identity.ipv4Address, "recovered:publication-a"), + ] ) - try server.start() - defer { server.stop() } - try await Task.sleep(for: .milliseconds(250)) + } - let session = VNCSessionController() - session.connect( - host: identity.ipv4Address, - port: port, - username: "", - password: "", - clipboardEnabled: false + @Test @MainActor + func addressRotationRepublishesInsteadOfRecoveringAnUncertainDesktop() async throws { + let first = desktopIdentity(name: "rotating-host", address: "100.64.12.51") + let second = TailnetIdentity( + tailnetName: first.tailnetName, + loginName: first.loginName, + dnsName: first.dnsName, + hostName: first.hostName, + ipv4Address: "100.64.12.52", + userID: first.userID + ) + #expect(first != second) + #expect( + CrabfleetDesktopRegistration.hostID(identity: first) + == CrabfleetDesktopRegistration.hostID(identity: second) + ) + let registration = IdentityAwareAmbiguousDesktopRegistration( + uncertainPublicationIDs: ["publication-a"] + ) + var publicationIDs = ["publication-a", "publication-b"] + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { publicationIDs.removeFirst() } ) - defer { session.disconnect() } - let clock = ContinuousClock() - let deadline = clock.now.advanced(by: .seconds(15)) - while clock.now < deadline { - if session.phase == .connected && session.framebufferUpdateCount > 0 { break } - try await Task.sleep(for: .milliseconds(25)) + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: first, port: 5_901) } - #expect(session.phase == .connected) - #expect(session.framebufferUpdateCount > 0) - #expect(session.framebuffer?.size.width == 64) - #expect(session.framebuffer?.size.height == 64) + try await lifecycle.publish(identity: second, port: 5_901) + try await lifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register(first.ipv4Address, 5_901, "publication-a"), + .register(second.ipv4Address, 5_901, "publication-b"), + .recover(first.ipv4Address, "publication-a"), + .unregister(first.ipv4Address, "recovered:publication-a"), + .unregister(second.ipv4Address, "token:publication-b"), + ] + ) } @Test @MainActor - func syncsUTF8ClipboardAndNegotiatesResizeOverLoopback() async throws { - // Full-protocol end-to-end: the production server and the RoyalVNCKit - // client exchange handshake, Tight frames, Extended Clipboard, and - // ExtendedDesktopSize over a real TCP connection on loopback. The - // tailnet-specific pieces (address binding, whois) are injected. - let identity = TailnetIdentity( - tailnetName: "example.com", - loginName: "tester@example.com", - dnsName: "workstation.example.ts.net.", - hostName: "Workstation", - ipv4Address: "127.0.0.1", - userID: 42 + func portRotationRepublishesInsteadOfRecoveringAnUncertainDesktop() async throws { + let identity = desktopIdentity(name: "rotating-port", address: "100.64.12.53") + let registration = IdentityAwareAmbiguousDesktopRegistration( + uncertainPublicationIDs: ["publication-a"] ) - let capture = MacScreenCapture() - let jpeg = try #require(testJPEG()) - await capture.frameStore.update( - .init(jpegData: jpeg, sequence: 1, width: 64, height: 64) + var publicationIDs = ["publication-a", "publication-b"] + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { publicationIDs.removeFirst() } ) - let hostPasteboard = NSPasteboard(name: .init("CrabfleetMacTests.host.\(UUID().uuidString)")) - hostPasteboard.clearContents() - let hostClipboard = HostClipboardBridge(pasteboard: hostPasteboard, pollingInterval: 0.02) + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + try await lifecycle.publish(identity: identity, port: 5_902) + try await lifecycle.removePublishedIdentities() - let port: UInt16 = 5_921 - let server = TailnetRFBServer( - identity: identity, - runner: StaticTailscaleRunner(output: ""), - capture: capture, - descriptor: .init( - displayID: 0, - displayBounds: CGRect(x: 0, y: 0, width: 64, height: 64), - frameWidth: 64, - frameHeight: 64, - sourcePixelWidth: 256, - sourcePixelHeight: 256 - ), - input: NoopRemoteInput(), - clipboard: hostClipboard, - peerAuthorizer: LoopbackPeerAuthorizer(), - port: port, - eventHandler: { _ in } + #expect( + await registration.events + == [ + .register(identity.ipv4Address, 5_901, "publication-a"), + .register(identity.ipv4Address, 5_902, "publication-b"), + .recover(identity.ipv4Address, "publication-a"), + .unregister(identity.ipv4Address, "recovered:publication-a"), + .unregister(identity.ipv4Address, "token:publication-b"), + ] ) - try server.start() - defer { server.stop() } - try await Task.sleep(for: .milliseconds(250)) + } - let viewerPasteboard = NSPasteboard( - name: .init("CrabfleetMacTests.viewer.\(UUID().uuidString)") + @Test @MainActor + func ambiguousDesktopCleanupPreservesANewerPublisher() async throws { + let identity = desktopIdentity(name: "shared-host", address: "100.64.12.47") + let registration = TwoProcessDesktopRegistration(lostPublicationID: "publication-a") + let firstLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "publication-a" } ) - viewerPasteboard.clearContents() - let coordinator = ClipboardCoordinator(pasteboard: viewerPasteboard, pollingInterval: 0.02) - let session = VNCSessionController(targetID: "smoke", clipboardCoordinator: coordinator) - coordinator.focus(session: session, targetID: "smoke") - session.connect( - host: identity.ipv4Address, - port: port, - username: "", - password: "" + let secondLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "publication-b" } ) - defer { session.disconnect() } - // The Extended Clipboard caps handshake must complete on the client. - try await waitFor { session.connection?.supportsUTF8Clipboard == true } + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await firstLifecycle.publish(identity: identity, port: 5_901) + } + try await secondLifecycle.publish(identity: identity, port: 5_901) + try await firstLifecycle.removePublishedIdentities() - // Viewer to host: emoji only survives the extended UTF-8 path. - viewerPasteboard.clearContents() - viewerPasteboard.setString("client copy 🚀", forType: .string) - try await waitFor { hostPasteboard.string(forType: .string) == "client copy 🚀" } + #expect(await registration.activePublicationID == "publication-b") + #expect( + await registration.events + == [ + .register("publication-a"), + .register("publication-b"), + .recover("publication-a"), + ] + ) - // Host to server-cut-text: the host push lands on the viewer pasteboard. - hostPasteboard.clearContents() - hostPasteboard.setString("host copy 🦀", forType: .string) - try await waitFor { viewerPasteboard.string(forType: .string) == "host copy 🦀" } + try await secondLifecycle.removePublishedIdentities() + #expect(await registration.activePublicationID == nil) + } - // The ExtendedDesktopSize announce must unlock client resize requests. - try await waitFor { session.requestDesktopSize(.init(width: 128, height: 128)) } + @Test @MainActor + func retainedLegacyCleanupDoesNotDeleteANewerPublisher() async throws { + let identity = desktopIdentity(name: "legacy-host", address: "100.64.12.60") + let registration = RetainedLegacyDesktopRegistration() + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) - // The stream must keep flowing after the resize exchange (this test - // fixture has no live capture stream, so the server answers the resize - // with an out-of-resources status and continues serving frames). An - // unchanged frame is deduplicated into rectangle-free heartbeats, so new - // content must arrive as a fresh framebuffer update. - let updates = session.framebufferUpdateCount - await capture.frameStore.update( - .init(jpegData: jpeg, sequence: 2, width: 64, height: 64) + try await lifecycle.publish(identity: identity, port: 5_901) + await #expect(throws: DesktopRegistrationTestError.failed) { + try await lifecycle.removePublishedIdentities() + } + await registration.publishNewerEndpoint() + + try await lifecycle.removePublishedIdentities() + + #expect(await registration.activeEndpoint == "newer-publisher") + #expect(await registration.events == [.register, .unregister]) + } + + @Test @MainActor + func desktopPublicationReplacementUsesSanitizedHostIDAcrossIdentityChanges() async throws { + let first = desktopIdentity(name: "shared-host", address: "100.64.12.48") + let second = TailnetIdentity( + tailnetName: first.tailnetName, + loginName: first.loginName, + dnsName: first.dnsName, + hostName: "Shared Host Renamed", + ipv4Address: "100.64.12.49", + userID: first.userID + ) + #expect(first != second) + #expect( + CrabfleetDesktopRegistration.hostID(identity: first) + == CrabfleetDesktopRegistration.hostID(identity: second) + ) + let registration = MutableIdentityDesktopRegistration() + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + + try await lifecycle.publish(identity: first, port: 5_901) + try await lifecycle.publish(identity: second, port: 5_901) + try await lifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register(first.ipv4Address), + .register(second.ipv4Address), + .unregister(second.ipv4Address, "token:\(second.ipv4Address)"), + ] ) - try await waitFor { session.framebufferUpdateCount > updates } - #expect(session.phase == .connected) } - @MainActor - private func waitFor( - timeout: Duration = .seconds(15), - _ condition: @escaping @MainActor () -> Bool - ) async throws { - let clock = ContinuousClock() - let deadline = clock.now.advanced(by: timeout) - while clock.now < deadline { - if condition() { return } - try await Task.sleep(for: .milliseconds(25)) + @Test @MainActor + func failedDesktopRemovalSurvivesLaterIdentityChanges() async throws { + let first = desktopIdentity(name: "first-host", address: "100.64.12.41") + let second = desktopIdentity(name: "second-host", address: "100.64.12.42") + let registration = RecordingDesktopRegistration( + unregisterFailures: [first.dnsName: 2] + ) + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + + try await lifecycle.publish(identity: first, port: 5_901) + await #expect(throws: DesktopRegistrationTestError.failed) { + try await lifecycle.removePublishedIdentities() } - #expect(condition()) - } - private func statusDocument() throws -> TailscaleStatusDocument { - try JSONDecoder().decode(TailscaleStatusDocument.self, from: Data(statusJSON().utf8)) + try await lifecycle.publish(identity: second, port: 5_901) + await #expect(throws: DesktopRegistrationTestError.failed) { + try await lifecycle.removePublishedIdentities() + } + try await lifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register(first.dnsName), + .unregister(first.dnsName, "token:\(first.dnsName)"), + .register(second.dnsName), + .unregister(first.dnsName, "token:\(first.dnsName)"), + .unregister(second.dnsName, "token:\(second.dnsName)"), + .unregister(first.dnsName, "token:\(first.dnsName)"), + ] + ) } - private func statusJSON() -> String { - """ - { - "BackendState": "Running", - "CurrentTailnet": { "Name": "example.com" }, - "Self": { - "DNSName": "workstation.example.ts.net.", - "HostName": "Workstation", + @Test @MainActor + func terminationRetriesRetainedDesktopCleanup() async throws { + let identity = desktopIdentity(name: "retry-cleanup", address: "100.64.12.45") + let registration = RecordingDesktopRegistration( + unregisterFailures: [identity.dnsName: 1] + ) + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + try await lifecycle.publish(identity: identity, port: 5_901) + await #expect(throws: DesktopRegistrationTestError.failed) { + try await lifecycle.removePublishedIdentities() + } + + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: StaticTailscaleRunner(output: statusJSON()), + desktopRegistration: registration, + registrationLifecycle: lifecycle, + defaults: defaults + ) + + #expect(await controller.stopAndWaitForCleanup()) + + #expect(controller.registryPhase == .notPublished) + #expect( + await registration.events + == [ + .register(identity.dnsName), + .unregister(identity.dnsName, "token:\(identity.dnsName)"), + .unregister(identity.dnsName, "token:\(identity.dnsName)"), + ] + ) + } + + @Test @MainActor + func applicationDelegateOwnsTheShareControllerUsedByTheApp() throws { + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: StaticTailscaleRunner(output: statusJSON()), + desktopRegistration: nil, + defaults: defaults + ) + let delegate = CrabfleetApplicationDelegate(shareController: controller) + + #expect(delegate.shareController === controller) + } + + @Test @MainActor + func applicationTerminationCancelsPendingAutoShareStartup() async throws { + let runner = CountingTailscaleRunner(output: statusJSON()) + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: nil, + defaults: defaults + ) + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) }, + isAutoShareRequested: { true }, + autoShareDelay: .seconds(30) + ) + let application = NSApplication.shared + + delegate.applicationDidFinishLaunching( + Notification(name: NSApplication.didFinishLaunchingNotification) + ) + #expect(delegate.applicationShouldTerminate(application) == .terminateLater) + #expect(await waitUntilAsync { replies == [true] }) + try await Task.sleep(for: .milliseconds(50)) + + #expect(await runner.callCount == 0) + #expect(controller.phase == .idle) + } + + @Test @MainActor + func applicationTerminationCancelsAutoShareDuringInitialRefresh() async throws { + let runner = SequencedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: nil, + defaults: defaults + ) + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) }, + isAutoShareRequested: { true }, + autoShareDelay: .zero + ) + let application = NSApplication.shared + + delegate.applicationDidFinishLaunching( + Notification(name: NSApplication.didFinishLaunchingNotification) + ) + #expect(await waitUntilAsync { await runner.callCount == 1 }) + #expect(delegate.applicationShouldTerminate(application) == .terminateLater) + #expect(await waitUntilAsync { replies == [true] }) + + await runner.resumeNext( + .success(.init(standardOutput: statusJSON(), standardError: "")) + ) + let continuedPreflight = await waitUntilAsync(timeout: .milliseconds(200)) { + await runner.callCount > 1 + } + if continuedPreflight { + await runner.resumeNext(.failure(CancellationError())) + } + + #expect(!continuedPreflight) + #expect(!controller.isRefreshing) + #expect(controller.phase == .idle) + } + + @Test @MainActor + func applicationTerminationWaitsForPrivateShareCleanup() async throws { + let registration = SuspendedDesktopCleanupRegistration() + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + let identity = desktopIdentity(name: "termination-cleanup", address: "100.64.12.44") + try await lifecycle.publish(identity: identity, port: 5_901) + + let runner = SuspendedTailscaleRunner() + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: runner, + desktopRegistration: registration, + registrationLifecycle: lifecycle, + defaults: defaults + ) + let startTask = Task { await controller.start() } + #expect(await waitUntilAsync { await runner.hasStarted }) + + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) } + ) + + #expect(delegate.applicationShouldTerminate(NSApplication.shared) == .terminateLater) + #expect(await waitUntilAsync { await registration.hasStartedUnregistration }) + #expect(replies.isEmpty) + + await runner.resume( + .success(.init(standardOutput: statusJSON(), standardError: "")) + ) + await startTask.value + await registration.finishUnregistration() + + #expect(await waitUntilAsync { replies == [true] }) + #expect(controller.phase == .idle) + #expect(controller.registryPhase == .notPublished) + } + + @Test @MainActor + func idleApplicationTerminationContinuesWhenRecoveryServerIsUnavailable() async throws { + try await assertIdleApplicationTerminationContinues { _ in + throw URLError(.cannotConnectToHost) + } + } + + @Test @MainActor + func idleApplicationTerminationContinuesWhenRecoverySessionHasExpired() async throws { + try await assertIdleApplicationTerminationContinues { request in + let responseURL = try #require(request.url) + return ( + Data(), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 401, + httpVersion: nil, + headerFields: nil + )) + ) + } + } + + @Test @MainActor + func applicationTerminationRetainsAmbiguousPublicationForRelaunchCleanup() async throws { + let identity = desktopIdentity(name: "durable-cleanup", address: "100.64.12.54") + let registration = RecoverableAmbiguousDesktopRegistration(recoverFailures: 1) + let suiteName = "CrabfleetMacTests.\(UUID().uuidString)" + let defaults = try #require(UserDefaults(suiteName: suiteName)) + defer { defaults.removePersistentDomain(forName: suiteName) } + let stateStore = UserDefaultsDesktopHostRegistrationStateStore(defaults: defaults) + let recoveryScope = desktopRecoveryScope() + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "durable-publication" }, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + let controller = PrivateMacShareController( + runner: StaticTailscaleRunner(output: statusJSON()), + desktopRegistration: registration, + registrationLifecycle: lifecycle, + defaults: defaults + ) + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) } + ) + + #expect(delegate.applicationShouldTerminate(NSApplication.shared) == .terminateLater) + #expect(await waitUntilAsync { replies == [true] }) + if case .failed = controller.registryPhase { + // Expected: cleanup failed, but its exact retry identity was persisted. + } else { + Issue.record("expected failed registry cleanup state") + } + + let reloadedLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + try await reloadedLifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register("durable-publication"), + .recover("durable-publication"), + .recover("durable-publication"), + .unregister("recovered:durable-publication"), + ] + ) + } + + @Test @MainActor + func persistedCleanupRecoversOwnershipWithoutStoringTheToken() async throws { + let identity = desktopIdentity(name: "persisted-cleanup", address: "100.64.12.56") + let registration = IdentityAwareAmbiguousDesktopRegistration(uncertainPublicationIDs: []) + let stateStore = ToggleDesktopRegistrationStateStore() + let recoveryScope = desktopRecoveryScope() + do { + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "persisted-publication" }, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + try await lifecycle.publish(identity: identity, port: 5_901) + } + + let persistedText = String(decoding: try #require(stateStore.data), as: UTF8.self) + #expect(!persistedText.contains("token:persisted-publication")) + + let reloadedLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + try await reloadedLifecycle.removePublishedIdentities() + + #expect( + await registration.events + == [ + .register(identity.ipv4Address, 5_901, "persisted-publication"), + .recover(identity.ipv4Address, "persisted-publication"), + .unregister(identity.ipv4Address, "recovered:persisted-publication"), + ] + ) + } + + @Test @MainActor + func persistenceFailureDoesNotAbortDesktopUnregister() async throws { + let identity = desktopIdentity(name: "persist-failure", address: "100.64.12.63") + let registration = RecordingDesktopRegistration() + let stateStore = ToggleDesktopRegistrationStateStore() + let recoveryScope = desktopRecoveryScope() + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + try await lifecycle.publish(identity: identity, port: 5_901) + stateStore.failsWrites = true + + await #expect(throws: DesktopRegistrationTestError.failed) { + try await lifecycle.removePublishedIdentities() + } + + #expect( + await registration.events + == [ + .register(identity.dnsName), + .unregister(identity.dnsName, "token:\(identity.dnsName)"), + ] + ) + } + + @Test @MainActor + func loadingPersistedLegacyCleanupClearsUnsafeState() async throws { + let identity = desktopIdentity(name: "persisted-legacy", address: "100.64.12.61") + let stateStore = ToggleDesktopRegistrationStateStore() + let recoveryScope = desktopRecoveryScope() + let persistedIdentity: [String: Any] = [ + "tailnetName": identity.tailnetName, + "loginName": identity.loginName, + "dnsName": identity.dnsName, + "hostName": identity.hostName, + "ipv4Address": identity.ipv4Address, + "userID": identity.userID, + ] + let legacyRegistration: [String: Any] = [ + "persistedIdentity": persistedIdentity, + "hostID": CrabfleetDesktopRegistration.hostID(identity: identity), + "publicationID": "legacy-publication", + "usesLegacyCleanup": true, + ] + let data = try JSONSerialization.data(withJSONObject: [ + "uncertainRegistrations": [], + "publishedRegistration": legacyRegistration, + "pendingRemovals": [legacyRegistration], + ]) + try stateStore.save(data, scope: recoveryScope) + let registration = RecordingDesktopRegistration() + var recoveryScopeRequests = 0 + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { + recoveryScopeRequests += 1 + return recoveryScope + } + ) + + try await lifecycle.removePublishedIdentities() + + #expect(stateStore.data(for: recoveryScope) == nil) + #expect(await registration.events.isEmpty) + + let reloadedLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { + recoveryScopeRequests += 1 + return recoveryScope + } + ) + try await reloadedLifecycle.removePublishedIdentities() + + #expect(recoveryScopeRequests == 1) + #expect(await registration.events.isEmpty) + } + + @Test @MainActor + func applicationTerminationContinuesWhenDurableRecoveryCannotBeUpdated() async throws { + let identity = desktopIdentity(name: "unsaved-cleanup", address: "100.64.12.55") + let registration = RecoverableAmbiguousDesktopRegistration() + let stateStore = ToggleDesktopRegistrationStateStore() + let recoveryScope = desktopRecoveryScope() + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "unsaved-publication" }, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + stateStore.failsWrites = true + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: StaticTailscaleRunner(output: statusJSON()), + desktopRegistration: registration, + registrationLifecycle: lifecycle, + defaults: defaults + ) + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) } + ) + + #expect(delegate.applicationShouldTerminate(NSApplication.shared) == .terminateLater) + #expect(await waitUntilAsync { replies == [true] }) + #expect(stateStore.data != nil) + if case .failed = controller.registryPhase { + // Expected: the existing durable retry identity remains available after relaunch. + } else { + Issue.record("expected failed registry persistence state") + } + } + + @Test @MainActor + func applicationTerminationIsCancelledForUnpersistedActiveCleanup() async throws { + let identity = desktopIdentity(name: "active-cleanup", address: "100.64.12.59") + let registration = RecordingDesktopRegistration( + unregisterFailures: [identity.dnsName: 2] + ) + let lifecycle = DesktopHostRegistrationLifecycle(registration: registration) + try await lifecycle.publish(identity: identity, port: 5_901) + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: StaticTailscaleRunner(output: statusJSON()), + desktopRegistration: registration, + registrationLifecycle: lifecycle, + defaults: defaults + ) + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) } + ) + + #expect(delegate.applicationShouldTerminate(NSApplication.shared) == .terminateLater) + #expect(await waitUntilAsync { replies == [false] }) + if case .failed = controller.registryPhase { + // Expected: no durable retry state exists for the active registration. + } else { + Issue.record("expected failed active cleanup state") + } + } + + @Test @MainActor + func persistedRecoveryDoesNotCrossDeployments() async throws { + try await assertPersistedRecoveryIsScoped( + originalScope: desktopRecoveryScope(), + otherScope: desktopRecoveryScope(origin: "https://other.example") + ) + } + + @Test @MainActor + func persistedRecoveryDoesNotCrossAccounts() async throws { + try await assertPersistedRecoveryIsScoped( + originalScope: desktopRecoveryScope(), + otherScope: desktopRecoveryScope(ownerSubject: "github:other") + ) + } + + @Test @MainActor + func definitiveRegistrationHTTPFailureDoesNotBecomeRecoveryIntent() async throws { + try await assertDefinitiveRegistrationFailureClearsIntent(.httpStatus(403)) + } + + @Test @MainActor + func redirectedRegistrationDoesNotBecomeRecoveryIntent() async throws { + try await assertDefinitiveRegistrationFailureClearsIntent(.redirect) + } + + @Test + func privateShareCanStartViewOnlyWithoutAccessibility() { + #expect( + PrivateMacSharePermissionPolicy.canStart( + identityAvailable: true, + screenRecordingGranted: true + )) + #expect( + !PrivateMacSharePermissionPolicy.canStart( + identityAvailable: false, + screenRecordingGranted: true + )) + #expect( + !PrivateMacSharePermissionPolicy.canStart( + identityAvailable: true, + screenRecordingGranted: false + )) + } + + @Test + func recognizesExplicitPrivateShareLaunchMode() { + #expect( + PrivateMacShareLaunchMode.isEnabled( + arguments: ["CrabfleetMac", "--share-this-mac"], environment: [:])) + #expect( + PrivateMacShareLaunchMode.isEnabled( + arguments: ["CrabfleetMac"], environment: ["CRABFLEET_AUTO_SHARE": "1"])) + #expect( + !PrivateMacShareLaunchMode.isEnabled( + arguments: ["CrabfleetMac"], environment: ["CRABFLEET_AUTO_SHARE": "true"])) + #expect(!PrivateMacShareLaunchMode.isEnabled(arguments: ["CrabfleetMac"], environment: [:])) + } + + @Test + func parsesExplicitVNCConnectionLaunchMode() throws { + let explicitAddress = try VNCConnectionLaunchMode.address( + arguments: ["CrabfleetMac", "--connect", "vnc://100.64.0.8:5901"], + environment: ["CRABFLEET_AUTO_CONNECT": "vnc://ignored.example:5900"] + ) + let explicit = try #require(explicitAddress) + #expect(explicit.host == "100.64.0.8") + #expect(explicit.port == 5_901) + + let environmentAddress = try VNCConnectionLaunchMode.address( + arguments: ["CrabfleetMac"], + environment: ["CRABFLEET_AUTO_CONNECT": "viewer.example:5999"] + ) + let environment = try #require(environmentAddress) + #expect(environment.host == "viewer.example") + #expect(environment.port == 5_999) + #expect( + try VNCConnectionLaunchMode.address( + arguments: ["CrabfleetMac"], environment: [:]) == nil) + } + + @Test + func rejectsMissingOrCredentialedVNCConnectionLaunchAddress() { + #expect(throws: VNCAddressError.missingHost) { + try VNCConnectionLaunchMode.address( + arguments: ["CrabfleetMac", "--connect"], environment: [:]) + } + #expect(throws: VNCAddressError.embeddedPassword) { + try VNCConnectionLaunchMode.address( + arguments: ["CrabfleetMac", "--connect", "vnc://user:secret@example.test"], + environment: [:] + ) + } + } + + @Test + func acceptsOnlineUserOnActiveTailnet() throws { + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + #expect(identity.tailnetName == "example.com") + #expect(identity.loginName == "operator@example.com") + #expect(identity.ipv4Address == "100.64.12.34") + #expect(identity.vncAddress(port: 5901) == "vnc://100.64.12.34:5901") + } + + @Test + func derivesGenericStableDesktopHostID() { + let identity = TailnetIdentity( + tailnetName: "example.com", + loginName: "operator@example.com", + dnsName: "workstation-1.example.ts.net", + hostName: "Workstation", + ipv4Address: "100.64.12.34", + userID: 42 + ) + #expect(CrabfleetDesktopRegistration.hostID(identity: identity) == "workstation-1") + + let fallback = TailnetIdentity( + tailnetName: identity.tailnetName, + loginName: identity.loginName, + dnsName: "", + hostName: identity.hostName, + ipv4Address: identity.ipv4Address, + userID: identity.userID + ) + #expect(CrabfleetDesktopRegistration.hostID(identity: fallback) == "mac-100-64-12-34") + } + + @Test + func acceptsOnlySecureCrabfleetAPIURLs() throws { + #expect( + CrabfleetDesktopRegistration.isSecureAPIURL( + try #require(URL(string: "https://fleet.example/api/fleet")))) + #expect( + CrabfleetDesktopRegistration.isSecureAPIURL( + try #require(URL(string: "http://127.0.0.1:8787")))) + #expect( + !CrabfleetDesktopRegistration.isSecureAPIURL( + try #require(URL(string: "http://fleet.example")))) + #expect( + !CrabfleetDesktopRegistration.isSecureAPIURL( + try #require(URL(string: "https://user@fleet.example")))) + #expect( + !CrabfleetDesktopRegistration.isSecureAPIURL( + try #require(URL(string: "https://fleet.example?token=value")))) + } + + @Test + func buildsAuthenticatedDesktopRegistrationRequest() throws { + let registration = try #require( + CrabfleetDesktopRegistration(environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ])) + let identity = TailnetIdentity( + tailnetName: "example.com", + loginName: "operator@example.com", + dnsName: "workstation-1.example.ts.net", + hostName: "Workstation", + ipv4Address: "100.64.12.34", + userID: 42 + ) + + let request = try registration.registrationRequest( + identity: identity, + port: 5901, + publicationID: "publication-id" + ) + #expect(request.url?.absoluteString == "https://fleet.example/api/desktop-hosts/workstation-1") + #expect(request.httpMethod == "PUT") + #expect(request.value(forHTTPHeaderField: "Cookie") == "crabbox_session=secret") + #expect( + request.value(forHTTPHeaderField: CrabfleetDesktopRegistration.ownershipModeHeader) + == CrabfleetDesktopRegistration.tokenOwnershipMode + ) + #expect( + request.value(forHTTPHeaderField: CrabfleetDesktopRegistration.publicationIDHeader) + == "publication-id" + ) + let body = try #require(request.httpBody) + let json = try #require(JSONSerialization.jsonObject(with: body) as? [String: Any]) + #expect(json["name"] as? String == "Workstation") + #expect(json["address"] as? String == "100.64.12.34") + #expect(json["port"] as? Int == 5901) + + let removal = try registration.removalRequest( + identity: identity, + ownershipToken: "desktop-ownership-token" + ) + #expect(removal.url == request.url) + #expect(removal.httpMethod == "DELETE") + #expect(removal.value(forHTTPHeaderField: "Cookie") == "crabbox_session=secret") + #expect( + removal.value(forHTTPHeaderField: "X-Crabfleet-Ownership-Token") + == "desktop-ownership-token" + ) + #expect(removal.httpBody == nil) + } + + @Test + func desktopRegistrationScopesRecoveryToNormalizedOriginAndStableOwner() async throws { + let transport = DesktopRegistrationTransport { request in + let responseURL = try #require(request.url) + #expect(responseURL.host?.lowercased() == "fleet.example") + #expect(responseURL.path == "/api/native/v1/session") + #expect(request.httpMethod == "GET") + #expect(request.value(forHTTPHeaderField: "Cookie") == "crabbox_session=secret") + return ( + Data(#"{"user":{"subject":"github:123"}}"#.utf8), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://FLEET.EXAMPLE:443/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + + #expect( + try await registration.recoveryScope() + == DesktopHostRegistrationRecoveryScope( + apiOrigin: "https://fleet.example", + ownerSubject: "github:123" + ) + ) + } + + @Test + func desktopRegistrationReturnsTheServerOwnershipToken() async throws { + let transport = DesktopRegistrationTransport { request in + let responseURL = try #require(request.url) + return ( + Data(#"{"host":{"id":"workstation"},"ownershipToken":"server-ownership-token"}"#.utf8), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + #expect( + try await registration.register( + identity: identity, + port: 5_901, + publicationID: "publication-id" + ) + == "server-ownership-token" + ) + } + + @Test + func desktopRegistrationFallsBackToLegacyCleanupForOldServers() async throws { + let transport = DesktopRegistrationTransport { request in + let responseURL = try #require(request.url) + return ( + Data(#"{"host":{"id":"workstation"}}"#.utf8), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + #expect( + try await registration.register( + identity: identity, + port: 5_901, + publicationID: "publication-id" + ) == nil + ) + let removal = try registration.removalRequest(identity: identity, ownershipToken: nil) + #expect(removal.value(forHTTPHeaderField: "X-Crabfleet-Ownership-Token") == nil) + } + + @Test + func desktopRegistrationRecoversOnlyTheMatchingPublication() async throws { + let transport = DesktopRegistrationTransport { request in + let responseURL = try #require(request.url) + #expect(request.httpMethod == "POST") + #expect(responseURL.query == "recover=1") + let body = try #require(request.httpBody) + let json = try #require(JSONSerialization.jsonObject(with: body) as? [String: Any]) + #expect(json["publicationID"] as? String == "publication-id") + return ( + Data(#"{"ownershipToken":"server-ownership-token"}"#.utf8), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + #expect( + try await registration.recover( + identity: identity, + publicationID: "publication-id" + ) == "server-ownership-token" + ) + } + + @Test + func desktopRegistrationTreatsMissingRecoveryRouteAsUncertain() async throws { + let transport = DesktopRegistrationTransport { request in + let responseURL = try #require(request.url) + return ( + Data(), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 404, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await registration.recover( + identity: identity, + publicationID: "publication-id" + ) + } + } + + @Test @MainActor + func legacyRecoveryRoutePreservesUncertainPublicationWithoutDeletingNewerPublisher() + async throws + { + let transport = LegacyDesktopServerTransport() + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "legacy-publication" } + ) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + await transport.publishNewerEndpoint() + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.removePublishedIdentities() + } + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await lifecycle.removePublishedIdentities() + } + + #expect(await transport.activeEndpoint == "newer-publisher") + #expect(await transport.events == [.register, .recover, .recover]) + } + + @Test + func desktopRegistrationTreatsMalformedCommittedResponsesAsUncertain() async throws { + let transport = DesktopRegistrationTransport { request in + let responseURL = try #require(request.url) + return ( + Data(#"{"host":{"id":"workstation"},"ownershipToken":null}"#.utf8), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await registration.register( + identity: identity, + port: 5_901, + publicationID: "publication-id" + ) + } + } + + @Test + func desktopRegistrationTreatsTransportFailuresAsUncertain() async throws { + let transport = DesktopRegistrationTransport { _ in + throw URLError(.timedOut) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await registration.register( + identity: identity, + port: 5_901, + publicationID: "publication-id" + ) + } + } + + @Test + func desktopRegistrationRejectsRedirectedResponses() async throws { + let redirectedURL = try #require(URL(string: "https://login.example.test/desktop-host")) + let transport = DesktopRegistrationTransport { _ in + ( + Data(), + try #require( + HTTPURLResponse( + url: redirectedURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + } + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + + await #expect(throws: DesktopHostRegistrationError.redirectRejected) { + try await registration.register( + identity: identity, + port: 5_901, + publicationID: "publication-id" + ) + } + } + + @Test + func rejectsInvalidTailnetAndIdentityFields() throws { + var value = statusJSON() + value = value.replacingOccurrences( + of: #""Name": "example.com""#, with: #""Name": """#) + let missingTailnet = try JSONDecoder().decode( + TailscaleStatusDocument.self, + from: Data(value.utf8) + ) + #expect(throws: PrivateMacShareError.invalidTailnetIdentity) { + try TailnetIdentityPolicy.identity(from: missingTailnet) + } + + value = statusJSON().replacingOccurrences( + of: "operator@example.com", + with: "" + ) + let missingUser = try JSONDecoder().decode( + TailscaleStatusDocument.self, + from: Data(value.utf8) + ) + #expect(throws: PrivateMacShareError.invalidTailnetUser) { + try TailnetIdentityPolicy.identity(from: missingUser) + } + } + + @Test + func recognizesOnlyTailscaleIPv4Range() { + #expect(TailnetIdentityPolicy.isTailscaleIPv4("100.64.0.1")) + #expect(TailnetIdentityPolicy.isTailscaleIPv4("100.127.255.254")) + #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.63.255.255")) + #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.128.0.1")) + #expect(!TailnetIdentityPolicy.isTailscaleIPv4("10.0.0.1")) + #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.64.invalid.1.2")) + #expect(!TailnetIdentityPolicy.isTailscaleIPv4("100.64..1")) + } + + @Test + func validatesBoundedTailnetIdentityFields() { + #expect(TailnetIdentityPolicy.isValidTailnetName("example.com")) + #expect(TailnetIdentityPolicy.isValidTailnetName("example.github")) + #expect(!TailnetIdentityPolicy.isValidTailnetName("")) + #expect(!TailnetIdentityPolicy.isValidTailnetName(" example.com")) + #expect(!TailnetIdentityPolicy.isValidTailnetName("bad\nname")) + #expect(!TailnetIdentityPolicy.isValidTailnetName(String(repeating: "a", count: 254))) + + #expect(TailnetIdentityPolicy.isValidLogin("operator@example.com")) + #expect(TailnetIdentityPolicy.isValidLogin("github-user")) + #expect(!TailnetIdentityPolicy.isValidLogin("")) + #expect(!TailnetIdentityPolicy.isValidLogin("github-user ")) + #expect(!TailnetIdentityPolicy.isValidLogin("bad\u{0}login")) + #expect(!TailnetIdentityPolicy.isValidLogin(String(repeating: "a", count: 321))) + } + + @Test + func authorizesOnlySameTailnetUserAndExactPeerAddress() async throws { + let identity = try TailnetIdentityPolicy.identity(from: statusDocument()) + let accepted = TailnetPeerAuthorizer( + runner: StaticTailscaleRunner(output: whoisJSON(login: identity.loginName)), + expectedIdentity: identity + ) + #expect(await accepted.authorize(remoteAddress: "100.100.10.20")) + + let otherUser = TailnetPeerAuthorizer( + runner: StaticTailscaleRunner(output: whoisJSON(login: "other@example.com")), + expectedIdentity: identity + ) + let otherUserID = TailnetPeerAuthorizer( + runner: StaticTailscaleRunner( + output: whoisJSON(login: identity.loginName, userID: 43)), + expectedIdentity: identity + ) + let otherAddress = TailnetPeerAuthorizer( + runner: StaticTailscaleRunner( + output: whoisJSON(login: identity.loginName, addresses: ["100.100.10.21/32"])), + expectedIdentity: identity + ) + let unauthorizedNode = TailnetPeerAuthorizer( + runner: StaticTailscaleRunner( + output: whoisJSON(login: identity.loginName, machineAuthorized: false)), + expectedIdentity: identity + ) + #expect(!(await otherUser.authorize(remoteAddress: "100.100.10.20"))) + #expect(!(await otherUserID.authorize(remoteAddress: "100.100.10.20"))) + #expect(!(await otherAddress.authorize(remoteAddress: "100.100.10.20"))) + #expect(!(await unauthorizedNode.authorize(remoteAddress: "100.100.10.20"))) + #expect(!(await accepted.authorize(remoteAddress: "192.168.1.4"))) + #expect(!(await accepted.authorize(remoteAddress: identity.ipv4Address))) + } + + @Test @MainActor + func expiresIncompleteRFBHandshakeAndReleasesInput() async throws { + let identity = TailnetIdentity( + tailnetName: "example.com", + loginName: "tester@example.com", + dnsName: "workstation.example.ts.net.", + hostName: "Workstation", + ipv4Address: "127.0.0.1", + userID: 42 + ) + let capture = MacScreenCapture() + let input = RemoteInputRecorder() + let events = RFBEventRecorder() + let port: UInt16 = 5_923 + let server = TailnetRFBServer( + identity: identity, + runner: StaticTailscaleRunner(output: ""), + capture: capture, + descriptor: .init( + displayID: 0, + displayBounds: CGRect(x: 0, y: 0, width: 64, height: 64), + frameWidth: 64, + frameHeight: 64, + sourcePixelWidth: 64, + sourcePixelHeight: 64 + ), + input: input, + peerAuthorizer: LoopbackPeerAuthorizer(), + port: port, + handshakeTimeout: .milliseconds(100), + eventHandler: { events.append($0) } + ) + try server.start() + defer { server.stop() } + try await Task.sleep(for: .milliseconds(100)) + + let connection = NWConnection( + host: "127.0.0.1", + port: try #require(NWEndpoint.Port(rawValue: port)), + using: .tcp + ) + connection.start(queue: .global(qos: .userInitiated)) + defer { connection.cancel() } + + try await waitFor { + events.values.contains { + if case .sessionFailed(let message) = $0 { + return message.contains("handshake timed out") + } + return false + } + } + #expect(input.releaseCount == 1) + } + + @Test + func keepsNewestCapturedFrameWhenUpdatesArriveOutOfOrder() async throws { + let store = CapturedDesktopFrameStore() + await store.update(.init(jpegData: Data([2]), sequence: 2, width: 2, height: 2)) + await store.update(.init(jpegData: Data([1]), sequence: 1, width: 2, height: 2)) + + #expect(await store.latest()?.sequence == 2) + } + + @Test + func buildsTightJPEGFramebufferUpdate() throws { + let jpeg = Data([0xFF, 0xD8, 0xFF, 0xD9]) + let frame = CapturedDesktopFrame(jpegData: jpeg, sequence: 7, width: 1_600, height: 900) + let packet = try RFBWire.tightJPEGUpdate(frame: frame) + + #expect(packet[0] == 0) + #expect(packet.readUInt16(at: 2) == 1) + #expect(packet.readUInt16(at: 8) == 1_600) + #expect(packet.readUInt16(at: 10) == 900) + #expect(packet.readInt32(at: 12) == RFBWire.tightEncoding) + #expect(packet[16] == 0x90) + #expect(packet[17] == 4) + #expect(packet.suffix(4) == jpeg) + } + + @Test + func encodesTightCompactLengths() { + #expect(RFBWire.tightCompactLength(0) == Data([0x00])) + #expect(RFBWire.tightCompactLength(127) == Data([0x7F])) + #expect(RFBWire.tightCompactLength(128) == Data([0x80, 0x01])) + #expect(RFBWire.tightCompactLength(16_383) == Data([0xFF, 0x7F])) + #expect(RFBWire.tightCompactLength(16_384) == Data([0x80, 0x80, 0x01])) + } + + @Test + func scalesCaptureWithinBoundedEvenDimensions() { + let retina = MacScreenCapture.captureDimensions(sourceWidth: 5_120, sourceHeight: 2_880) + #expect(retina.width == 2_560) + #expect(retina.height == 1_440) + #expect(retina.width.isMultiple(of: 2)) + #expect(retina.height.isMultiple(of: 2)) + + let small = MacScreenCapture.captureDimensions(sourceWidth: 1_280, sourceHeight: 800) + #expect(small.width == 1_280) + #expect(small.height == 800) + } + + @Test + func mapsRFBKeysymsToMacKeys() { + #expect(MacRemoteInputController.keyCode(for: 0x61) != nil) + #expect(MacRemoteInputController.keyCode(for: 0xFF51) != nil) + #expect(MacRemoteInputController.keyCode(for: 0xFFE7) != nil) + #expect(MacRemoteInputController.keyCode(for: 0x1F980) == nil) + } + + @Test + func inputSessionFinishWaitsForProducersAndRejectsLateInput() { + let input = BlockingRemoteInputRecorder() + let gate = RemoteInputSessionGate(input: input, viewOnly: false) + let producerFinished = DispatchSemaphore(value: 0) + let finishAttempted = DispatchSemaphore(value: 0) + let finishCompleted = DispatchSemaphore(value: 0) + + DispatchQueue.global().async { + gate.keyEvent(down: true, keysym: 0x61) + producerFinished.signal() + } + #expect(input.waitForKeyEntry()) + + DispatchQueue.global().async { + finishAttempted.signal() + gate.finish() + finishCompleted.signal() + } + #expect(finishAttempted.wait(timeout: .now() + 1) == .success) + #expect(input.events == [.key(down: true, keysym: 0x61)]) + #expect(finishCompleted.wait(timeout: .now() + 0.01) == .timedOut) + + input.allowKeyReturn() + #expect(producerFinished.wait(timeout: .now() + 1) == .success) + #expect(finishCompleted.wait(timeout: .now() + 1) == .success) + #expect(input.events == [.key(down: true, keysym: 0x61), .release]) + + gate.keyEvent(down: false, keysym: 0x61) + gate.pointerEvent(buttonMask: 0x01, x: 1, y: 1) + #expect(input.events == [.key(down: true, keysym: 0x61), .release]) + } + + @Test + func retriesHeldInputReleaseAfterAccessibilityReturns() async { + let trust = AccessibilityTrust(granted: true) + let events = RemoteInputEventRecorder() + let controller = MacRemoteInputController( + descriptor: CapturedDisplayDescriptor( + displayID: 1, + displayBounds: CGRect(x: 0, y: 0, width: 100, height: 100), + frameWidth: 100, + frameHeight: 100, + sourcePixelWidth: 100, + sourcePixelHeight: 100 + ), + accessibilityGranted: { trust.isGranted() }, + pendingReleaseRetryDelay: .milliseconds(10), + keyEventPoster: { down, keysym in + events.append(.key(down: down, keysym: keysym)) + }, + mouseEventPoster: { type, _, button in + events.append(.mouse(type: type, button: button)) + } + ) + + controller.keyEvent(down: true, keysym: 0x61) + controller.pointerEvent(buttonMask: 0x01, x: 50, y: 50) + #expect(await waitUntilAsync { + events.contains(.key(down: true, keysym: 0x61)) + && events.contains(.mouse(type: .leftMouseDown, button: .left)) + }) + + let checksBeforeRevocation = trust.checkCount + trust.setGranted(false) + controller.releaseAllInput() + #expect(await waitUntilAsync { + trust.checkCount > checksBeforeRevocation + }) + #expect(!events.contains(.key(down: false, keysym: 0x61))) + #expect(!events.contains(.mouse(type: .leftMouseUp, button: .left))) + + trust.setGranted(true) + #expect(await waitUntilAsync { + events.contains(.key(down: false, keysym: 0x61)) + && events.contains(.mouse(type: .leftMouseUp, button: .left)) + }) + } + + @Test + func pendingInputReleaseRetainsControllerThroughTeardown() async { + let trust = AccessibilityTrust(granted: true) + let events = RemoteInputEventRecorder() + var controller: MacRemoteInputController? = MacRemoteInputController( + descriptor: CapturedDisplayDescriptor( + displayID: 1, + displayBounds: CGRect(x: 0, y: 0, width: 100, height: 100), + frameWidth: 100, + frameHeight: 100, + sourcePixelWidth: 100, + sourcePixelHeight: 100 + ), + accessibilityGranted: { trust.isGranted() }, + pendingReleaseRetryDelay: .milliseconds(10), + keyEventPoster: { down, keysym in + events.append(.key(down: down, keysym: keysym)) + } + ) + weak var retainedController = controller + + controller?.keyEvent(down: true, keysym: 0x61) + #expect(await waitUntilAsync { + events.contains(.key(down: true, keysym: 0x61)) + }) + let checksBeforeRevocation = trust.checkCount + trust.setGranted(false) + controller?.releaseAllInput() + #expect(await waitUntilAsync { + trust.checkCount > checksBeforeRevocation + }) + controller = nil + + #expect(retainedController != nil) + trust.setGranted(true) + #expect(await waitUntilAsync { + events.contains(.key(down: false, keysym: 0x61)) + }) + #expect(await waitUntilAsync { retainedController == nil }) + } + + @Test + func pendingInputReleaseStopsRetryingAfterTeardownBudgetExpires() async { + let trust = AccessibilityTrust(granted: true) + let events = RemoteInputEventRecorder() + var controller: MacRemoteInputController? = MacRemoteInputController( + descriptor: CapturedDisplayDescriptor( + displayID: 1, + displayBounds: CGRect(x: 0, y: 0, width: 100, height: 100), + frameWidth: 100, + frameHeight: 100, + sourcePixelWidth: 100, + sourcePixelHeight: 100 + ), + accessibilityGranted: { trust.isGranted() }, + pendingReleaseRetryDelay: .milliseconds(10), + pendingReleaseRetryLimit: 2, + keyEventPoster: { down, keysym in + events.append(.key(down: down, keysym: keysym)) + } + ) + weak var retainedController = controller + + controller?.keyEvent(down: true, keysym: 0x61) + #expect(await waitUntilAsync { + events.contains(.key(down: true, keysym: 0x61)) + }) + trust.setGranted(false) + let checksBeforeRelease = trust.checkCount + controller?.releaseAllInput() + controller = nil + + #expect(await waitUntilAsync { + trust.checkCount >= checksBeforeRelease + 3 + }) + #expect(await waitUntilAsync { retainedController == nil }) + #expect(!events.contains(.key(down: false, keysym: 0x61))) + } + + @Test + func emptyInputReleaseDoesNotRetainController() async { + let trust = AccessibilityTrust(granted: false) + var controller: MacRemoteInputController? = MacRemoteInputController( + descriptor: CapturedDisplayDescriptor( + displayID: 1, + displayBounds: CGRect(x: 0, y: 0, width: 100, height: 100), + frameWidth: 100, + frameHeight: 100, + sourcePixelWidth: 100, + sourcePixelHeight: 100 + ), + accessibilityGranted: { trust.isGranted() }, + pendingReleaseRetryDelay: .milliseconds(10) + ) + weak var retainedController = controller + let checksBeforeRelease = trust.checkCount + + controller?.releaseAllInput() + controller = nil + + #expect(await waitUntilAsync { retainedController == nil }) + #expect(trust.checkCount == checksBeforeRelease) + } + + @Test + func decodesX11UnicodeKeysymsForMacInput() { + #expect(MacRemoteInputController.unicodeScalar(for: 0x0100_03BB) == "λ") + #expect(MacRemoteInputController.unicodeScalar(for: 0x0101_F980) == "🦀") + #expect(MacRemoteInputController.unicodeScalar(for: 0x0111_0000) == nil) + } + + @Test @MainActor + func servesRoyalVNCKitOverTheCurrentTailnet() async throws { + guard ProcessInfo.processInfo.environment["CRABFLEET_TAILNET_RFB_SMOKE"] == "1" else { + return + } + + let runner = try SystemTailscaleCommandRunner() + let status = try await runner.run(arguments: ["status", "--json"]) + let document = try JSONDecoder().decode( + TailscaleStatusDocument.self, + from: Data(status.standardOutput.utf8) + ) + let identity = try TailnetIdentityPolicy.identity(from: document) + let capture = MacScreenCapture() + let jpeg = try #require(testJPEG()) + await capture.frameStore.update( + .init(jpegData: jpeg, sequence: 1, width: 64, height: 64) + ) + + let port: UInt16 = 5_909 + let server = TailnetRFBServer( + identity: identity, + runner: runner, + capture: capture, + descriptor: .init( + displayID: 0, + displayBounds: CGRect(x: 0, y: 0, width: 64, height: 64), + frameWidth: 64, + frameHeight: 64, + sourcePixelWidth: 64, + sourcePixelHeight: 64 + ), + input: NoopRemoteInput(), + port: port, + eventHandler: { _ in } + ) + try server.start() + defer { server.stop() } + try await Task.sleep(for: .milliseconds(250)) + + let session = VNCSessionController() + session.connect( + host: identity.ipv4Address, + port: port, + username: "", + password: "", + clipboardEnabled: false + ) + defer { session.disconnect() } + + let clock = ContinuousClock() + let deadline = clock.now.advanced(by: .seconds(15)) + while clock.now < deadline { + if session.phase == .connected && session.framebufferUpdateCount > 0 { break } + try await Task.sleep(for: .milliseconds(25)) + } + #expect(session.phase == .connected) + #expect(session.framebufferUpdateCount > 0) + #expect(session.framebuffer?.size.width == 64) + #expect(session.framebuffer?.size.height == 64) + } + + @Test @MainActor + func syncsUTF8ClipboardAndNegotiatesResizeOverLoopback() async throws { + // Full-protocol end-to-end: the production server and the RoyalVNCKit + // client exchange handshake, Tight frames, Extended Clipboard, and + // ExtendedDesktopSize over a real TCP connection on loopback. The + // tailnet-specific pieces (address binding, whois) are injected. + let identity = TailnetIdentity( + tailnetName: "example.com", + loginName: "tester@example.com", + dnsName: "workstation.example.ts.net.", + hostName: "Workstation", + ipv4Address: "127.0.0.1", + userID: 42 + ) + let capture = MacScreenCapture() + let jpeg = try #require(testJPEG()) + await capture.frameStore.update( + .init(jpegData: jpeg, sequence: 1, width: 64, height: 64) + ) + + let hostPasteboard = NSPasteboard(name: .init("CrabfleetMacTests.host.\(UUID().uuidString)")) + hostPasteboard.clearContents() + let hostClipboard = HostClipboardBridge(pasteboard: hostPasteboard, pollingInterval: 0.02) + + let port: UInt16 = 5_921 + let server = TailnetRFBServer( + identity: identity, + runner: StaticTailscaleRunner(output: ""), + capture: capture, + descriptor: .init( + displayID: 0, + displayBounds: CGRect(x: 0, y: 0, width: 64, height: 64), + frameWidth: 64, + frameHeight: 64, + sourcePixelWidth: 256, + sourcePixelHeight: 256 + ), + input: NoopRemoteInput(), + clipboard: hostClipboard, + peerAuthorizer: LoopbackPeerAuthorizer(), + port: port, + eventHandler: { _ in } + ) + try server.start() + defer { server.stop() } + try await Task.sleep(for: .milliseconds(250)) + + let viewerPasteboard = NSPasteboard( + name: .init("CrabfleetMacTests.viewer.\(UUID().uuidString)") + ) + viewerPasteboard.clearContents() + let coordinator = ClipboardCoordinator(pasteboard: viewerPasteboard, pollingInterval: 0.02) + let session = VNCSessionController(targetID: "smoke", clipboardCoordinator: coordinator) + coordinator.focus(session: session, targetID: "smoke") + session.connect( + host: identity.ipv4Address, + port: port, + username: "", + password: "" + ) + defer { session.disconnect() } + + // The Extended Clipboard caps handshake must complete on the client. + try await waitFor { session.connection?.supportsUTF8Clipboard == true } + + // Viewer to host: emoji only survives the extended UTF-8 path. + viewerPasteboard.clearContents() + viewerPasteboard.setString("client copy 🚀", forType: .string) + try await waitFor { hostPasteboard.string(forType: .string) == "client copy 🚀" } + + // Host to server-cut-text: the host push lands on the viewer pasteboard. + hostPasteboard.clearContents() + hostPasteboard.setString("host copy 🦀", forType: .string) + try await waitFor { viewerPasteboard.string(forType: .string) == "host copy 🦀" } + + // The ExtendedDesktopSize announce must unlock client resize requests. + try await waitFor { session.requestDesktopSize(.init(width: 128, height: 128)) } + + // The stream must keep flowing after the resize exchange (this test + // fixture has no live capture stream, so the server answers the resize + // with an out-of-resources status and continues serving frames). An + // unchanged frame is deduplicated into rectangle-free heartbeats, so new + // content must arrive as a fresh framebuffer update. + let updates = session.framebufferUpdateCount + await capture.frameStore.update( + .init(jpegData: jpeg, sequence: 2, width: 64, height: 64) + ) + try await waitFor { session.framebufferUpdateCount > updates } + #expect(session.phase == .connected) + } + + @MainActor + private func waitFor( + timeout: Duration = .seconds(15), + _ condition: @escaping @MainActor () -> Bool + ) async throws { + let clock = ContinuousClock() + let deadline = clock.now.advanced(by: timeout) + while clock.now < deadline { + if condition() { return } + try await Task.sleep(for: .milliseconds(25)) + } + #expect(condition()) + } + + private func statusDocument() throws -> TailscaleStatusDocument { + try JSONDecoder().decode(TailscaleStatusDocument.self, from: Data(statusJSON().utf8)) + } + + private func desktopIdentity(name: String, address: String) -> TailnetIdentity { + TailnetIdentity( + tailnetName: "example.com", + loginName: "operator@example.com", + dnsName: "\(name).example.ts.net", + hostName: name, + ipv4Address: address, + userID: 42 + ) + } + + private func statusJSON() -> String { + """ + { + "BackendState": "Running", + "CurrentTailnet": { "Name": "example.com" }, + "Self": { + "DNSName": "workstation.example.ts.net.", + "HostName": "Workstation", "Online": true, "TailscaleIPs": ["100.64.12.34", "fd7a:115c:a1e0::1"], "UserID": 42 @@ -534,72 +2261,959 @@ struct PrivateMacShareTests { "42": { "LoginName": "operator@example.com" } } } - """ + """ + } + + private func whoisJSON( + login: String, + userID: Int64 = 42, + addresses: [String] = ["100.100.10.20/32"], + machineAuthorized: Bool = true + ) -> String { + let encodedAddresses = addresses.map { "\"\($0)\"" }.joined(separator: ", ") + return """ + { + "Node": { + "Addresses": [\(encodedAddresses)], + "MachineAuthorized": \(machineAuthorized), + "User": \(userID) + }, + "UserProfile": { "LoginName": "\(login)" } + } + """ + } + + private func testJPEG() -> Data? { + guard + let bitmap = NSBitmapImageRep( + bitmapDataPlanes: nil, + pixelsWide: 64, + pixelsHigh: 64, + bitsPerSample: 8, + samplesPerPixel: 4, + hasAlpha: true, + isPlanar: false, + colorSpaceName: .deviceRGB, + bytesPerRow: 0, + bitsPerPixel: 0 + ) + else { return nil } + bitmap.setColor( + NSColor(deviceRed: 0.25, green: 0.9, blue: 0.7, alpha: 1), + atX: 16, + y: 16 + ) + bitmap.setColor( + NSColor(deviceRed: 1, green: 0.55, blue: 0.2, alpha: 1), + atX: 48, + y: 48 + ) + return bitmap.representation(using: .jpeg, properties: [.compressionFactor: 0.8]) + } + + @Test + func buildScriptRecreatesTheAppBundleBeforeAssembly() throws { + let testFile = URL(fileURLWithPath: #filePath) + let script = + testFile + .deletingLastPathComponent() + .deletingLastPathComponent() + .deletingLastPathComponent() + .appendingPathComponent("scripts/build-app.sh") + let contents = try String(contentsOf: script, encoding: .utf8) + let removal = try #require(contents.range(of: "rm -rf \"$app_dir\"")) + let assembly = try #require(contents.range(of: "mkdir -p \"$macos_dir\" \"$resources_dir\"")) + #expect(removal.lowerBound < assembly.lowerBound) + } + + @MainActor + private func assertPersistedRecoveryIsScoped( + originalScope: DesktopHostRegistrationRecoveryScope, + otherScope: DesktopHostRegistrationRecoveryScope + ) async throws { + let identity = desktopIdentity(name: "scoped-recovery", address: "100.64.12.57") + let registration = RecoverableAmbiguousDesktopRegistration() + let stateStore = ToggleDesktopRegistrationStateStore() + let originalLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { "scoped-publication" }, + stateStore: stateStore, + recoveryScopeProvider: { originalScope } + ) + await #expect(throws: DesktopHostRegistrationResultUncertainError.self) { + try await originalLifecycle.publish(identity: identity, port: 5_901) + } + + let otherLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { otherScope } + ) + try await otherLifecycle.removePublishedIdentities() + #expect(await registration.events == [.register("scoped-publication")]) + + let reloadedLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { originalScope } + ) + try await reloadedLifecycle.removePublishedIdentities() + #expect( + await registration.events + == [ + .register("scoped-publication"), + .recover("scoped-publication"), + .unregister("recovered:scoped-publication"), + ] + ) + } + + @MainActor + private func assertIdleApplicationTerminationContinues( + transportHandler: @escaping (URLRequest) throws -> (Data, HTTPURLResponse) + ) async throws { + var recoveryScopeRequests = 0 + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: DesktopRegistrationTransport { request in + recoveryScopeRequests += 1 + return try transportHandler(request) + } + )) + let defaults = try #require( + UserDefaults(suiteName: "CrabfleetMacTests.\(UUID().uuidString)") + ) + let controller = PrivateMacShareController( + runner: StaticTailscaleRunner(output: statusJSON()), + desktopRegistration: registration, + defaults: defaults + ) + var replies: [Bool] = [] + let delegate = CrabfleetApplicationDelegate( + shareController: controller, + replyToTerminationRequest: { replies.append($0) } + ) + + #expect(delegate.applicationShouldTerminate(NSApplication.shared) == .terminateLater) + #expect(await waitUntilAsync { replies == [true] }) + #expect(recoveryScopeRequests == 0) + #expect(controller.phase == .idle) + #expect(controller.registryPhase == .notPublished) + } + + @MainActor + private func assertDefinitiveRegistrationFailureClearsIntent( + _ failure: DefinitiveRegistrationFailureTransport.Failure + ) async throws { + let transport = DefinitiveRegistrationFailureTransport(failure: failure) + let registration = try #require( + CrabfleetDesktopRegistration( + environment: [ + "CRABFLEET_API_URL": "https://fleet.example/api/fleet", + "CRABFLEET_SESSION_COOKIE": "crabbox_session=secret", + ], + transport: transport + )) + let stateStore = ToggleDesktopRegistrationStateStore() + let recoveryScope = desktopRecoveryScope() + var publicationIDs = ["publication-a", "publication-b"] + let lifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + createPublicationID: { publicationIDs.removeFirst() }, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + let identity = desktopIdentity(name: "definitive-failure", address: "100.64.12.58") + + await #expect(throws: DesktopHostRegistrationError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + await #expect(throws: DesktopHostRegistrationError.self) { + try await lifecycle.publish(identity: identity, port: 5_901) + } + #expect(await transport.publicationIDs == ["publication-a", "publication-b"]) + #expect(stateStore.data(for: recoveryScope) == nil) + + let reloadedLifecycle = DesktopHostRegistrationLifecycle( + registration: registration, + stateStore: stateStore, + recoveryScopeProvider: { recoveryScope } + ) + try await reloadedLifecycle.removePublishedIdentities() + #expect(await transport.publicationIDs == ["publication-a", "publication-b"]) + } +} + +private struct StaticTailscaleRunner: TailscaleCommandRunning { + let output: String + + func run(arguments: [String]) async throws -> TailscaleCommandResult { + .init(standardOutput: output, standardError: "") + } +} + +private actor CountingTailscaleRunner: TailscaleCommandRunning { + let output: String + private(set) var callCount = 0 + + init(output: String) { + self.output = output + } + + func run(arguments: [String]) async throws -> TailscaleCommandResult { + callCount += 1 + return .init(standardOutput: output, standardError: "") + } +} + +private struct NoopRemoteInput: RemoteInputForwarding { + func keyEvent(down: Bool, keysym: UInt32) {} + func pointerEvent(buttonMask: UInt8, x: UInt16, y: UInt16) {} +} + +private struct LoopbackPeerAuthorizer: TailnetPeerAuthorizing { + func authorize(remoteAddress: String) async -> Bool { + remoteAddress == "127.0.0.1" + } +} + +private actor SuspendedTailscaleRunner: TailscaleCommandRunning { + private var continuation: CheckedContinuation? + private(set) var hasStarted = false + + func run(arguments: [String]) async throws -> TailscaleCommandResult { + hasStarted = true + return try await withCheckedThrowingContinuation { continuation in + self.continuation = continuation + } + } + + func resume(_ result: Result) { + continuation?.resume(with: result) + continuation = nil + } +} + +private actor SequencedTailscaleRunner: TailscaleCommandRunning { + private var continuations: [CheckedContinuation] = [] + private(set) var callCount = 0 + + func run(arguments: [String]) async throws -> TailscaleCommandResult { + callCount += 1 + return try await withCheckedThrowingContinuation { continuation in + continuations.append(continuation) + } + } + + func resumeNext(_ result: Result) { + guard !continuations.isEmpty else { return } + continuations.removeFirst().resume(with: result) + } +} + +private actor AsyncInvocationState { + private(set) var started = false + private(set) var finished = false + + func markStarted() { + started = true + } + + func markFinished() { + finished = true + } +} + +private actor SuspendedAsyncOperation { + private var continuation: CheckedContinuation? + private(set) var invocationCount = 0 + + func run() async { + invocationCount += 1 + await withCheckedContinuation { continuation in + self.continuation = continuation + } + } + + func finish() { + continuation?.resume() + continuation = nil + } +} + +private actor SuspendedDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case registerStarted + case registerFinished + case unregisterStarted + } + + private var registrationContinuation: CheckedContinuation? + private(set) var events: [Event] = [] + + var hasStartedRegistration: Bool { + registrationContinuation != nil + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.registerStarted) + await withCheckedContinuation { continuation in + registrationContinuation = continuation + } + events.append(.registerFinished) + return "registration-token" } - private func whoisJSON( - login: String, - userID: Int64 = 42, - addresses: [String] = ["100.100.10.20/32"], - machineAuthorized: Bool = true - ) -> String { - let encodedAddresses = addresses.map { "\"\($0)\"" }.joined(separator: ", ") - return """ - { - "Node": { - "Addresses": [\(encodedAddresses)], - "MachineAuthorized": \(machineAuthorized), - "User": \(userID) - }, - "UserProfile": { "LoginName": "\(login)" } - } - """ + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + nil } - private func testJPEG() -> Data? { - guard - let bitmap = NSBitmapImageRep( - bitmapDataPlanes: nil, - pixelsWide: 64, - pixelsHigh: 64, - bitsPerSample: 8, - samplesPerPixel: 4, - hasAlpha: true, - isPlanar: false, - colorSpaceName: .deviceRGB, - bytesPerRow: 0, - bitsPerPixel: 0 - ) - else { return nil } - bitmap.setColor( - NSColor(deviceRed: 0.25, green: 0.9, blue: 0.7, alpha: 1), - atX: 16, - y: 16 - ) - bitmap.setColor( - NSColor(deviceRed: 1, green: 0.55, blue: 0.2, alpha: 1), - atX: 48, - y: 48 - ) - return bitmap.representation(using: .jpeg, properties: [.compressionFactor: 0.8]) + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + #expect(ownershipToken == "registration-token") + events.append(.unregisterStarted) + } + + func finishRegistration() { + registrationContinuation?.resume() + registrationContinuation = nil } } -private struct StaticTailscaleRunner: TailscaleCommandRunning { - let output: String +private enum DesktopRegistrationTestError: Error { + case failed +} - func run(arguments: [String]) async throws -> TailscaleCommandResult { - .init(standardOutput: output, standardError: "") +private actor RecordingDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register(String) + case unregister(String, String?) + } + + private var registerFailures: [String: Int] + private var unregisterFailures: [String: Int] + private(set) var events: [Event] = [] + + init( + registerFailures: [String: Int] = [:], + unregisterFailures: [String: Int] = [:] + ) { + self.registerFailures = registerFailures + self.unregisterFailures = unregisterFailures + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register(identity.dnsName)) + if consumeFailure(for: identity.dnsName, from: ®isterFailures) { + throw DesktopRegistrationTestError.failed + } + return "token:\(identity.dnsName)" + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + nil + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + events.append(.unregister(identity.dnsName, ownershipToken)) + if consumeFailure(for: identity.dnsName, from: &unregisterFailures) { + throw DesktopRegistrationTestError.failed + } + } + + private func consumeFailure( + for identity: String, + from failures: inout [String: Int] + ) -> Bool { + guard let remaining = failures[identity], remaining > 0 else { return false } + failures[identity] = remaining - 1 + return true } } -private struct NoopRemoteInput: RemoteInputForwarding { +private actor MutableIdentityDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register(String) + case unregister(String, String?) + } + + private(set) var events: [Event] = [] + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register(identity.ipv4Address)) + return "token:\(identity.ipv4Address)" + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + nil + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + events.append(.unregister(identity.ipv4Address, ownershipToken)) + } +} + +private actor AmbiguousDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register(String) + case recover(String) + case unregister(String, String?) + } + + private(set) var events: [Event] = [] + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register(identity.dnsName)) + throw DesktopHostRegistrationResultUncertainError(message: "response lost") + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + events.append(.recover(identity.dnsName)) + return "recovered-token" + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + events.append(.unregister(identity.dnsName, ownershipToken)) + } +} + +private actor NegativeRecoveryDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register + case recover + case unregister + } + + private(set) var activeEndpoint: String? + private(set) var events: [Event] = [] + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register) + throw DesktopHostRegistrationResultUncertainError(message: "response lost") + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + events.append(.recover) + return nil + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + #expect(ownershipToken == nil) + events.append(.unregister) + activeEndpoint = nil + } + + func publishNewerEndpoint() { + activeEndpoint = "newer-publisher" + } +} + +private actor RecoverableAmbiguousDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register(String) + case recover(String) + case unregister(String) + } + + private var recoverFailures: Int + private(set) var events: [Event] = [] + + init(recoverFailures: Int = 0) { + self.recoverFailures = recoverFailures + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register(publicationID)) + throw DesktopHostRegistrationResultUncertainError(message: "response lost") + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + events.append(.recover(publicationID)) + if recoverFailures > 0 { + recoverFailures -= 1 + throw DesktopRegistrationTestError.failed + } + return "recovered:\(publicationID)" + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + events.append(.unregister(ownershipToken ?? "")) + } +} + +@MainActor +private final class ToggleDesktopRegistrationStateStore: + DesktopHostRegistrationStateStoring +{ + var failsWrites = false + private var dataByScope: [DesktopHostRegistrationRecoveryScope: Data] = [:] + var data: Data? { dataByScope.values.first } + + func containsState() -> Bool { + !dataByScope.isEmpty + } + + func load(scope: DesktopHostRegistrationRecoveryScope) throws -> Data? { + dataByScope[scope] + } + + func save(_ data: Data?, scope: DesktopHostRegistrationRecoveryScope) throws { + if failsWrites { + throw DesktopRegistrationTestError.failed + } + dataByScope[scope] = data + } + + func data(for scope: DesktopHostRegistrationRecoveryScope) -> Data? { + dataByScope[scope] + } +} + +private actor IdentityAwareAmbiguousDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register(String, UInt16, String) + case recover(String, String) + case unregister(String, String?) + } + + private let uncertainPublicationIDs: Set + private(set) var events: [Event] = [] + + init(uncertainPublicationIDs: Set) { + self.uncertainPublicationIDs = uncertainPublicationIDs + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register(identity.ipv4Address, port, publicationID)) + if uncertainPublicationIDs.contains(publicationID) { + throw DesktopHostRegistrationResultUncertainError(message: "response lost") + } + return "token:\(publicationID)" + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + events.append(.recover(identity.ipv4Address, publicationID)) + return "recovered:\(publicationID)" + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + events.append(.unregister(identity.ipv4Address, ownershipToken)) + } +} + +private actor TwoProcessDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register(String) + case recover(String) + case unregister(String) + } + + private let lostPublicationID: String + private var activeOwnershipToken: String? + private(set) var activePublicationID: String? + private(set) var events: [Event] = [] + + init(lostPublicationID: String) { + self.lostPublicationID = lostPublicationID + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register(publicationID)) + let ownershipToken = "token:\(publicationID)" + activePublicationID = publicationID + activeOwnershipToken = ownershipToken + if publicationID == lostPublicationID { + throw DesktopHostRegistrationResultUncertainError(message: "response lost") + } + return ownershipToken + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + events.append(.recover(publicationID)) + guard activePublicationID == publicationID else { return nil } + return activeOwnershipToken + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + guard ownershipToken == activeOwnershipToken else { return } + events.append(.unregister(ownershipToken ?? "")) + activePublicationID = nil + activeOwnershipToken = nil + } +} + +private actor RetainedLegacyDesktopRegistration: DesktopHostRegistering { + enum Event: Equatable { + case register + case unregister + } + + private var failNextUnregister = true + private(set) var activeEndpoint: String? + private(set) var events: [Event] = [] + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + events.append(.register) + activeEndpoint = "legacy-publisher" + return nil + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + nil + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + #expect(ownershipToken == nil) + events.append(.unregister) + if failNextUnregister { + failNextUnregister = false + throw DesktopRegistrationTestError.failed + } + activeEndpoint = nil + } + + func publishNewerEndpoint() { + activeEndpoint = "newer-publisher" + } +} + +private actor SuspendedDesktopCleanupRegistration: DesktopHostRegistering { + private var unregistrationContinuation: CheckedContinuation? + + var hasStartedUnregistration: Bool { + unregistrationContinuation != nil + } + + func register( + identity: TailnetIdentity, + port: UInt16, + publicationID: String + ) async throws -> String? { + "slow-cleanup-token" + } + + func recover(identity: TailnetIdentity, publicationID: String) async throws -> String? { + nil + } + + func unregister(identity: TailnetIdentity, ownershipToken: String?) async throws { + #expect(ownershipToken == "slow-cleanup-token") + await withCheckedContinuation { continuation in + unregistrationContinuation = continuation + } + } + + func finishUnregistration() { + unregistrationContinuation?.resume() + unregistrationContinuation = nil + } +} + +private final class RemoteInputRecorder: RemoteInputForwarding, @unchecked Sendable { + private let lock = NSLock() + private var releases = 0 + + var releaseCount: Int { + lock.lock() + defer { lock.unlock() } + return releases + } + func keyEvent(down: Bool, keysym: UInt32) {} func pointerEvent(buttonMask: UInt8, x: UInt16, y: UInt16) {} + + func releaseAllInput() { + lock.lock() + releases += 1 + lock.unlock() + } } -private struct LoopbackPeerAuthorizer: TailnetPeerAuthorizing { - func authorize(remoteAddress: String) async -> Bool { - remoteAddress == "127.0.0.1" +private final class BlockingRemoteInputRecorder: RemoteInputForwarding, @unchecked Sendable { + enum Event: Equatable { + case key(down: Bool, keysym: UInt32) + case pointer + case release + } + + private let lock = NSLock() + private let keyEntered = DispatchSemaphore(value: 0) + private let keyMayReturn = DispatchSemaphore(value: 0) + private var storage: [Event] = [] + + var events: [Event] { + lock.lock() + defer { lock.unlock() } + return storage + } + + func keyEvent(down: Bool, keysym: UInt32) { + lock.lock() + storage.append(.key(down: down, keysym: keysym)) + lock.unlock() + keyEntered.signal() + keyMayReturn.wait() + } + + func pointerEvent(buttonMask: UInt8, x: UInt16, y: UInt16) { + lock.lock() + storage.append(.pointer) + lock.unlock() + } + + func releaseAllInput() { + lock.lock() + storage.append(.release) + lock.unlock() } + + func waitForKeyEntry() -> Bool { + keyEntered.wait(timeout: .now() + 1) == .success + } + + func allowKeyReturn() { + keyMayReturn.signal() + } +} + +private final class RFBEventRecorder: @unchecked Sendable { + private let lock = NSLock() + private var storage: [TailnetRFBServerEvent] = [] + + var values: [TailnetRFBServerEvent] { + lock.lock() + defer { lock.unlock() } + return storage + } + + func append(_ event: TailnetRFBServerEvent) { + lock.lock() + storage.append(event) + lock.unlock() + } +} + +private enum RecordedRemoteInputEvent: Equatable { + case key(down: Bool, keysym: UInt32) + case mouse(type: CGEventType, button: CGMouseButton) +} + +private final class RemoteInputEventRecorder: @unchecked Sendable { + private let lock = NSLock() + private var events: [RecordedRemoteInputEvent] = [] + + func append(_ event: RecordedRemoteInputEvent) { + lock.lock() + events.append(event) + lock.unlock() + } + + func contains(_ event: RecordedRemoteInputEvent) -> Bool { + lock.lock() + defer { lock.unlock() } + return events.contains(event) + } +} + +private final class AccessibilityTrust: @unchecked Sendable { + private let lock = NSLock() + private var granted: Bool + private var checks = 0 + + init(granted: Bool) { + self.granted = granted + } + + var checkCount: Int { + lock.lock() + defer { lock.unlock() } + return checks + } + + func isGranted() -> Bool { + lock.lock() + defer { lock.unlock() } + checks += 1 + return granted + } + + func setGranted(_ granted: Bool) { + lock.lock() + self.granted = granted + lock.unlock() + } +} + +private final class DesktopRegistrationTransport: HTTPDataTransport { + private let handler: (URLRequest) throws -> (Data, HTTPURLResponse) + + init(handler: @escaping (URLRequest) throws -> (Data, HTTPURLResponse)) { + self.handler = handler + } + + func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse) { + try handler(request) + } + + func close() {} +} + +private actor DefinitiveRegistrationFailureTransport: HTTPDataTransport { + enum Failure { + case httpStatus(Int) + case redirect + } + + let failure: Failure + private(set) var publicationIDs: [String] = [] + + init(failure: Failure) { + self.failure = failure + } + + func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse) { + let requestURL = try #require(request.url) + let publicationID = try #require( + request.value(forHTTPHeaderField: CrabfleetDesktopRegistration.publicationIDHeader) + ) + publicationIDs.append(publicationID) + let responseURL: URL + let statusCode: Int + switch failure { + case .httpStatus(let status): + responseURL = requestURL + statusCode = status + case .redirect: + responseURL = try #require(URL(string: "https://login.example.test/desktop-host")) + statusCode = 200 + } + return ( + Data(), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: statusCode, + httpVersion: nil, + headerFields: nil + )) + ) + } + + nonisolated func close() {} +} + +private actor LegacyDesktopServerTransport: HTTPDataTransport { + enum Event: Equatable { + case register + case recover + case unregister + } + + private(set) var activeEndpoint: String? + private(set) var events: [Event] = [] + + func data(for request: URLRequest) async throws -> (Data, HTTPURLResponse) { + let responseURL = try #require(request.url) + switch request.httpMethod { + case "PUT": + events.append(.register) + activeEndpoint = "legacy-publisher" + throw URLError(.networkConnectionLost) + case "POST": + events.append(.recover) + return ( + Data(), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 404, + httpVersion: nil, + headerFields: nil + )) + ) + case "DELETE": + events.append(.unregister) + activeEndpoint = nil + return ( + Data(), + try #require( + HTTPURLResponse( + url: responseURL, + statusCode: 200, + httpVersion: nil, + headerFields: nil + )) + ) + default: + Issue.record("unexpected legacy desktop request method") + throw URLError(.badURL) + } + } + + func publishNewerEndpoint() { + activeEndpoint = "newer-publisher" + } + + nonisolated func close() {} +} + +private func waitUntilAsync( + timeout: Duration = .seconds(2), + condition: @escaping () async -> Bool +) async -> Bool { + let clock = ContinuousClock() + let deadline = clock.now.advanced(by: timeout) + while clock.now < deadline { + if await condition() { return true } + try? await Task.sleep(for: .milliseconds(10)) + } + return await condition() +} + +private func desktopRecoveryScope( + origin: String = "https://fleet.example", + ownerSubject: String = "github:123" +) -> DesktopHostRegistrationRecoveryScope { + DesktopHostRegistrationRecoveryScope( + apiOrigin: origin, + ownerSubject: ownerSubject + ) } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Compression/ZlibStream.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Compression/ZlibStream.swift index a0797cbd..cffbba17 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Compression/ZlibStream.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Compression/ZlibStream.swift @@ -38,6 +38,12 @@ extension ZlibStream { func decompressedData(compressedData: Data, maximumOutputSize: Int) throws -> Data { + guard maximumOutputSize >= 0 else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataOverflow + )) + } + let stream = self.stream let flush = ZlibFlush.noFlush @@ -64,47 +70,55 @@ extension ZlibStream { stream.availIn = .init(compressedSize) while true { -// print("AVAILIN: \(stream.availIn)") - - if stream.availIn <= 0 { - break - } - - var isDone = false - stream.nextOut = buffer stream.availOut = .init(bufferSize) -// print("AVAIL OUT BEFORE INFLATE: \(stream.availOut)") + let inputBefore = stream.availIn + let outputBefore = stream.totalOut + let isDone: Bool - // Inflate another chunk. do { isDone = try stream.inflate(flush: flush) - - if stream.availOut >= 0 { - let availOut: UInt = .init(stream.availOut) - - // print("AVAIL OUT AFTER INFLATE: \(availOut)") - - let actualOut = bufferSize - availOut - - if actualOut > 0 { - guard decompressedData.count <= maximumOutputSize - Int(actualOut) else { - throw VNCError.protocol(.zlibDecompress( - underlyingError: ZlibStreamError.decompressedDataOverflow - )) - } - decompressedData.append(buffer, - count: .init(actualOut)) - } + } catch let error as ZlibError { + if case .bufferError = error, inputBefore == 0 { + break } + throw VNCError.protocol(.zlibDecompress(underlyingError: error)) } catch { throw VNCError.protocol(.zlibDecompress(underlyingError: error)) } + let actualOut = bufferSize - UInt(stream.availOut) + + if actualOut > 0 { + let actualOutCount = Int(actualOut) + guard decompressedData.count <= maximumOutputSize, + actualOutCount <= maximumOutputSize - decompressedData.count else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataOverflow + )) + } + decompressedData.append(buffer, count: actualOutCount) + } + if isDone { + guard stream.availIn == 0 else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataLengthMismatch + )) + } + break + } + + if stream.availIn == 0, stream.availOut > 0 { break } + + guard stream.availIn < inputBefore || stream.totalOut > outputBefore else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataLengthMismatch + )) + } } } @@ -135,6 +149,8 @@ extension ZlibStream { throw VNCError.protocol(.zlibDecompress(underlyingError: nil)) } + var overflowByte: UInt8 = 0 + while true { let doneBytes = stream.totalOut let remainingBytes = uncompressedSize - doneBytes @@ -143,26 +159,59 @@ extension ZlibStream { throw VNCError.protocol(.zlibDecompress(underlyingError: ZlibStreamError.decompressedDataOverflow)) } - stream.nextOut = decompressedDataBytes.advanced(by: .init(doneBytes)) - stream.availOut = .init(remainingBytes) + let inputBefore = stream.availIn + let outputBefore = stream.totalOut + let isDone: Bool - if remainingBytes <= 0 { - break - } - - // Inflate another chunk. do { - let isDone = try stream.inflate(flush: flush) - - if isDone { + if remainingBytes > 0 { + stream.nextOut = decompressedDataBytes.advanced(by: .init(doneBytes)) + stream.availOut = .init(remainingBytes) + isDone = try stream.inflate(flush: flush) + } else { + isDone = try withUnsafeMutablePointer(to: &overflowByte) { overflowPtr in + stream.nextOut = overflowPtr + stream.availOut = 1 + return try stream.inflate(flush: flush) + } + } + } catch let error as ZlibError { + if case .bufferError = error, inputBefore == 0 { break } + throw VNCError.protocol(.zlibDecompress(underlyingError: error)) } catch { throw VNCError.protocol(.zlibDecompress(underlyingError: error)) } + + guard stream.totalOut <= uncompressedSize else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataOverflow + )) + } + + if isDone { + guard stream.availIn == 0 else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataLengthMismatch + )) + } + break + } + + if stream.availIn == 0, stream.availOut > 0 { + break + } + + guard stream.availIn < inputBefore || stream.totalOut > outputBefore else { + throw VNCError.protocol(.zlibDecompress( + underlyingError: ZlibStreamError.decompressedDataLengthMismatch + )) + } } - guard stream.totalOut == uncompressedSize else { + guard stream.totalOut == uncompressedSize, + stream.availIn == 0 else { throw VNCError.protocol(.zlibDecompress(underlyingError: ZlibStreamError.decompressedDataLengthMismatch)) } } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Encryption/BigNum.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Encryption/BigNum.swift index 56ced73e..829e5f31 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Encryption/BigNum.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Encryption/BigNum.swift @@ -20,6 +20,15 @@ final class BigNum { } extension BigNum { + var isGreaterThanOne: Bool { + self.bigInt > 1 + } + + func isValidDiffieHellmanElement(modulus: BigNum) -> Bool { + guard modulus.bigInt > 2 else { return false } + return self.bigInt > 1 && self.bigInt < modulus.bigInt - 1 + } + var isZero: Bool { let isIt = self.bigInt == 0 diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Framebuffer/VNCFramebuffer.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Framebuffer/VNCFramebuffer.swift index 63c0689b..086a12dc 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Framebuffer/VNCFramebuffer.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Framebuffer/VNCFramebuffer.swift @@ -731,10 +731,34 @@ private extension VNCFramebuffer { return } - var data = bufferData(ofRegion: sourceRegion) + let data = bufferData(ofRegion: sourceRegion) + let bytesPerPixel = destinationProperties.bytesPerPixel + let rowByteCount = Int(destinationRegion.width) * bytesPerPixel + let destinationX = Int(destinationRegion.x) + let destinationY = Int(destinationRegion.y) + + guard data.count == rowByteCount * Int(destinationRegion.height) else { + logger.logError("Invalid internal framebuffer data length for CopyRect") + return + } + + data.withUnsafeBytes { sourceBytes in + guard let sourceBase = sourceBytes.baseAddress else { return } + + for row in 0.. Data { diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Network/NWConnection+NetworkConnection.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Network/NWConnection+NetworkConnection.swift index 9ad2b099..1e54c744 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Network/NWConnection+NetworkConnection.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Network/NWConnection+NetworkConnection.swift @@ -76,36 +76,45 @@ extension NWConnection: NetworkConnectionReading { maximumLength: Int) async throws -> Data { return try await withCheckedThrowingContinuation { continuation in receive(minimumIncompleteLength: minimumLength, maximumLength: maximumLength) { content, _, isComplete, error in - guard !isComplete else { - continuation.resume(throwing: VNCError.connection(.closed)) - - return - } - - guard error == nil else { - continuation.resume(throwing: error!) - - return - } - - guard let content else { - continuation.resume(throwing: VNCError.protocol(.noData)) - - return + do { + continuation.resume(returning: try Self.validateReadContent( + content, + isComplete: isComplete, + error: error, + minimumLength: minimumLength, + maximumLength: maximumLength + )) + } catch { + continuation.resume(throwing: error) } + } + } + } - let receivedLength = content.count - - guard receivedLength >= minimumLength, - receivedLength <= maximumLength else { - continuation.resume(throwing: VNCError.protocol(.invalidData)) - - return - } + static func validateReadContent( + _ content: Data?, + isComplete: Bool, + error: Error?, + minimumLength: Int, + maximumLength: Int + ) throws -> Data { + if let error { + throw error + } - continuation.resume(returning: content) + if let content, !content.isEmpty { + guard content.count >= minimumLength, + content.count <= maximumLength else { + throw VNCError.protocol(.invalidData) } + return content + } + + if isComplete { + throw VNCError.connection(.closed) } + + throw VNCError.protocol(.noData) } } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/TightEncoding.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/TightEncoding.swift index 4fa2c14b..e98645b4 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/TightEncoding.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/TightEncoding.swift @@ -333,7 +333,7 @@ private extension VNCProtocol.TightEncoding { case gradient = 2 } - func resetZStreamsIfNeeded(control: UInt8, + func resetZStreamsIfNeeded(control: UInt8, logger: VNCLogger) { for idx in 0..<4 { let mask = UInt8(1 << idx) @@ -349,28 +349,30 @@ private extension VNCProtocol.TightEncoding { } } } +} +extension VNCProtocol.TightEncoding { func readCompactLength(connection: NetworkConnectionReading, logger: VNCLogger) async throws -> Int { - var length = 0 - var shift = 0 - - for _ in 0..<3 { -// logger.logDebug("Reading Tight Compact Length") - - let byte = try await connection.readUInt8() - length |= Int(byte & 0x7F) << shift - - if (byte & 0x80) == 0 { - return length - } + let first = try await connection.readUInt8() + var length = Int(first & 0x7F) + guard (first & 0x80) != 0 else { + return length + } - shift += 7 + let second = try await connection.readUInt8() + length |= Int(second & 0x7F) << 7 + guard (second & 0x80) != 0 else { + return length } - throw VNCError.protocol(.invalidData) + let third = try await connection.readUInt8() + length |= Int(third) << 14 + return length } +} +private extension VNCProtocol.TightEncoding { func readBuffered(connection: NetworkConnectionReading, length: Int, logger: VNCLogger) async throws -> Data { diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/ZRLEEncoding.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/ZRLEEncoding.swift index f291d7ed..1fd80e43 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/ZRLEEncoding.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Encodings/Frame/ZRLEEncoding.swift @@ -68,7 +68,10 @@ extension VNCProtocol.ZRLEEncoding { let decompressedData = try zStream.decompressedData( compressedData: compressedData, - maximumOutputSize: VNCProtocolLimits.maximumFramebufferBytes + maximumOutputSize: Self.maximumInflatedSize( + width: Int(rectangle.width), + height: Int(rectangle.height) + ) ) let stream = DataStream(data: decompressedData) @@ -156,6 +159,35 @@ extension VNCProtocol.ZRLEEncoding { framebuffer.didUpdate(region: rectangle.region) } + + static func maximumInflatedSize(width: Int, height: Int) -> Int { + guard width > 0, height > 0 else { return 0 } + + let tileSize = Int(Self.tileSize) + var maximumSize = 0 + + for tileY in stride(from: 0, to: height, by: tileSize) { + let tileHeight = min(tileSize, height - tileY) + + for tileX in stride(from: 0, to: width, by: tileSize) { + let tileWidth = min(tileSize, width - tileX) + let pixelCount = tileWidth * tileHeight + let packedPaletteBytes = 16 * 3 + ((tileWidth * 4 + 7) / 8) * tileHeight + let rlePaletteBytes = 127 * 3 + pixelCount * 2 + let tilePayloadBytes = max( + pixelCount * 3, + 3, + packedPaletteBytes, + pixelCount * 4, + rlePaletteBytes + ) + + maximumSize += 1 + tilePayloadBytes + } + } + + return maximumSize + } } private extension VNCProtocol.ZRLEEncoding { @@ -234,6 +266,12 @@ private extension VNCProtocol.ZRLEEncoding { } let indexInPalette = (Int(encoded) >> shift) & mask + guard indexInPalette < Int(paletteSize) else { + throw VNCError.protocol(.zrlePaletteIndexOverflow( + paletteIndex: indexInPalette, + paletteSize: paletteSize + )) + } let sourceStartIndex = indexInPalette * 4 diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Messages/ClientToServer/ClientFence.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Messages/ClientToServer/ClientFence.swift new file mode 100644 index 00000000..4e158f5c --- /dev/null +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Messages/ClientToServer/ClientFence.swift @@ -0,0 +1,45 @@ +#if canImport(FoundationEssentials) +import FoundationEssentials +#else +import Foundation +#endif + +extension VNCProtocol { + struct FenceFlags: OptionSet { + let rawValue: UInt32 + + static let blockBefore = Self(rawValue: 1 << 0) + static let blockAfter = Self(rawValue: 1 << 1) + static let syncNext = Self(rawValue: 1 << 2) + static let request = Self(rawValue: 1 << 31) + + static let supported: Self = [.blockBefore, .blockAfter, .syncNext, .request] + } + + struct ClientFence: VNCSendableMessage { + static let maximumPayloadLength = 64 + static let messageType: UInt8 = 248 + + var messageType: UInt8 { Self.messageType } + let flags: FenceFlags + let payload: Data + } +} + +extension VNCProtocol.ClientFence { + var data: Data { + precondition(payload.count <= Self.maximumPayloadLength) + + var data = Data(capacity: 9 + payload.count) + data.append(messageType) + data.appendPadding(length: 3) + data.append(flags.rawValue, bigEndian: true) + data.append(UInt8(payload.count)) + data.append(payload) + return data + } + + func send(connection: NetworkConnectionWriting) async throws { + try await connection.write(data: data) + } +} diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Messages/ServerToClient/ServerFence.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Messages/ServerToClient/ServerFence.swift new file mode 100644 index 00000000..e369d93b --- /dev/null +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/Messages/ServerToClient/ServerFence.swift @@ -0,0 +1,28 @@ +#if canImport(FoundationEssentials) +import FoundationEssentials +#else +import Foundation +#endif + +extension VNCProtocol { + struct ServerFence: VNCReceivableMessage { + static let messageType: UInt8 = 248 + + let messageType: UInt8 + let flags: FenceFlags + let payload: Data + } +} + +extension VNCProtocol.ServerFence { + static func receive(connection: NetworkConnectionReading) async throws -> Self { + try await connection.readPadding(length: 3) + let flags = VNCProtocol.FenceFlags(rawValue: try await connection.readUInt32()) + let payloadLength = Int(try await connection.readUInt8()) + guard payloadLength <= VNCProtocol.ClientFence.maximumPayloadLength else { + throw VNCError.protocol(.invalidData) + } + let payload = try await connection.read(length: payloadLength) + return .init(messageType: messageType, flags: flags, payload: payload) + } +} diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/AppleRemoteDesktop/ARDDiffieHellmanKeyAgreement.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/AppleRemoteDesktop/ARDDiffieHellmanKeyAgreement.swift index 0a23e98b..7b7a5be7 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/AppleRemoteDesktop/ARDDiffieHellmanKeyAgreement.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/AppleRemoteDesktop/ARDDiffieHellmanKeyAgreement.swift @@ -3,9 +3,38 @@ import FoundationEssentials #else import Foundation #endif +import CryptoSwift extension VNCProtocol.ARDAuthentication { struct DiffieHellmanKeyAgreement { + struct ValidatedSafePrimeCache { + let capacity: Int + private var insertionOrder = [Data]() + private var values = Set() + + init(capacity: Int) { + precondition(capacity > 0) + self.capacity = capacity + } + + var count: Int { + values.count + } + + func contains(_ value: Data) -> Bool { + values.contains(value) + } + + mutating func insert(_ value: Data) { + guard values.insert(value).inserted else { return } + + insertionOrder.append(value) + if values.count > capacity { + values.remove(insertionOrder.removeFirst()) + } + } + } + let publicKey: Data let privateKey: Data let secretKey: Data @@ -19,7 +48,10 @@ extension VNCProtocol.ARDAuthentication { generator: Data, peerKey: Data, keyLength: Int) { - guard keyLength > 0 else { + guard Self.validParameters(prime: prime, + generator: generator, + peerKey: peerKey, + keyLength: keyLength) else { return nil } @@ -47,11 +79,71 @@ extension VNCProtocol.ARDAuthentication { } private extension VNCProtocol.ARDAuthentication.DiffieHellmanKeyAgreement { + static let safePrimeCondition = NSCondition() + static var validatedSafePrimes = ValidatedSafePrimeCache(capacity: 32) + static var rejectedSafePrimes = Set() + static var safePrimeValidations = Set() + struct KeyPair { let publicKey: Data let privateKey: Data } + static func validParameters(prime: Data, + generator: Data, + peerKey: Data, + keyLength: Int) -> Bool { + guard keyLength >= 128, + keyLength <= 512, + prime.count == keyLength, + peerKey.count == keyLength, + let bigPrime = BigNum(data: prime), + bigPrime.bitsCount == keyLength * 8, + let bigGenerator = BigNum(data: generator), + bigGenerator.isValidDiffieHellmanElement(modulus: bigPrime), + let bigPeerKey = BigNum(data: peerKey), + bigPeerKey.isValidDiffieHellmanElement(modulus: bigPrime), + Self.isSafePrime(prime) else { + return false + } + + return true + } + + static func isSafePrime(_ data: Data) -> Bool { + safePrimeCondition.lock() + while safePrimeValidations.contains(data) { + safePrimeCondition.wait() + } + if validatedSafePrimes.contains(data) { + safePrimeCondition.unlock() + return true + } + if rejectedSafePrimes.contains(data) { + safePrimeCondition.unlock() + return false + } + safePrimeValidations.insert(data) + safePrimeCondition.unlock() + + let prime = CS.BigUInt(data) + let valid = prime.isPrime(rounds: 16) && ((prime - 1) >> 1).isPrime(rounds: 16) + + safePrimeCondition.lock() + safePrimeValidations.remove(data) + if valid { + validatedSafePrimes.insert(data) + } else { + if rejectedSafePrimes.count >= 32, let evicted = rejectedSafePrimes.first { + rejectedSafePrimes.remove(evicted) + } + rejectedSafePrimes.insert(data) + } + safePrimeCondition.broadcast() + safePrimeCondition.unlock() + return valid + } + static func generateKeyPair(generator: Data, prime: Data, keyLength: Int) -> KeyPair? { @@ -59,6 +151,7 @@ private extension VNCProtocol.ARDAuthentication.DiffieHellmanKeyAgreement { let bigPubKey = BigNum() guard let bigPrime = BigNum(data: prime), + bigPrime.isGreaterThanOne, let bigGenerator = BigNum(data: generator) else { return nil } @@ -77,7 +170,7 @@ private extension VNCProtocol.ARDAuthentication.DiffieHellmanKeyAgreement { x: bigPrivKey, p: bigPrime) - guard modSuccess else { + guard modSuccess, !bigPubKey.isZero else { return nil } @@ -115,7 +208,7 @@ private extension VNCProtocol.ARDAuthentication.DiffieHellmanKeyAgreement { x: bigPrivKey, p: bigPrime) - guard modSuccess else { + guard modSuccess, !bigSharedKey.isZero else { return nil } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCBigNum.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCBigNum.swift index ac31c275..89c024d7 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCBigNum.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCBigNum.swift @@ -40,57 +40,61 @@ extension VNCProtocol.UltraVNCMSLogonIIAuthentication.DiffieHellmanKeyAgreement static func addM64(x: UInt64, y: UInt64, m: UInt64) -> UInt64 { - let part = Int64(x + y < x - ? (-1 % .init(m) + 1) % .init(m) - : 0) + guard m != 0 else { return 0 } - let partU: UInt64 = numericCast(part) + let reducedX = x % m + let reducedY = y % m + let distanceToModulus = m - reducedY - let result: UInt64 = (x + y) % m + partU + if reducedX >= distanceToModulus { + return reducedX - distanceToModulus + } - return result + return reducedX + reducedY } /// (x * y) % m static func mulM64(x: UInt64, y: UInt64, m: UInt64) -> UInt64 { - var y = y - var r = UInt64(0) - var x = UInt64(0) + guard m != 0 else { return 0 } - repeat { - x>>=1 + var multiplicand = x % m + var multiplier = y % m + var result = UInt64(0) - if x & 1 != 0 { - r = addM64(x: r, y: y, m: m) + while multiplier > 0 { + if multiplier & 1 != 0 { + result = addM64(x: result, y: multiplicand, m: m) } - y = addM64(x: y, y: y, m: m) - } while x > 0 + multiplier >>= 1 + multiplicand = addM64(x: multiplicand, y: multiplicand, m: m) + } - return r + return result } /// (x ^ y) % m static func powM64(b: UInt64, e: UInt64, m: UInt64) -> UInt64 { - var b = b - var r = UInt64(0) - var e = UInt64(0) + guard m != 0 else { return 0 } - repeat { - e>>=1 + var base = b % m + var exponent = e + var result = UInt64(1) % m - if e & 1 != 0 { - r = mulM64(x: r, y: b, m: m) + while exponent > 0 { + if exponent & 1 != 0 { + result = mulM64(x: result, y: base, m: m) } - b = mulM64(x: b, y: b, m: m) - } while e > 0 + exponent >>= 1 + base = mulM64(x: base, y: base, m: m) + } - return r + return result } } } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCMSLogonIIDiffieHellmanKeyAgreement.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCMSLogonIIDiffieHellmanKeyAgreement.swift index e068c791..79c6716a 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCMSLogonIIDiffieHellmanKeyAgreement.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/Protocol/SecurityTypes/UltraVNCMSLogonII/UltraVNCMSLogonIIDiffieHellmanKeyAgreement.swift @@ -44,13 +44,15 @@ private extension VNCProtocol.UltraVNCMSLogonIIAuthentication.DiffieHellmanKeyAg static func generateKeyPair(generator: Data, modulus: Data) -> KeyPair? { let generatorNum = UltraVNCBigNum.dataToBigNum(generator) - guard generatorNum < maxNum else { return nil } - let modulusNum = UltraVNCBigNum.dataToBigNum(modulus) - guard modulusNum < maxNum else { return nil } + guard modulusNum > 3, + modulusNum < maxNum, + generatorNum > 1, + generatorNum <= modulusNum - 2 else { + return nil + } - let privNum = UltraVNCBigNum.randomBigNum(max: .init(maxNum)) - guard privNum < maxNum else { return nil } + let privNum = UInt64.random(in: 2..<(modulusNum - 1)) let privData = UltraVNCBigNum.bigNumToData(privNum) @@ -73,7 +75,14 @@ private extension VNCProtocol.UltraVNCMSLogonIIAuthentication.DiffieHellmanKeyAg let modulusNum = UltraVNCBigNum.dataToBigNum(modulus) let respNum = UltraVNCBigNum.dataToBigNum(resp) - guard respNum < maxNum else { return nil } + guard modulusNum > 3, + modulusNum < maxNum, + privNum > 1, + privNum < modulusNum, + respNum > 1, + respNum <= modulusNum - 2 else { + return nil + } let keyNum = UltraVNCBigNum.powM64(b: .init(respNum), e: .init(privNum), diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/State.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/State.swift index 1b08e904..423c11bf 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/State.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/State.swift @@ -18,6 +18,8 @@ extension VNCConnection { private var _pixelFormat: VNCProtocol.PixelFormat? private var _desktopName: String? private var _incrementalUpdatesEnabled = false + private var _areFencesSupported = false + private var _pixelFormatTransitionFenceFlags: VNCProtocol.FenceFlags = [] private var _areContinuousUpdatesSupported = false private var _areContinuousUpdatesEnabled = false private var _extendedClipboardServerCaps: VNCExtendedClipboardCaps? @@ -70,6 +72,16 @@ extension VNCConnection { set { withLock { _incrementalUpdatesEnabled = newValue } } } + var areFencesSupported: Bool { + get { withLock { _areFencesSupported } } + set { withLock { _areFencesSupported = newValue } } + } + + var pixelFormatTransitionFenceFlags: VNCProtocol.FenceFlags { + get { withLock { _pixelFormatTransitionFenceFlags } } + set { withLock { _pixelFormatTransitionFenceFlags = newValue } } + } + var areContinuousUpdatesSupported: Bool { get { withLock { _areContinuousUpdatesSupported } } set { withLock { _areContinuousUpdatesSupported = newValue } } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+API.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+API.swift index 22d3be7f..5a25e3a4 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+API.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+API.swift @@ -4,6 +4,79 @@ import FoundationEssentials import Foundation #endif +private struct PixelFormatTransitionMessage: VNCSendableMessage { + let synchronizationFenceMessage: VNCProtocol.ClientFence? + let pixelFormatMessage: VNCProtocol.SetPixelFormat + let encodingsMessage: VNCProtocol.SetEncodings + let completionFenceMessage: VNCProtocol.ClientFence? + let willSend: () throws -> Void + let didSend: () -> Void + let willSendFence: () -> Void + let didSendFence: () -> Void + let didFailFence: () -> Void + + var messageType: UInt8 { + synchronizationFenceMessage?.messageType ?? pixelFormatMessage.messageType + } + var data: Data { + (synchronizationFenceMessage?.data ?? Data()) + + pixelFormatMessage.data + + encodingsMessage.data + + (completionFenceMessage?.data ?? Data()) + } + + func send(connection: NetworkConnectionWriting) async throws { + if synchronizationFenceMessage == nil { + try willSend() + } else { + willSendFence() + } + do { + try await connection.write(data: data) + } catch { + if synchronizationFenceMessage != nil { + didFailFence() + } + throw error + } + if synchronizationFenceMessage == nil { + didSend() + } else { + didSendFence() + } + } +} + +private struct PixelFormatTransition { + let pixelFormat: VNCProtocol.PixelFormat + let fenceFlags: VNCProtocol.FenceFlags + let synchronizationFencePayload: Data? + let completionFencePayload: Data? + let sequence: UInt64 +} + +private struct FenceCapabilityProbeMessage: VNCSendableMessage { + let fenceMessage: VNCProtocol.ClientFence + let pixelFormatMessage: VNCProtocol.SetPixelFormat + let completionFenceMessage: VNCProtocol.ClientFence + let didSend: () -> Void + + var messageType: UInt8 { fenceMessage.messageType } + var data: Data { + fenceMessage.data + pixelFormatMessage.data + completionFenceMessage.data + } + + func send(connection: NetworkConnectionWriting) async throws { + try await connection.write(data: data) + didSend() + } +} + +private enum PixelFormatFenceProbePayload { + static let capability = Data("royalvnc-pixel-format".utf8) + static let completion = Data("royalvnc-pixel-format-complete".utf8) +} + // MARK: - Connect/Disconnect public extension VNCConnection { #if canImport(ObjectiveC) @@ -26,19 +99,15 @@ public extension VNCConnection { @objc #endif func updateColorDepth(_ colorDepth: Settings.ColorDepth) { - guard let framebuffer = framebuffer else { return } - - let newPixelFormat = VNCProtocol.PixelFormat(depth: colorDepth.rawValue) - - state.pixelFormat = newPixelFormat - - let sendPixelFormatMessage = VNCProtocol.SetPixelFormat(pixelFormat: newPixelFormat) - - enqueueClientToServerMessage(sendPixelFormatMessage) + withLifecycleLock { + guard connectionState.status == .connected, + framebuffer != nil else { + return + } - recreateFramebuffer(size: framebuffer.size, - screens: framebuffer.screens, - pixelFormat: newPixelFormat) + let newPixelFormat = VNCProtocol.PixelFormat(depth: colorDepth.rawValue) + requestPixelFormatTransition(newPixelFormat) + } } /// Requests a single-screen desktop matching the viewer viewport. @@ -73,6 +142,424 @@ public extension VNCConnection { } } +extension VNCConnection { + private func requestPixelFormatTransition(_ pixelFormat: VNCProtocol.PixelFormat) { + framebufferRequestLock.lock() + pendingPixelFormatTransition = pixelFormat + framebufferRequestGeneration &+= 1 + framebufferPacingTask?.cancel() + framebufferPacingTask = nil + let transition = takePendingPixelFormatTransitionLocked() + framebufferRequestLock.unlock() + + if let transition { + enqueuePixelFormatTransition(transition) + } + } + + private func takePendingPixelFormatTransitionLocked() -> PixelFormatTransition? { + guard !isPixelFormatTransitionInFlight, + pixelFormatFenceCapabilityProbePayload == nil, + let pixelFormat = pendingPixelFormatTransition else { + return nil + } + + let supportedFenceFlags = state.pixelFormatTransitionFenceFlags + let supportsSynchronizedBoundary = + supportedFenceFlags.contains(.blockBefore) + && supportedFenceFlags.contains(.syncNext) + var fenceFlags: VNCProtocol.FenceFlags = [] + if state.areContinuousUpdatesEnabled { + guard supportsSynchronizedBoundary else { + if !state.areFencesSupported { + schedulePixelFormatFenceNegotiationTimeoutLocked() + } else { + pendingPixelFormatTransition = nil + } + return nil + } + fenceFlags = [.request, .blockBefore, .syncNext] + } else if framebufferUpdateRequestOutstanding { + if supportsSynchronizedBoundary { + fenceFlags = [.request, .blockBefore, .syncNext] + } else { + if !state.areFencesSupported { + schedulePixelFormatFenceNegotiationTimeoutLocked() + } + return nil + } + } + + cancelPixelFormatFenceNegotiationTimeoutLocked() + pendingPixelFormatTransition = nil + isPixelFormatTransitionInFlight = true + pixelFormatTransitionInFlight = pixelFormat + nextPixelFormatChangeSequence &+= 1 + let sequence = nextPixelFormatChangeSequence + pixelFormatTransitionInFlightSequence = sequence + let synchronizationFencePayload: Data? + let completionFencePayload: Data? + if !fenceFlags.isEmpty { + pixelFormatTransitionFenceSequence &+= 1 + var sequence = pixelFormatTransitionFenceSequence.bigEndian + synchronizationFencePayload = withUnsafeBytes(of: &sequence) { Data($0) } + pixelFormatTransitionFenceSequence &+= 1 + sequence = pixelFormatTransitionFenceSequence.bigEndian + completionFencePayload = withUnsafeBytes(of: &sequence) { Data($0) } + pixelFormatTransitionFencePayload = completionFencePayload + pixelFormatTransitionRequiredFenceFlags = [.blockBefore] + pixelFormatTransitionFenceWasSent = false + } else { + synchronizationFencePayload = nil + completionFencePayload = nil + pixelFormatTransitionRequiredFenceFlags = [] + pixelFormatTransitionFenceWasSent = false + } + return PixelFormatTransition( + pixelFormat: pixelFormat, + fenceFlags: fenceFlags, + synchronizationFencePayload: synchronizationFencePayload, + completionFencePayload: completionFencePayload, + sequence: sequence + ) + } + + private func enqueuePixelFormatTransition(_ transition: PixelFormatTransition) { + let encodingTypes: [VNCEncodingType] + do { + encodingTypes = try orderedEncodingTypes(pixelFormat: transition.pixelFormat) + } catch { + handleBreakingError(error) + return + } + let synchronizationFenceMessage = transition.synchronizationFencePayload.map { + VNCProtocol.ClientFence(flags: transition.fenceFlags, payload: $0) + } + let completionFenceMessage = transition.completionFencePayload.map { + VNCProtocol.ClientFence(flags: [.request, .blockBefore], payload: $0) + } + let message = PixelFormatTransitionMessage( + synchronizationFenceMessage: synchronizationFenceMessage, + pixelFormatMessage: VNCProtocol.SetPixelFormat(pixelFormat: transition.pixelFormat), + encodingsMessage: VNCProtocol.SetEncodings(encodingTypes: encodingTypes), + completionFenceMessage: completionFenceMessage, + willSend: { [weak self] in + try self?.beginPixelFormatTransition( + transition.pixelFormat, + sequence: transition.sequence + ) + }, + didSend: { [weak self] in + self?.completePixelFormatTransition() + }, + willSendFence: { [weak self] in + guard let payload = transition.completionFencePayload else { return } + self?.beginPixelFormatTransitionFenceWrite(payload: payload) + }, + didSendFence: { [weak self] in + guard let payload = transition.completionFencePayload else { return } + self?.schedulePixelFormatTransitionDeadline(payload: payload) + }, + didFailFence: { [weak self] in + guard let payload = transition.completionFencePayload else { return } + self?.cancelPixelFormatTransitionFenceWrite(payload: payload) + } + ) + + enqueueClientToServerMessage(message) + } + + private func beginPixelFormatTransitionFenceWrite(payload: Data) { + framebufferRequestLock.lock() + defer { framebufferRequestLock.unlock() } + guard isPixelFormatTransitionInFlight, + pixelFormatTransitionFencePayload == payload else { + return + } + pixelFormatTransitionFenceWasSent = true + } + + private func schedulePixelFormatTransitionDeadline(payload: Data) { + framebufferRequestLock.lock() + defer { framebufferRequestLock.unlock() } + guard isPixelFormatTransitionInFlight, + pixelFormatTransitionFencePayload == payload, + pixelFormatTransitionFenceWasSent else { + return + } + + schedulePixelFormatTransitionDeadlineIfReadyLocked(payload: payload) + } + + private func cancelPixelFormatTransitionFenceWrite(payload: Data) { + framebufferRequestLock.lock() + defer { framebufferRequestLock.unlock() } + guard pixelFormatTransitionFencePayload == payload else { return } + pixelFormatTransitionFenceWasSent = false + cancelPixelFormatTransitionDeadlineLocked() + } + + private func schedulePixelFormatTransitionDeadlineIfReadyLocked(payload: Data) { + guard pixelFormatTransitionFenceWasSent, + !framebufferUpdateRequestOutstanding, + pixelFormatTransitionDeadlineTask == nil else { + return + } + + pixelFormatTransitionDeadlineTask = Task { [weak self] in + do { + try await Task.sleep(nanoseconds: 5_000_000_000) + } catch { + return + } + self?.expirePixelFormatTransitionDeadline(payload: payload) + } + } + + private func cancelPixelFormatTransitionDeadlineLocked() { + pixelFormatTransitionDeadlineTask?.cancel() + pixelFormatTransitionDeadlineTask = nil + } + + func expirePixelFormatTransitionDeadline(payload: Data) { + framebufferRequestLock.lock() + guard isPixelFormatTransitionInFlight, + pixelFormatTransitionFencePayload == payload, + pixelFormatTransitionDeadlineTask != nil else { + framebufferRequestLock.unlock() + return + } + pixelFormatTransitionDeadlineTask = nil + framebufferRequestLock.unlock() + + handleBreakingError(VNCError.protocol(.pixelFormatTransitionTimedOut)) + } + + func publishPixelFormatFenceSupport() -> Bool { + framebufferRequestLock.lock() + defer { framebufferRequestLock.unlock() } + guard !state.areFencesSupported else { return false } + + cancelPixelFormatFenceNegotiationTimeoutLocked() + if pixelFormatFenceCapabilityProbePayload == nil, + state.pixelFormat != nil { + pixelFormatFenceCapabilityProbePayload = PixelFormatFenceProbePayload.capability + } + state.areFencesSupported = true + return true + } + + func enqueuePixelFormatFenceSupportProbe() { + framebufferRequestLock.lock() + guard let payload = pixelFormatFenceCapabilityProbePayload, + let pixelFormat = state.pixelFormat else { + framebufferRequestLock.unlock() + return + } + nextPixelFormatChangeSequence &+= 1 + pixelFormatFenceCapabilityProbeSequence = nextPixelFormatChangeSequence + framebufferRequestLock.unlock() + + enqueueClientToServerMessage( + FenceCapabilityProbeMessage( + fenceMessage: VNCProtocol.ClientFence( + flags: [.request, .blockBefore, .blockAfter, .syncNext], + payload: payload + ), + pixelFormatMessage: VNCProtocol.SetPixelFormat(pixelFormat: pixelFormat), + completionFenceMessage: VNCProtocol.ClientFence( + flags: [.request, .blockBefore], + payload: PixelFormatFenceProbePayload.completion + ), + didSend: { [weak self] in + self?.didSendPixelFormatFenceCapabilityProbe(payload: payload) + } + ) + ) + } + + private func didSendPixelFormatFenceCapabilityProbe(payload: Data) { + framebufferRequestLock.lock() + defer { framebufferRequestLock.unlock() } + guard pixelFormatFenceCapabilityProbePayload == payload else { return } + schedulePixelFormatFenceNegotiationTimeoutLocked() + } + + private func schedulePixelFormatFenceNegotiationTimeoutLocked() { + guard pixelFormatFenceNegotiationTask == nil else { return } + + pixelFormatFenceNegotiationTask = Task { [weak self] in + do { + try await Task.sleep(nanoseconds: 1_000_000_000) + } catch { + return + } + self?.expirePixelFormatFenceNegotiation() + } + } + + private func cancelPixelFormatFenceNegotiationTimeoutLocked() { + pixelFormatFenceNegotiationTask?.cancel() + pixelFormatFenceNegotiationTask = nil + } + + func expirePixelFormatFenceNegotiation() { + framebufferRequestLock.lock() + cancelPixelFormatFenceNegotiationTimeoutLocked() + let probeTimedOut = pixelFormatFenceCapabilityProbePayload != nil + let negotiationTimedOut = + !state.areFencesSupported + && pendingPixelFormatTransition != nil + && !isPixelFormatTransitionInFlight + guard probeTimedOut || negotiationTimedOut else { + framebufferRequestLock.unlock() + return + } + if probeTimedOut { + expiredPixelFormatFenceCapabilityProbePayload = pixelFormatFenceCapabilityProbePayload + pixelFormatFenceCapabilityProbePayload = nil + framebufferRequestLock.unlock() + handleBreakingError(VNCError.protocol(.pixelFormatTransitionTimedOut)) + logger.logDebug("Fence capability negotiation timed out") + return + } + let isWaitingForLegacyFramebufferBoundary = + negotiationTimedOut + && !state.areContinuousUpdatesEnabled + && framebufferUpdateRequestOutstanding + if negotiationTimedOut && !isWaitingForLegacyFramebufferBoundary { + pendingPixelFormatTransition = nil + } + let shouldResumeUpdates = + pendingPixelFormatTransition == nil + && !framebufferUpdateRequestOutstanding + && !isPixelFormatTransitionInFlight + framebufferRequestLock.unlock() + + if shouldResumeUpdates { + scheduleNextFramebufferUpdate() + } + logger.logDebug("Fence capability negotiation timed out") + } + + func completePixelFormatFence(_ fence: VNCProtocol.ServerFence) throws { + framebufferRequestLock.lock() + if fence.payload == pixelFormatFenceCapabilityProbePayload + || fence.payload == expiredPixelFormatFenceCapabilityProbePayload { + guard let sequence = pixelFormatFenceCapabilityProbeSequence else { + framebufferRequestLock.unlock() + throw VNCError.protocol(.invalidData) + } + cancelPixelFormatFenceNegotiationTimeoutLocked() + if fence.payload == PixelFormatFenceProbePayload.completion { + guard fence.flags.contains(.blockBefore) else { + framebufferRequestLock.unlock() + throw VNCError.protocol(.invalidData) + } + pixelFormatFenceCapabilityProbePayload = nil + expiredPixelFormatFenceCapabilityProbePayload = nil + pixelFormatFenceCapabilityProbeSequence = nil + let transition = takePendingPixelFormatTransitionLocked() + framebufferRequestLock.unlock() + + try resetZRLECompressionState(for: sequence) + if let transition { + enqueuePixelFormatTransition(transition) + } + return + } + state.pixelFormatTransitionFenceFlags = fence.flags.intersection([ + .blockBefore, + .blockAfter, + .syncNext + ]) + guard state.pixelFormatTransitionFenceFlags.contains(.syncNext) else { + let requiredFlags: VNCProtocol.FenceFlags = [.blockBefore, .blockAfter] + guard state.pixelFormatTransitionFenceFlags.intersection(requiredFlags) + == requiredFlags else { + framebufferRequestLock.unlock() + throw VNCError.protocol(.invalidData) + } + pixelFormatFenceCapabilityProbePayload = PixelFormatFenceProbePayload.completion + expiredPixelFormatFenceCapabilityProbePayload = nil + schedulePixelFormatFenceNegotiationTimeoutLocked() + framebufferRequestLock.unlock() + return + } + pixelFormatFenceCapabilityProbePayload = nil + expiredPixelFormatFenceCapabilityProbePayload = nil + pixelFormatFenceCapabilityProbeSequence = nil + let transition = takePendingPixelFormatTransitionLocked() + framebufferRequestLock.unlock() + + try resetZRLECompressionState(for: sequence) + if let transition { + enqueuePixelFormatTransition(transition) + } + return + } + guard fence.payload == pixelFormatTransitionFencePayload else { + framebufferRequestLock.unlock() + return + } + guard pixelFormatTransitionFenceWasSent else { + framebufferRequestLock.unlock() + throw VNCError.protocol(.invalidData) + } + let requiredFlags = pixelFormatTransitionRequiredFenceFlags + guard fence.flags.intersection(requiredFlags) == requiredFlags else { + framebufferRequestLock.unlock() + throw VNCError.protocol(.invalidData) + } + guard let pixelFormat = pixelFormatTransitionInFlight, + let sequence = pixelFormatTransitionInFlightSequence else { + framebufferRequestLock.unlock() + throw VNCError.protocol(.invalidData) + } + cancelPixelFormatTransitionDeadlineLocked() + pixelFormatTransitionFenceWasSent = false + pixelFormatTransitionFencePayload = nil + pixelFormatTransitionRequiredFenceFlags = [] + framebufferRequestLock.unlock() + + try beginPixelFormatTransition(pixelFormat, sequence: sequence) + completePixelFormatTransition() + } + + private func beginPixelFormatTransition( + _ pixelFormat: VNCProtocol.PixelFormat, + sequence: UInt64 + ) throws { + try withLifecycleLock { + guard connectionState.status == .connected, + let framebuffer = framebuffer else { + return + } + + try resetZRLECompressionState(for: sequence) + state.pixelFormat = pixelFormat + recreateFramebuffer(size: framebuffer.size, + screens: framebuffer.screens, + pixelFormat: pixelFormat) + } + } + + private func completePixelFormatTransition() { + framebufferRequestLock.lock() + isPixelFormatTransitionInFlight = false + pixelFormatTransitionInFlight = nil + pixelFormatTransitionInFlightSequence = nil + let nextTransition = takePendingPixelFormatTransitionLocked() + framebufferRequestLock.unlock() + + if let nextTransition { + enqueuePixelFormatTransition(nextTransition) + } else { + scheduleNextFramebufferUpdate() + } + } +} + // MARK: - Mouse Input public extension VNCConnection { #if canImport(ObjectiveC) @@ -240,7 +727,11 @@ extension VNCConnection { func reserveFramebufferUpdateRequest() -> Bool { framebufferRequestLock.lock() defer { framebufferRequestLock.unlock() } - guard !framebufferUpdateRequestOutstanding else { return false } + guard !framebufferUpdateRequestOutstanding, + pendingPixelFormatTransition == nil, + !isPixelFormatTransitionInFlight else { + return false + } framebufferUpdateRequestOutstanding = true return true } @@ -248,14 +739,24 @@ extension VNCConnection { func completeFramebufferUpdateRequest() { framebufferRequestLock.lock() framebufferUpdateRequestOutstanding = false + if let payload = pixelFormatTransitionFencePayload { + schedulePixelFormatTransitionDeadlineIfReadyLocked(payload: payload) + } + let transition = takePendingPixelFormatTransitionLocked() framebufferRequestLock.unlock() + + if let transition { + enqueuePixelFormatTransition(transition) + } } func scheduleNextFramebufferUpdate() { framebufferRequestLock.lock() framebufferPacingTask?.cancel() - guard !framebufferUpdateRequestOutstanding else { + guard !framebufferUpdateRequestOutstanding, + pendingPixelFormatTransition == nil, + !isPixelFormatTransitionInFlight else { framebufferPacingTask = nil framebufferRequestLock.unlock() return @@ -300,6 +801,18 @@ extension VNCConnection { framebufferPacingTask?.cancel() framebufferPacingTask = nil framebufferUpdateRequestOutstanding = false + pendingPixelFormatTransition = nil + isPixelFormatTransitionInFlight = false + pixelFormatTransitionInFlight = nil + pixelFormatTransitionInFlightSequence = nil + pixelFormatTransitionFencePayload = nil + pixelFormatTransitionRequiredFenceFlags = [] + pixelFormatTransitionFenceWasSent = false + cancelPixelFormatTransitionDeadlineLocked() + pixelFormatFenceCapabilityProbePayload = nil + expiredPixelFormatFenceCapabilityProbePayload = nil + pixelFormatFenceCapabilityProbeSequence = nil + cancelPixelFormatFenceNegotiationTimeoutLocked() framebufferRequestLock.unlock() } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Delegate.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Delegate.swift index 533fa606..ee629746 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Delegate.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Delegate.swift @@ -136,18 +136,8 @@ extension VNCConnection { private extension VNCConnection { func askDelegateForCredential(authenticationType: VNCAuthenticationType) async throws -> VNCCredential { - let credential: VNCCredential? = await withCheckedContinuation { continuation in - DispatchQueue.main.async { [weak self] in - guard let self, let delegate = self.delegate else { - continuation.resume(returning: nil) - return - } - - delegate.connection(self, credentialFor: authenticationType) { credential in - continuation.resume(returning: credential) - } - } - } + let request = beginCredentialRequest(authenticationType: authenticationType) + let credential = await request.value() guard let credential else { throw VNCError.authentication(.noAuthenticationDataProvided) @@ -156,3 +146,29 @@ private extension VNCConnection { return credential } } + +extension VNCConnection { + func beginCredentialRequest(authenticationType: VNCAuthenticationType) -> PendingCredentialRequest { + let requestID = UUID() + let request = PendingCredentialRequest { [weak self] in + self?.removeCredentialRequest(id: requestID) + } + guard registerCredentialRequest(request, id: requestID) else { + request.resolve(with: nil) + return request + } + + DispatchQueue.main.async { [weak self, request] in + guard request.pending else { return } + guard let self, let delegate = self.delegate else { + request.resolve(with: nil) + return + } + + delegate.connection(self, credentialFor: authenticationType) { [request] credential in + request.resolve(with: credential) + } + } + return request + } +} diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Receive.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Receive.swift index 70cdc958..9de4541c 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Receive.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection+Receive.swift @@ -55,6 +55,9 @@ private extension VNCConnection { case VNCProtocol.EndOfContinuousUpdates.messageType: try await handleEndOfContinuousUpdatesMessage() + case VNCProtocol.ServerFence.messageType: + try await handleServerFenceMessage() + default: throw VNCError.protocol(.unsupportedServerToClientMessage(messageType: messageType)) } @@ -195,9 +198,32 @@ private extension VNCConnection { func handleEndOfContinuousUpdatesMessage() async throws { didReceiveEndOfContinuousUpdates() } + + func handleServerFenceMessage() async throws { + let fence = try await VNCProtocol.ServerFence.receive(connection: connection) + try handleServerFence(fence) + } } extension VNCConnection { + func handleServerFence(_ fence: VNCProtocol.ServerFence) throws { + let first = publishPixelFormatFenceSupport() + + if fence.flags.contains(.request) { + let responseFlags = fence.flags.intersection(.blockBefore) + enqueueClientToServerMessage( + VNCProtocol.ClientFence(flags: responseFlags, payload: fence.payload) + ) + } else { + try completePixelFormatFence(fence) + } + + if first { + logger.logDebug("Fence supported (server sent ServerFence)") + enqueuePixelFormatFenceSupportProbe() + } + } + func didReceiveEndOfContinuousUpdates() { let first = !state.areContinuousUpdatesSupported diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection.swift index cfa15ec8..f666aa3f 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Connection/VNCConnection.swift @@ -110,12 +110,32 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { var framebufferRequestGeneration: UInt64 = 0 var framebufferUpdateRequestOutstanding = false var framebufferPacingTask: Task? + var pendingPixelFormatTransition: VNCProtocol.PixelFormat? + var isPixelFormatTransitionInFlight = false + var pixelFormatTransitionInFlight: VNCProtocol.PixelFormat? + var pixelFormatTransitionInFlightSequence: UInt64? + var pixelFormatTransitionFenceSequence: UInt64 = 0 + var pixelFormatTransitionFencePayload: Data? + var pixelFormatTransitionRequiredFenceFlags: VNCProtocol.FenceFlags = [] + var pixelFormatTransitionFenceWasSent = false + var pixelFormatTransitionDeadlineTask: Task? + var pixelFormatFenceCapabilityProbePayload: Data? + var expiredPixelFormatFenceCapabilityProbePayload: Data? + var pixelFormatFenceCapabilityProbeSequence: UInt64? + var pixelFormatFenceNegotiationTask: Task? + var nextPixelFormatChangeSequence: UInt64 = 0 private let queue = DispatchQueue(label: "com.royalapps.royalvnc.connectionqueue", attributes: .concurrent) private let lifecycleLock = NSRecursiveLock() + private let credentialRequestLock = NSLock() + private var pendingCredentialRequests = [ + UUID: WeakCredentialRequest + ]() private let sharedZStream: ZlibStream private let sharedZRLEZStream: ZlibStream + private let zrleCompressionStateLock = NSLock() + private var lastZRLECompressionResetSequence: UInt64 = 0 // MARK: - Internal Properties let taskPriority = TaskPriority.high @@ -219,7 +239,7 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { return enabledEncodings }() - func orderedEncodingTypes() throws -> [VNCEncodingType] { + func orderedEncodingTypes(pixelFormat: VNCProtocol.PixelFormat? = nil) throws -> [VNCEncodingType] { // Frame Encodings (Required) var encs: [VNCEncodingType] = [ VNCFrameEncodingType.copyRect.rawValue @@ -227,22 +247,23 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { // Frame Encodings (Customizable) var customizedFrameEncodings = settings.frameEncodings.map({ $0.rawValue }) + let negotiatedPixelFormat = pixelFormat ?? state.pixelFormat // TODO: Remove once we support ZRLE for non-24-bit pixel formats - if let pixelFormat = state.pixelFormat, + if let pixelFormat = negotiatedPixelFormat, customizedFrameEncodings.contains(VNCFrameEncodingType.zrle.rawValue), !VNCProtocol.ZRLEEncoding.supportsPixelFormat(pixelFormat) { customizedFrameEncodings.removeAll(where: { $0 == VNCFrameEncodingType.zrle.rawValue }) } - if let pixelFormat = state.pixelFormat, + if let pixelFormat = negotiatedPixelFormat, customizedFrameEncodings.contains(VNCFrameEncodingType.tight.rawValue), !VNCProtocol.TightEncoding.supportsPixelFormat(pixelFormat) { customizedFrameEncodings.removeAll(where: { $0 == VNCFrameEncodingType.tight.rawValue }) } #if canImport(VideoToolbox) - if let pixelFormat = state.pixelFormat, + if let pixelFormat = negotiatedPixelFormat, customizedFrameEncodings.contains(VNCFrameEncodingType.openH264.rawValue), !VNCProtocol.OpenH264Encoding.supportsPixelFormat(pixelFormat) { customizedFrameEncodings.removeAll(where: { $0 == VNCFrameEncodingType.openH264.rawValue }) @@ -261,6 +282,7 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { // Pseudo Encodings encs.append(contentsOf: [ VNCPseudoEncodingType.lastRect.rawValue, + VNCPseudoEncodingType.fence.rawValue, VNCPseudoEncodingType.continuousUpdates.rawValue, VNCPseudoEncodingType.extendedDesktopSize.rawValue, VNCPseudoEncodingType.desktopSize.rawValue, @@ -289,6 +311,15 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { return uniqueEncs } + func resetZRLECompressionState(for sequence: UInt64) throws { + zrleCompressionStateLock.lock() + defer { zrleCompressionStateLock.unlock() } + guard sequence > lastZRLECompressionResetSequence else { return } + + try sharedZRLEZStream.reset() + lastZRLECompressionResetSequence = sequence + } + // MARK: - Public Initializers public init(settings: Settings, logger: VNCLogger, @@ -370,6 +401,7 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { deinit { let _self = self + _self.cancelPendingCredentialRequests() _self.clipboardMonitor.delegate = nil _self.clipboardDelegate = nil @@ -379,14 +411,19 @@ public final class VNCConnection: NSObjectOrAnyObject, @unchecked Sendable { // MARK: - Internal Connection State API extension VNCConnection { - func beginConnecting() { + @discardableResult + func beginConnecting() -> Bool { lifecycleLock.lock() defer { lifecycleLock.unlock() } - guard !state.disconnectRequested else { return } + guard !state.disconnectRequested, + connectionState.status == .disconnected else { + return false + } updateConnectionState(.connecting) connection.start(queue: queue) + return true } func beginDisconnecting(error: Error? = nil) { @@ -397,6 +434,7 @@ extension VNCConnection { updateConnectionState(.disconnecting) handshakeTask?.cancel() handshakeTask = nil + cancelPendingCredentialRequests() receiveTask?.cancel() receiveTask = nil sendTask?.cancel() @@ -418,6 +456,47 @@ extension VNCConnection { beginDisconnecting(error: error) } + func withLifecycleLock(_ operation: () throws -> T) rethrows -> T { + lifecycleLock.lock() + defer { lifecycleLock.unlock() } + return try operation() + } + + func registerCredentialRequest( + _ request: PendingCredentialRequest, + id: UUID + ) -> Bool { + lifecycleLock.lock() + defer { lifecycleLock.unlock() } + + guard !state.disconnectRequested else { + return false + } + + credentialRequestLock.lock() + defer { credentialRequestLock.unlock() } + + pendingCredentialRequests[id] = WeakCredentialRequest(request) + return true + } + + func removeCredentialRequest(id: UUID) { + credentialRequestLock.lock() + pendingCredentialRequests.removeValue(forKey: id) + credentialRequestLock.unlock() + } + + func cancelPendingCredentialRequests() { + credentialRequestLock.lock() + let requests = pendingCredentialRequests.values.compactMap(\.request) + pendingCredentialRequests.removeAll() + credentialRequestLock.unlock() + + for request in requests { + request.resolve(with: nil) + } + } + func updateConnectionState(_ newConnectionState: ConnectionState) { self.connectionState = newConnectionState @@ -439,6 +518,71 @@ extension VNCConnection { } } +final class PendingCredentialRequest: @unchecked Sendable { + private let lock = NSLock() + private var continuation: CheckedContinuation? + private var isResolved = false + private var resolvedCredential: VNCCredential? + private var onResolution: (() -> Void)? + + init(onResolution: @escaping () -> Void) { + self.onResolution = onResolution + } + + var pending: Bool { + lock.lock() + defer { lock.unlock() } + return !isResolved + } + + func value() async -> VNCCredential? { + await withTaskCancellationHandler { + await withCheckedContinuation { continuation in + lock.lock() + if isResolved { + let credential = resolvedCredential + resolvedCredential = nil + lock.unlock() + continuation.resume(returning: credential) + } else { + self.continuation = continuation + lock.unlock() + } + } + } onCancel: { + resolve(with: nil) + } + } + + func resolve(with credential: VNCCredential?) { + lock.lock() + guard !isResolved else { + lock.unlock() + return + } + isResolved = true + let continuation = self.continuation + self.continuation = nil + if continuation == nil { + resolvedCredential = credential + } + let onResolution = self.onResolution + self.onResolution = nil + lock.unlock() + + continuation?.resume(returning: credential) + onResolution?() + } +} + +private final class WeakCredentialRequest { + weak var request: PendingCredentialRequest? + + init(_ request: PendingCredentialRequest) { + self.request = request + } +} + // MARK: - Connection State Change Handling private extension VNCConnection { func connectionStatusDidChange(_ newState: NetworkConnectionStatus) { diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Cursor/VNCCursor.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Cursor/VNCCursor.swift index 1f4a1366..6df75d2d 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Cursor/VNCCursor.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Cursor/VNCCursor.swift @@ -96,10 +96,8 @@ public extension VNCCursor { return } - // TODO: This assumes BGRA32 which might not be the case. - GraphicsUtils.copyBGRAtoRGBA(srcBuffer: ptrAddr, - dstBuffer: destinationPixelBuffer, - byteCount: byteCount) + destinationPixelBuffer.copyMemory(from: ptrAddr, + byteCount: byteCount) } } } diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Error/ProtocolError.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Error/ProtocolError.swift index 1448c415..8004e272 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Error/ProtocolError.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Error/ProtocolError.swift @@ -16,6 +16,7 @@ public extension VNCError { case framebufferUpdateReceivedWithoutFramebuffer case framebufferFailedToCreateIOSurface case setColourMapEntriesReceivedWithoutFramebuffer + case pixelFormatTransitionTimedOut case frameDecode(encodingType: VNCEncodingType, underlyingError: Error?) case zlibDecompress(underlyingError: Error?) case zrleInvalidSubencoding(subencoding: UInt8) @@ -50,6 +51,8 @@ public extension VNCError { return "Failed to create IOSurface for Framebuffer." case .setColourMapEntriesReceivedWithoutFramebuffer: return "A Set Colour Map Entries request has been retrieved but no Framebuffer has been created yet." + case .pixelFormatTransitionTimedOut: + return "The server did not acknowledge the pixel format transition." case .frameDecode(let encodingType, let underlyingError): return VNCError.combinedErrorDescription("An error occurred while decoding a Framebuffer Update Message (Encoding Type: \(encodingType)).", underlyingError: underlyingError) diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode+ObjC.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode+ObjC.swift index 630d69a8..4c86e3fb 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode+ObjC.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode+ObjC.swift @@ -86,7 +86,7 @@ public final class _ObjC_VNCKeyCode: NSObject { @objc public static let ansiKeypadEnter = X11KeySymbols.XK_KP_Enter @objc - public static let ansiKeypadDecimal = X11KeySymbols.XK_KP_Separator + public static let ansiKeypadDecimal = X11KeySymbols.XK_KP_Decimal @objc public static let f1 = X11KeySymbols.XK_F1 diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode.swift index 7cffc973..ef9659c4 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/Input/VNCKeyCode.swift @@ -48,7 +48,7 @@ public struct VNCKeyCode: Equatable { public static let ansiKeypadMinus = VNCKeyCode(X11KeySymbols.XK_KP_Subtract) public static let ansiKeypadPlus = VNCKeyCode(X11KeySymbols.XK_KP_Add) public static let ansiKeypadEnter = VNCKeyCode(X11KeySymbols.XK_KP_Enter) - public static let ansiKeypadDecimal = VNCKeyCode(X11KeySymbols.XK_KP_Separator) + public static let ansiKeypadDecimal = VNCKeyCode(X11KeySymbols.XK_KP_Decimal) public static let f1 = VNCKeyCode(X11KeySymbols.XK_F1) public static let f2 = VNCKeyCode(X11KeySymbols.XK_F2) @@ -159,10 +159,15 @@ public extension VNCKeyCode { var codes = [VNCKeyCode]() - for scalar in character.unicodeScalars { - let unicodeValue = scalar.value + for scalar in character.unicodeScalars { + let unicodeValue = scalar.value + let keySym = + (0x0020...0x007e).contains(unicodeValue) + || (0x00a0...0x00ff).contains(unicodeValue) + ? unicodeValue + : 0x0100_0000 | unicodeValue - codes.append(.init(unicodeValue)) + codes.append(.init(keySym)) } return codes diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/VNCPseudoEncodingType.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/VNCPseudoEncodingType.swift index d0300acb..97b6620d 100644 --- a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/VNCPseudoEncodingType.swift +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Sources/RoyalVNCKit/SDK/VNCPseudoEncodingType.swift @@ -8,6 +8,7 @@ public enum VNCPseudoEncodingType: VNCEncodingType { case lastRect = -224 case cursor = -239 case desktopName = -307 + case fence = -312 case continuousUpdates = -313 case desktopSize = -223 case extendedDesktopSize = -308 @@ -51,6 +52,8 @@ extension VNCPseudoEncodingType: CustomStringConvertible { "Cursor" case .desktopName: "Desktop Name" + case .fence: + "Fence" case .continuousUpdates: "Continuous Updates" case .desktopSize: diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Tests/RoyalVNCKitTests/AuditFindingsTests.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Tests/RoyalVNCKitTests/AuditFindingsTests.swift new file mode 100644 index 00000000..b577e7fa --- /dev/null +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Tests/RoyalVNCKitTests/AuditFindingsTests.swift @@ -0,0 +1,1298 @@ +import Foundation +import Testing + +#if canImport(Network) +import Network +#endif + +@testable import RoyalVNCKit + +struct AuditFindingsTests { + @Test + func rejectsOutOfRangeZRLEPackedPaletteIndex() async throws { + let framebuffer = try makeFramebuffer(width: 1, height: 1, depth: 24) + let encoding = VNCProtocol.ZRLEEncoding(zStream: ZlibStream()) + let rectangle = VNCProtocol.Rectangle( + xPosition: 0, + yPosition: 0, + width: 1, + height: 1, + encodingType: Int32(VNCFrameEncodingType.zrle.rawValue.rawValue) + ) + + var payload = Data([3]) + payload.append(contentsOf: [ + 0, 0, 0, + 64, 64, 64, + 128, 128, 128, + 0xC0, + ]) + let compressed = try ZlibOneShot.deflate(payload) + var compressedLength = UInt32(compressed.count).bigEndian + var wire = withUnsafeBytes(of: &compressedLength) { Data($0) } + wire.append(compressed) + + await #expect(throws: (any Error).self) { + try await encoding.decodeRectangle( + rectangle, + framebuffer: framebuffer, + connection: AuditBufferConnection(wire), + logger: VNCPrintLogger() + ) + } + } + + @Test + func derivesZRLEInflationLimitFromRectangleGeometry() { + #expect(VNCProtocol.ZRLEEncoding.maximumInflatedSize(width: 1, height: 1) == 384) + #expect(VNCProtocol.ZRLEEncoding.maximumInflatedSize(width: 64, height: 64) == 16_385) + #expect(VNCProtocol.ZRLEEncoding.maximumInflatedSize(width: 65, height: 1) == 894) + } + + @Test + func drainsPendingZlibOutputAfterConsumingAllInput() throws { + let expected = Data(repeating: 0xA5, count: 204_800) + let compressed = try ZlibOneShot.deflate(expected) + + let actual = try ZlibStream().decompressedData( + compressedData: compressed, + maximumOutputSize: expected.count + ) + + #expect(actual == expected) + } + + @Test + func rejectsInflatedChunkLargerThanOutputLimitWithoutTrapping() throws { + let compressed = try ZlibOneShot.deflate(Data(repeating: 0xA5, count: 385)) + + #expect(throws: (any Error).self) { + _ = try ZlibStream().decompressedData( + compressedData: compressed, + maximumOutputSize: 384 + ) + } + } + + @Test + func consumesSyncFlushBytesAfterFixedSizeOutputIsFull() throws { + let firstCompressed = Data([ + 0x78, 0x9C, 0x72, 0x74, 0x1C, 0x05, 0xA3, 0x60, 0x14, + 0x0C, 0x77, 0x00, 0x00, 0x00, 0x00, 0xFF, 0xFF, + 0x00, 0x00, 0x00, 0xFF, 0xFF, + ]) + let secondCompressed = Data([ + 0x72, 0x1A, 0x05, 0xA3, 0x60, 0x14, 0x0C, + 0x7B, 0x00, 0x00, 0x00, 0x00, 0xFF, 0xFF, + ]) + let stream = ZlibStream() + + let first = try stream.decompressedData( + compressedData: firstCompressed, + uncompressedSize: 1_000 + ) + let second = try stream.decompressedData( + compressedData: secondCompressed, + uncompressedSize: 1_000 + ) + + #expect(first == Data(repeating: 0x41, count: 1_000)) + #expect(second == Data(repeating: 0x42, count: 1_000)) + } + + @Test + func readsAllEightBitsOfThirdTightLengthByte() async throws { + let encoding = VNCProtocol.TightEncoding() + + let length = try await encoding.readCompactLength( + connection: AuditBufferConnection(Data([0x80, 0x80, 0xFF])), + logger: VNCPrintLogger() + ) + + #expect(length == 0xFF << 14) + } + + @Test + func retainsLegacyPixelFormatTransitionUntilSlowFramebufferUpdateCompletes() async throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + connection.framebufferUpdateRequestOutstanding = true + + connection.updateColorDepth(.depth8Bit) + + #expect(connection.pendingPixelFormatTransition?.depth == 8) + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + #expect(connection.pixelFormatFenceNegotiationTask != nil) + + for _ in 0..<100 { + guard connection.pixelFormatFenceNegotiationTask != nil else { break } + try await Task.sleep(nanoseconds: 20_000_000) + } + + #expect(connection.state.pixelFormat?.depth == 24) + #expect(connection.state.pixelFormat?.depth == connection.framebuffer?.sourcePixelFormat.depth) + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + #expect(connection.pixelFormatFenceNegotiationTask == nil) + #expect(connection.pendingPixelFormatTransition?.depth == 8) + + connection.completeFramebufferUpdateRequest() + + let transition = try #require(connection.clientToServerMessageQueue.dequeue()) + try await transition.message.send(connection: AuditWritingConnection()) + #expect(connection.pendingPixelFormatTransition == nil) + #expect(connection.state.pixelFormat?.depth == 8) + #expect(connection.state.pixelFormat?.depth == connection.framebuffer?.sourcePixelFormat.depth) + } + + @Test + func preservesEarlyPixelFormatTransitionUntilFenceSupportArrives() throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + connection.framebufferUpdateRequestOutstanding = true + + connection.updateColorDepth(.depth8Bit) + + #expect(connection.pendingPixelFormatTransition?.depth == 8) + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockAfter, .syncNext], + payload: Data("support".utf8) + ) + ) + + #expect(connection.pendingPixelFormatTransition?.depth == 8) + #expect(connection.clientToServerMessageQueue.dequeue() != nil) + #expect(connection.clientToServerMessageQueue.dequeue() != nil) + } + + @Test + func publishesFenceNegotiationBeforeExposingSupport() async throws { + let logger = AuditCallbackLogger() + let connection = VNCConnection( + settings: makeSettings(), + logger: logger, + framebufferAllocator: VNCFramebufferMallocAllocator(), + context: nil + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.state.areContinuousUpdatesEnabled = true + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + logger.onDebug = { message in + guard message == "Fence supported (server sent ServerFence)" else { return } + connection.updateColorDepth(.depth8Bit) + } + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + + #expect(connection.state.areFencesSupported) + #expect(connection.pixelFormatFenceCapabilityProbePayload != nil) + #expect(connection.pendingPixelFormatTransition?.depth == 8) + #expect(connection.state.pixelFormat?.depth == 24) + + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityWriter = AuditWritingConnection() + try await capabilityProbe.message.send(connection: capabilityWriter) + let capabilityLength = Int(capabilityWriter.data[8]) + let capabilityPayload = Data(capabilityWriter.data[9..<(9 + capabilityLength)]) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: capabilityPayload + ) + ) + + #expect(connection.pendingPixelFormatTransition == nil) + #expect(connection.clientToServerMessageQueue.dequeue() != nil) + logger.onDebug = nil + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func startsFenceCapabilityTimeoutAfterDelayedProbeSend() async throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + + #expect(connection.pixelFormatFenceCapabilityProbePayload != nil) + #expect(connection.pixelFormatFenceNegotiationTask == nil) + + try await capabilityProbe.message.send( + connection: AuditWritingConnection(delayNanoseconds: 1_100_000_000) + ) + + #expect(connection.pixelFormatFenceCapabilityProbePayload != nil) + #expect(connection.pixelFormatFenceNegotiationTask != nil) + connection.expirePixelFormatFenceNegotiation() + } + + @Test + func disconnectsWhenFenceCapabilityProbeIsUnanswered() async throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + + connection.updateColorDepth(.depth8Bit) + #expect(connection.pendingPixelFormatTransition?.depth == 8) + + connection.expirePixelFormatFenceNegotiation() + + #expect(connection.pixelFormatFenceCapabilityProbePayload == nil) + #expect(connection.expiredPixelFormatFenceCapabilityProbePayload == nil) + #expect(connection.connectionState.status == .disconnected) + #expect(connection.pendingPixelFormatTransition == nil) + #expect(connection.state.pixelFormat?.depth == 24) + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + } + + @Test + func lateFenceCapabilityResponseCannotReviveTimedOutConnection() async throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let probePayload = try #require(connection.pixelFormatFenceCapabilityProbePayload) + + connection.updateColorDepth(.depth8Bit) + connection.expirePixelFormatFenceNegotiation() + #expect(connection.connectionState.status == .disconnected) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: probePayload + ) + ) + + #expect(connection.expiredPixelFormatFenceCapabilityProbePayload == nil) + #expect(connection.connectionState.status == .disconnected) + #expect(connection.state.pixelFormatTransitionFenceFlags.isEmpty) + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + } + + @Test + func rejectsPartialFenceBoundariesDuringContinuousUpdates() throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.state.areFencesSupported = true + connection.state.pixelFormatTransitionFenceFlags = [.blockAfter] + connection.state.areContinuousUpdatesEnabled = true + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + connection.framebufferUpdateRequestOutstanding = true + + connection.updateColorDepth(.depth8Bit) + + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + #expect(connection.pendingPixelFormatTransition == nil) + #expect(connection.state.pixelFormat?.depth == 24) + } + + @Test + func synchronizesPixelFormatTransitionWithFenceResponse() async throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + connection.framebufferUpdateRequestOutstanding = true + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + let supportResponse = try #require(connection.clientToServerMessageQueue.dequeue()) + let supportWriter = AuditWritingConnection() + try await supportResponse.message.send(connection: supportWriter) + #expect(supportWriter.data[0] == VNCProtocol.ClientFence.messageType) + #expect(supportWriter.data[4..<8] == Data([0, 0, 0, 1])) + #expect(supportWriter.data[8] == 7) + + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityWriter = AuditWritingConnection() + try await capabilityProbe.message.send(connection: capabilityWriter) + let capabilityLength = Int(capabilityWriter.data[8]) + let capabilityPayload = Data(capabilityWriter.data[9..<(9 + capabilityLength)]) + #expect(capabilityWriter.data[0] == VNCProtocol.ClientFence.messageType) + #expect(capabilityWriter.data[(9 + capabilityLength)] == 0) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: capabilityPayload + ) + ) + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + let writer = AuditWritingConnection { + #expect(connection.state.pixelFormat?.depth == 24) + } + try await queued.message.send(connection: writer) + + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect(writer.data[0] == VNCProtocol.ClientFence.messageType) + #expect(writer.data[4..<8] == Data([0x80, 0, 0, 5])) + #expect(writer.data[8] == 8) + #expect(writer.data[17] == VNCProtocol.SetPixelFormat(pixelFormat: framebuffer.sourcePixelFormat).messageType) + #expect(writer.data[37] == VNCProtocol.SetEncodings(encodingTypes: []).messageType) + let encodingValues = setEncodingValues(in: writer.data, at: 37) + #expect(encodingValues.contains(Int32(VNCFrameEncodingType.raw.rawValue.rawValue))) + let completionFenceOffset = setEncodingsEndOffset(in: writer.data, at: 37) + #expect(writer.data[completionFenceOffset] == VNCProtocol.ClientFence.messageType) + #expect( + writer.data[(completionFenceOffset + 4)..<(completionFenceOffset + 8)] + == Data([0x80, 0, 0, 1]) + ) + #expect(connection.state.pixelFormat?.depth == 24) + + let synchronizationPayload = Data(writer.data[9..<17]) + let completionPayload = Data( + writer.data[(completionFenceOffset + 9)..<(completionFenceOffset + 17)] + ) + connection.completeFramebufferUpdateRequest() + #expect(connection.pixelFormatTransitionDeadlineTask != nil) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: synchronizationPayload + ) + ) + #expect(connection.state.pixelFormat?.depth == 24) + #expect(connection.pixelFormatTransitionDeadlineTask != nil) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore], + payload: completionPayload + ) + ) + #expect(connection.state.pixelFormat?.depth == 8) + #expect(connection.state.pixelFormat?.depth == connection.framebuffer?.sourcePixelFormat.depth) + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect(connection.connectionState.status == .connected) + } + + @Test + func renegotiatesEncodingsForSynchronizedPixelFormatTransition() async throws { + let connection = try await makeFenceCapableConnection( + settings: makeSettings(frameEncodings: [.tight, .zrle, .openH264, .raw]) + ) + + let depth24Encodings = try connection.orderedEncodingTypes( + pixelFormat: VNCProtocol.PixelFormat(depth: 24) + ) + #expect(depth24Encodings.contains(VNCFrameEncodingType.tight.rawValue)) + #expect(depth24Encodings.contains(VNCFrameEncodingType.zrle.rawValue)) +#if canImport(VideoToolbox) + #expect(depth24Encodings.contains(VNCFrameEncodingType.openH264.rawValue)) +#endif + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + let writer = AuditWritingConnection { + #expect(connection.state.pixelFormat?.depth == 24) + } + try await queued.message.send(connection: writer) + + #expect(writer.data[0] == VNCProtocol.ClientFence.messageType) + #expect(writer.data[17] == VNCProtocol.SetPixelFormat(pixelFormat: VNCProtocol.PixelFormat(depth: 8)).messageType) + #expect(writer.data[37] == VNCProtocol.SetEncodings(encodingTypes: []).messageType) + let values = setEncodingValues(in: writer.data, at: 37) + #expect(!values.contains(Int32(VNCFrameEncodingType.tight.rawValue.rawValue))) + #expect(!values.contains(Int32(VNCFrameEncodingType.zrle.rawValue.rawValue))) + #expect(!values.contains(Int32(VNCFrameEncodingType.openH264.rawValue.rawValue))) + #expect(values.contains(Int32(VNCFrameEncodingType.copyRect.rawValue.rawValue))) + #expect(values.contains(Int32(VNCFrameEncodingType.raw.rawValue.rawValue))) + let completionFenceOffset = setEncodingsEndOffset(in: writer.data, at: 37) + #expect(writer.data[completionFenceOffset] == VNCProtocol.ClientFence.messageType) + + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func resetsZRLECompressionAtSynchronizedPixelFormatBoundaries() async throws { + let connection = VNCConnection( + settings: makeSettings(frameEncodings: [.zrle, .raw]), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + let zrle = try #require( + connection.encodings[VNCFrameEncodingType.zrle.rawValue] as? VNCProtocol.ZRLEEncoding + ) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + + let compressedChunks = continuousZlibChunks() + let first = try zrle.zStream.decompressedData( + compressedData: compressedChunks[0], + uncompressedSize: 1_000 + ) + #expect(first == Data(repeating: 0x41, count: 1_000)) + try await capabilityProbe.message.send(connection: AuditWritingConnection()) + let second = try zrle.zStream.decompressedData( + compressedData: compressedChunks[1], + uncompressedSize: 1_000 + ) + #expect(second == Data(repeating: 0x42, count: 1_000)) + + let capabilityPayload = try #require(connection.pixelFormatFenceCapabilityProbePayload) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: capabilityPayload + ) + ) + let restartedAfterProbe = try zrle.zStream.decompressedData( + compressedData: compressedChunks[0], + uncompressedSize: 1_000 + ) + #expect(restartedAfterProbe == Data(repeating: 0x41, count: 1_000)) + + connection.framebufferUpdateRequestOutstanding = true + connection.updateColorDepth(.depth8Bit) + let transition = try #require(connection.clientToServerMessageQueue.dequeue()) + let transitionWriter = AuditWritingConnection() + try await transition.message.send(connection: transitionWriter) + let continuedBeforeBoundary = try zrle.zStream.decompressedData( + compressedData: compressedChunks[1], + uncompressedSize: 1_000 + ) + #expect(continuedBeforeBoundary == Data(repeating: 0x42, count: 1_000)) + + let completionFenceOffset = setEncodingsEndOffset(in: transitionWriter.data, at: 37) + let synchronizationPayload = Data(transitionWriter.data[9..<17]) + let completionPayload = Data( + transitionWriter.data[(completionFenceOffset + 9)..<(completionFenceOffset + 17)] + ) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: synchronizationPayload + ) + ) + let continuedAfterSynchronizationFence = try zrle.zStream.decompressedData( + compressedData: compressedChunks[2], + uncompressedSize: 1_000 + ) + #expect(continuedAfterSynchronizationFence == Data(repeating: 0x43, count: 1_000)) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore], + payload: completionPayload + ) + ) + try verifyFreshZRLEStream(zrle.zStream, byte: 0x44) + + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func defersZRLEResetUntilTrailingFenceWhenSyncNextIsUnsupported() async throws { + let connection = VNCConnection( + settings: makeSettings(frameEncodings: [.zrle, .raw]), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + let zrle = try #require( + connection.encodings[VNCFrameEncodingType.zrle.rawValue] as? VNCProtocol.ZRLEEncoding + ) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .blockAfter], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityWriter = AuditWritingConnection() + try await capabilityProbe.message.send(connection: capabilityWriter) + + let compressedChunks = continuousZlibChunks() + let first = try zrle.zStream.decompressedData( + compressedData: compressedChunks[0], + uncompressedSize: 1_000 + ) + #expect(first == Data(repeating: 0x41, count: 1_000)) + let second = try zrle.zStream.decompressedData( + compressedData: compressedChunks[1], + uncompressedSize: 1_000 + ) + #expect(second == Data(repeating: 0x42, count: 1_000)) + + let capabilityLength = Int(capabilityWriter.data[8]) + let capabilityPayload = Data(capabilityWriter.data[9..<(9 + capabilityLength)]) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .blockAfter], + payload: capabilityPayload + ) + ) + + let continuedAfterPartialResponse = try zrle.zStream.decompressedData( + compressedData: compressedChunks[2], + uncompressedSize: 1_000 + ) + #expect(continuedAfterPartialResponse == Data(repeating: 0x43, count: 1_000)) + + let trailingFenceOffset = 9 + capabilityLength + 20 + #expect(capabilityWriter.data[trailingFenceOffset] == VNCProtocol.ClientFence.messageType) + #expect( + capabilityWriter.data[(trailingFenceOffset + 4)..<(trailingFenceOffset + 8)] + == Data([0x80, 0, 0, 1]) + ) + let trailingPayloadLength = Int(capabilityWriter.data[trailingFenceOffset + 8]) + let trailingPayload = Data( + capabilityWriter.data[ + (trailingFenceOffset + 9)..<(trailingFenceOffset + 9 + trailingPayloadLength) + ] + ) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore], + payload: trailingPayload + ) + ) + + try verifyFreshZRLEStream(zrle.zStream, byte: 0x44) + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func ignoresStaleZRLEResetFromLateCapabilityProbeResponse() async throws { + let connection = VNCConnection( + settings: makeSettings(frameEncodings: [.zrle, .raw]), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + let zrle = try #require( + connection.encodings[VNCFrameEncodingType.zrle.rawValue] as? VNCProtocol.ZRLEEncoding + ) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + try await capabilityProbe.message.send(connection: AuditWritingConnection()) + let capabilityPayload = try #require(connection.pixelFormatFenceCapabilityProbePayload) + + let compressedChunks = continuousZlibChunks() + _ = try zrle.zStream.decompressedData( + compressedData: compressedChunks[0], + uncompressedSize: 1_000 + ) + connection.updateColorDepth(.depth8Bit) + connection.expirePixelFormatFenceNegotiation() + #expect(connection.connectionState.status == .disconnected) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: capabilityPayload + ) + ) + let continuedAfterLateProbe = try zrle.zStream.decompressedData( + compressedData: compressedChunks[1], + uncompressedSize: 1_000 + ) + #expect(continuedAfterLateProbe == Data(repeating: 0x42, count: 1_000)) + + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func rejectsPixelFormatFenceResponseBeforeRequestIsSent() async throws { + let connection = try await makeFenceCapableConnection() + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + let payload = try #require(connection.pixelFormatTransitionFencePayload) + + #expect(throws: (any Error).self) { + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: payload + ) + ) + } + #expect(connection.state.pixelFormat?.depth == 24) + + try await queued.message.send(connection: AuditWritingConnection()) + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func acceptsPixelFormatFenceResponseWhileWriteCompletes() async throws { + let connection = try await makeFenceCapableConnection() + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + let payload = try #require(connection.pixelFormatTransitionFencePayload) + connection.completeFramebufferUpdateRequest() + var responseError: Error? + let writer = AuditWritingConnection { + do { + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: payload + ) + ) + } catch { + responseError = error + } + } + + try await queued.message.send(connection: writer) + + #expect(responseError == nil) + #expect(connection.state.pixelFormat?.depth == 8) + #expect(!connection.pixelFormatTransitionFenceWasSent) + } + + @Test + func rollsBackPixelFormatFenceWhenWriteFails() async throws { + let connection = try await makeFenceCapableConnection() + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + + await #expect(throws: AuditWriteError.self) { + try await queued.message.send(connection: AuditFailingWritingConnection()) + } + + #expect(connection.isPixelFormatTransitionInFlight) + #expect(!connection.pixelFormatTransitionFenceWasSent) + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func waitsForSlowFramebufferBoundaryBeforeArmingTransitionDeadline() async throws { + let connection = try await makeFenceCapableConnection() + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + try await queued.message.send(connection: AuditWritingConnection()) + + try await Task.sleep(nanoseconds: 100_000_000) + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect(connection.connectionState.status == .connected) + + connection.completeFramebufferUpdateRequest() + #expect(connection.pixelFormatTransitionDeadlineTask != nil) + + connection.cancelFramebufferUpdateScheduling() + } + + @Test + func disconnectsWhenPixelFormatTransitionFenceIsMissing() async throws { + let connection = try await makeFenceCapableConnection() + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + try await queued.message.send(connection: AuditWritingConnection()) + let payload = try #require(connection.pixelFormatTransitionFencePayload) + + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect(connection.isPixelFormatTransitionInFlight) + + connection.expirePixelFormatTransitionDeadline(payload: payload) + #expect(connection.connectionState.status == .connected) + + connection.completeFramebufferUpdateRequest() + #expect(connection.pixelFormatTransitionDeadlineTask != nil) + connection.expirePixelFormatTransitionDeadline(payload: payload) + + #expect(connection.connectionState.status == .disconnected) + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect(!connection.isPixelFormatTransitionInFlight) + } + + @Test + func cancellingFramebufferSchedulingInvalidatesPixelFormatTransitionDeadline() async throws { + let connection = try await makeFenceCapableConnection() + + connection.updateColorDepth(.depth8Bit) + let queued = try #require(connection.clientToServerMessageQueue.dequeue()) + try await queued.message.send(connection: AuditWritingConnection()) + let payload = try #require(connection.pixelFormatTransitionFencePayload) + + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + connection.completeFramebufferUpdateRequest() + #expect(connection.pixelFormatTransitionDeadlineTask != nil) + + connection.cancelFramebufferUpdateScheduling() + connection.expirePixelFormatTransitionDeadline(payload: payload) + + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect(connection.connectionState.status == .connected) + } + + @Test + func waitsForOutstandingUpdateWhenSyncNextIsUnsupported() async throws { + let connection = VNCConnection( + settings: makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + connection.framebufferUpdateRequestOutstanding = true + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request], + payload: Data() + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityWriter = AuditWritingConnection() + try await capabilityProbe.message.send(connection: capabilityWriter) + let capabilityLength = Int(capabilityWriter.data[8]) + let capabilityPayload = Data(capabilityWriter.data[9..<(9 + capabilityLength)]) + + connection.updateColorDepth(.depth8Bit) + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .blockAfter], + payload: capabilityPayload + ) + ) + + connection.completeFramebufferUpdateRequest() + #expect(connection.clientToServerMessageQueue.dequeue() == nil) + let trailingPayload = try #require(connection.pixelFormatFenceCapabilityProbePayload) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore], + payload: trailingPayload + ) + ) + + let transition = try #require(connection.clientToServerMessageQueue.dequeue()) + let transitionWriter = AuditWritingConnection { + #expect(connection.state.pixelFormat?.depth == 8) + } + try await transition.message.send(connection: transitionWriter) + #expect(connection.pixelFormatTransitionDeadlineTask == nil) + #expect( + transitionWriter.data[0] + == VNCProtocol.SetPixelFormat(pixelFormat: framebuffer.sourcePixelFormat).messageType + ) + #expect(transitionWriter.data[20] == VNCProtocol.SetEncodings(encodingTypes: []).messageType) + #expect(connection.state.pixelFormat?.depth == 8) + #expect(connection.state.pixelFormat?.depth == connection.framebuffer?.sourcePixelFormat.depth) + } + + @Test + func advertisesAndDecodesFenceExtension() async throws { + let connection = VNCConnection(settings: makeSettings()) + #expect(try connection.orderedEncodingTypes().contains(VNCPseudoEncodingType.fence.rawValue)) + + var body = Data([0, 0, 0]) + body.append(UInt32(0x8000_0004), bigEndian: true) + body.append(UInt8(3)) + body.append(Data([1, 2, 3])) + + let fence = try await VNCProtocol.ServerFence.receive( + connection: AuditBufferConnection(body) + ) + + #expect(fence.flags == [.request, .syncNext]) + #expect(fence.payload == Data([1, 2, 3])) + } + + @Test + func copyRectPreservesInternalFramebufferPixels() throws { + let framebuffer = try makeFramebuffer(width: 2, height: 1, depth: 16) + var redPixel = Data([0x00, 0x7C]) + framebuffer.update( + region: VNCRegion(x: 0, y: 0, width: 1, height: 1), + data: &redPixel + ) + + framebuffer.copy( + region: VNCRegion(x: 0, y: 0, width: 1, height: 1), + to: VNCRegion(x: 1, y: 0, width: 1, height: 1) + ) + + let pixels = Data(bytes: framebuffer.surfaceAddress, count: framebuffer.surfaceByteCount) + #expect(pixels[0..<4] == pixels[4..<8]) + } + +#if canImport(Network) + @Test + func returnsFinalNetworkContentBeforeReportingEOF() throws { + let finalContent = Data([1, 2, 3]) + + let received = try NWConnection.validateReadContent( + finalContent, + isComplete: true, + error: nil, + minimumLength: 1, + maximumLength: 3 + ) + + #expect(received == finalContent) + } +#endif + + @Test + func rejectsRepeatConnectAttempts() { + let connection = VNCConnection(settings: makeSettings()) + connection.connectionState = .connecting + + #expect(connection.beginConnecting() == false) + } + + @Test @MainActor + func disconnectCancelsPendingCredentialContinuation() async { + let connection = VNCConnection(settings: makeSettings()) + let delegate = PendingCredentialDelegate() + connection.delegate = delegate + + let credentialTask = Task { + try await connection.askDelegateForPasswordCredential(authenticationType: .vnc) + } + + while delegate.completion == nil { + await Task.yield() + } + + connection.disconnect() + + do { + _ = try await credentialTask.value + Issue.record("Expected disconnect to cancel the pending credential request") + } catch { + #expect(connection.connectionState.status == .disconnected) + } + + delegate.completion?(VNCPasswordCredential(password: "late")) + } + + @Test @MainActor + func unresolvedCredentialDelegateRequestDoesNotRetainConnection() async { + var connection: VNCConnection? = VNCConnection(settings: makeSettings()) + weak var weakConnection = connection + let delegate = PendingCredentialDelegate() + connection?.delegate = delegate + + let request = connection!.beginCredentialRequest(authenticationType: .vnc) + let credentialTask = Task { + await request.value() + } + + while delegate.completion == nil { + await Task.yield() + } + + connection = nil + for _ in 0..<100 where weakConnection != nil { + await Task.yield() + } + + #expect(weakConnection == nil) + #expect(await credentialTask.value == nil) + + delegate.completion?(VNCPasswordCredential(password: "late")) + } + + @Test @MainActor + func taskCancellationResolvesPendingCredentialRequest() async { + let connection = VNCConnection(settings: makeSettings()) + let delegate = PendingCredentialDelegate() + connection.delegate = delegate + + let credentialTask = Task { + try await connection.askDelegateForPasswordCredential(authenticationType: .vnc) + } + + while delegate.completion == nil { + await Task.yield() + } + + credentialTask.cancel() + + do { + _ = try await credentialTask.value + Issue.record("Expected task cancellation to resolve the pending credential request") + } catch { + #expect(credentialTask.isCancelled) + } + + delegate.completion?(VNCPasswordCredential(password: "late")) + } + + @Test + func releasesResolvedCredentialAfterValueIsConsumed() async { + let request = PendingCredentialRequest(onResolution: {}) + var credential: VNCPasswordCredential? = VNCPasswordCredential(password: "secret") + weak var weakCredential = credential + + request.resolve(with: credential) + credential = nil + var received = await request.value() as? VNCPasswordCredential + + #expect(received === weakCredential) + received = nil + #expect(weakCredential == nil) + } + + @Test + func preservesAlreadyRGBAFormattedCursorChannels() { + let cursor = VNCCursor( + imageData: Data([0x11, 0x22, 0x33, 0x44]), + size: VNCSize(width: 1, height: 1), + hotspot: .zero, + bitsPerComponent: 8, + bitsPerPixel: 32, + bytesPerPixel: 4 + ) + let destination = UnsafeMutableRawPointer.allocate(byteCount: 4, alignment: 1) + defer { destination.deallocate() } + + cursor.copyPixelDataToRGBA32(destinationPixelBuffer: destination) + + #expect(Data(bytes: destination, count: 4) == Data([0x11, 0x22, 0x33, 0x44])) + } + + private func makeFramebuffer(width: UInt16, height: UInt16, depth: UInt8) throws + -> VNCFramebuffer + { + try VNCFramebuffer( + logger: VNCPrintLogger(), + size: VNCSize(width: width, height: height), + screens: [], + pixelFormat: VNCProtocol.PixelFormat(depth: depth), + allocator: VNCFramebufferMallocAllocator() + ) + } + + private func makeSettings( + frameEncodings: [VNCFrameEncodingType] = [.raw] + ) -> VNCConnection.Settings { + VNCConnection.Settings( + isDebugLoggingEnabled: false, + hostname: "127.0.0.1", + port: 5900, + isShared: true, + isScalingEnabled: true, + useDisplayLink: false, + inputMode: .none, + isClipboardRedirectionEnabled: false, + colorDepth: .depth24Bit, + frameEncodings: frameEncodings + ) + } + + private func makeFenceCapableConnection( + settings: VNCConnection.Settings? = nil + ) async throws -> VNCConnection { + let connection = VNCConnection( + settings: settings ?? makeSettings(), + framebufferAllocator: VNCFramebufferMallocAllocator() + ) + let framebuffer = try makeFramebuffer(width: 2, height: 2, depth: 24) + connection.framebuffer = framebuffer + connection.state.pixelFormat = framebuffer.sourcePixelFormat + connection.connectionState = .connected + connection._framebufferUpdatePolicy = .paused + connection.framebufferUpdateRequestOutstanding = true + + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.request, .blockBefore, .syncNext], + payload: Data("support".utf8) + ) + ) + _ = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityProbe = try #require(connection.clientToServerMessageQueue.dequeue()) + let capabilityWriter = AuditWritingConnection() + try await capabilityProbe.message.send(connection: capabilityWriter) + let capabilityLength = Int(capabilityWriter.data[8]) + let capabilityPayload = Data(capabilityWriter.data[9..<(9 + capabilityLength)]) + try connection.handleServerFence( + VNCProtocol.ServerFence( + messageType: VNCProtocol.ServerFence.messageType, + flags: [.blockBefore, .syncNext], + payload: capabilityPayload + ) + ) + return connection + } + + private func setEncodingValues(in data: Data, at offset: Int) -> [Int32] { + let count = Int(data[offset + 2]) << 8 | Int(data[offset + 3]) + return (0.. Int { + offset + 4 + setEncodingValues(in: data, at: offset).count * 4 + } + + private func verifyFreshZRLEStream(_ stream: ZlibStream, byte: UInt8) throws { + let expected = Data(repeating: byte, count: 64) + let actual = try stream.decompressedData( + compressedData: ZlibOneShot.deflate(expected), + maximumOutputSize: expected.count + ) + #expect(actual == expected) + } + + private func continuousZlibChunks() -> [Data] { + [ + Data([ + 0x78, 0x9C, 0x72, 0x74, 0x1C, 0x05, 0xA3, 0x60, 0x14, + 0x0C, 0x77, 0, 0, 0, 0, 0xFF, 0xFF, + ]), + Data([ + 0x72, 0x1A, 0x05, 0xA3, 0x60, 0x14, 0x0C, + 0x7B, 0, 0, 0, 0, 0xFF, 0xFF, + ]), + Data([ + 0x72, 0x1E, 0x05, 0xA3, 0x60, 0x14, 0x0C, + 0x7B, 0, 0, 0, 0, 0xFF, 0xFF, + ]), + ] + } +} + +private final class AuditBufferConnection: NetworkConnectionReading { + private let data: Data + private var offset = 0 + + init(_ data: Data) { + self.data = data + } + + func read(minimumLength: Int, maximumLength: Int) async throws -> Data { + let remaining = data.count - offset + guard minimumLength > 0, + maximumLength >= minimumLength, + remaining >= minimumLength else { + throw VNCError.protocol(.noData) + } + + let count = min(maximumLength, remaining) + defer { offset += count } + return data.subdata(in: offset..<(offset + count)) + } +} + +private final class AuditWritingConnection: NetworkConnectionWriting { + var data = Data() + private let delayNanoseconds: UInt64 + private let onWrite: () -> Void + + init(delayNanoseconds: UInt64 = 0, onWrite: @escaping () -> Void = {}) { + self.delayNanoseconds = delayNanoseconds + self.onWrite = onWrite + } + + func write(data: Data) async throws { + if delayNanoseconds > 0 { + try await Task.sleep(nanoseconds: delayNanoseconds) + } + onWrite() + self.data.append(data) + } +} + +private final class AuditCallbackLogger: VNCLogger { + var isDebugLoggingEnabled = false + var onDebug: ((String) -> Void)? + + func logDebug(_ message: @autoclosure () -> String) { + onDebug?(message()) + } + + func logInfo(_ message: String) {} + func logWarning(_ message: String) {} + func logError(_ message: String) {} +} + +private struct AuditWriteError: Error {} + +private final class AuditFailingWritingConnection: NetworkConnectionWriting { + func write(data: Data) async throws { + throw AuditWriteError() + } +} + +@MainActor +private final class PendingCredentialDelegate: VNCConnectionDelegate { + var completion: ((VNCCredential?) -> Void)? + + func connection( + _ connection: VNCConnection, + stateDidChange connectionState: VNCConnection.ConnectionState + ) {} + + func connection( + _ connection: VNCConnection, + credentialFor authenticationType: VNCAuthenticationType, + completion: @escaping (VNCCredential?) -> Void + ) { + self.completion = completion + } + + func connection(_ connection: VNCConnection, didCreateFramebuffer framebuffer: VNCFramebuffer) {} + func connection(_ connection: VNCConnection, didResizeFramebuffer framebuffer: VNCFramebuffer) {} + + func connection( + _ connection: VNCConnection, + didUpdateFramebuffer framebuffer: VNCFramebuffer, + x: UInt16, + y: UInt16, + width: UInt16, + height: UInt16 + ) {} + + func connection(_ connection: VNCConnection, didUpdateCursor cursor: VNCCursor) {} +} diff --git a/macos/CrabfleetMac/Vendor/RoyalVNCKit/Tests/RoyalVNCKitTests/SecurityAndInputTests.swift b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Tests/RoyalVNCKitTests/SecurityAndInputTests.swift new file mode 100644 index 00000000..91701073 --- /dev/null +++ b/macos/CrabfleetMac/Vendor/RoyalVNCKit/Tests/RoyalVNCKitTests/SecurityAndInputTests.swift @@ -0,0 +1,262 @@ +import Foundation +import Testing + +@testable import RoyalVNCKit + +struct SecurityAndInputTests { + typealias ARDKeyAgreement = VNCProtocol.ARDAuthentication.DiffieHellmanKeyAgreement + typealias UltraVNCKeyAgreement = + VNCProtocol.UltraVNCMSLogonIIAuthentication.DiffieHellmanKeyAgreement + typealias UltraVNCBigNum = + UltraVNCKeyAgreement.UltraVNCBigNum + + @Test + func rejectsWeakAppleRemoteDesktopModuli() { + for prime in [ + Data(repeating: 0, count: 128), + Data(repeating: 1, count: 128), + Data([0]) + Data(repeating: 0xFF, count: 127), + Data([0x80]) + Data(repeating: 0, count: 127), + appleRemoteDesktopUnsafePrime, + ] { + let agreement = ARDKeyAgreement( + prime: prime, + generator: Data([2]), + peerKey: paddedARDValue(2), + keyLength: 128 + ) + + #expect(agreement.map { _ in true } == nil) + } + } + + @Test + func rejectsAppleRemoteDesktopElementsOutsideTheSafeRange() { + let prime = appleRemoteDesktopSafePrime + let primeMinusOne = prime.dropLast() + Data([0xFE]) + + for generator in [ + Data([0]), + Data([1]), + primeMinusOne, + prime, + ] { + #expect( + ARDKeyAgreement( + prime: prime, + generator: generator, + peerKey: paddedARDValue(2), + keyLength: 128 + ) == nil + ) + } + + for peerKey in [ + paddedARDValue(0), + paddedARDValue(1), + primeMinusOne, + prime, + ] { + #expect( + ARDKeyAgreement( + prime: prime, + generator: Data([2]), + peerKey: peerKey, + keyLength: 128 + ) == nil + ) + } + } + + @Test + func acceptsSafeAppleRemoteDesktopKeyMaterial() { + let agreement = ARDKeyAgreement( + prime: appleRemoteDesktopSafePrime, + generator: Data([2]), + peerKey: paddedARDValue(2), + keyLength: 128 + ) + + #expect(agreement?.publicKey.count == 128) + #expect(agreement?.publicKey.contains { $0 != 0 } == true) + #expect(agreement?.secretKey.count == 128) + #expect(agreement?.secretKey.contains { $0 != 0 } == true) + } + + @Test + func evictsValidatedAppleRemoteDesktopSafePrimesInInsertionOrder() { + var cache = ARDKeyAgreement.ValidatedSafePrimeCache(capacity: 3) + let values = (0..<5).map { Data([$0]) } + + for value in values.prefix(3) { + cache.insert(value) + } + cache.insert(values[0]) + cache.insert(values[3]) + + #expect(cache.count == 3) + #expect(!cache.contains(values[0])) + #expect(cache.contains(values[1])) + #expect(cache.contains(values[2])) + #expect(cache.contains(values[3])) + + cache.insert(values[4]) + + #expect(cache.count == 3) + #expect(!cache.contains(values[1])) + #expect(cache.contains(values[2])) + #expect(cache.contains(values[3])) + #expect(cache.contains(values[4])) + } + + @Test + func computesUltraVNCModularArithmeticKnownAnswers() { + #expect( + UltraVNCBigNum.addM64( + x: 0xffff_ffff_ffff_fffe, + y: 0xffff_ffff_ffff_fffd, + m: 0x61 + ) == 0x14 + ) + #expect( + UltraVNCBigNum.mulM64( + x: 0xffff_ffff_ffff_ffc5, + y: 0xffff_ffff_ffff_ffa3, + m: 0xffff_ffff_ffff_ff61 + ) == 0x19c8 + ) + #expect(UltraVNCBigNum.powM64(b: 4, e: 13, m: 497) == 445) + #expect( + UltraVNCBigNum.powM64( + b: 0xffff_ffff_ffff_ffc5, + e: 0x1_2345, + m: 0xffff_ffff_ffff_ff61 + ) == 0x34be_28a2_05bf_50b9 + ) + } + + @Test + func rejectsDegenerateUltraVNCKeyAgreementParameters() { + let eightBytes: (UInt64) -> Data = { value in + withUnsafeBytes(of: value.bigEndian) { Data($0) } + } + + for (generator, modulus, response) in [ + (2, 0, 3), + (2, 1, 3), + (1, 17, 3), + (17, 17, 3), + (2, 17, 1), + (2, 17, 17), + ] { + #expect( + VNCProtocol.UltraVNCMSLogonIIAuthentication.DiffieHellmanKeyAgreement( + generator: eightBytes(UInt64(generator)), + modulus: eightBytes(UInt64(modulus)), + resp: eightBytes(UInt64(response)) + ) == nil + ) + } + } + + @Test + func rejectsUltraVNCPMinusOneKeyAgreementElements() { + let modulus = ultraVNCValue(17) + + #expect( + UltraVNCKeyAgreement( + generator: ultraVNCValue(16), + modulus: modulus, + resp: ultraVNCValue(3) + ) == nil + ) + #expect( + UltraVNCKeyAgreement( + generator: ultraVNCValue(3), + modulus: modulus, + resp: ultraVNCValue(16) + ) == nil + ) + } + + @Test + func acceptsUltraVNCPMinusTwoKeyAgreementElements() { + let modulus = ultraVNCValue(17) + + #expect( + UltraVNCKeyAgreement( + generator: ultraVNCValue(15), + modulus: modulus, + resp: ultraVNCValue(3) + ) != nil + ) + #expect( + UltraVNCKeyAgreement( + generator: ultraVNCValue(3), + modulus: modulus, + resp: ultraVNCValue(15) + ) != nil + ) + } + + @Test + func encodesCharactersAsX11KeySyms() { + #expect(VNCKeyCode.withCharacter("A").map(\.rawValue) == [0x41]) + #expect(VNCKeyCode.withCharacter("é").map(\.rawValue) == [0xe9]) + #expect(VNCKeyCode.withCharacter("α").map(\.rawValue) == [0x0100_03b1]) + #expect(VNCKeyCode.withCharacter("🦀").map(\.rawValue) == [0x0101_f980]) + #expect( + VNCKeyCode.withCharacter("e\u{301}").map(\.rawValue) == [0x65, 0x0100_0301] + ) + } + + @Test + func mapsKeypadDecimalToTheDecimalKeysym() { + #expect(VNCKeyCode.ansiKeypadDecimal.rawValue == X11KeySymbols.XK_KP_Decimal) + #expect(VNCKeyCode.ansiKeypadDecimal.rawValue != X11KeySymbols.XK_KP_Separator) + + #if canImport(ObjectiveC) + #expect(_ObjC_VNCKeyCode.ansiKeypadDecimal == X11KeySymbols.XK_KP_Decimal) + #endif + } + + private func paddedARDValue(_ value: UInt8) -> Data { + Data(repeating: 0, count: 127) + Data([value]) + } + + private var appleRemoteDesktopSafePrime: Data { + hexadecimalData( + """ + C692B0343A9FC77AB54DD8F0912F24E657BACB3D4272E6525E624DCBAB26A479 + 904118111CCE782B6709522BD201F15C38EDF1B3E94DEAA7DEE91B4B4619607B + 3B76E1A1F9B65F6F545D42982FEE07F1F78D5855E9C490CAD9B45855F6BDEA7 + 5BF549643A572571B9F8073EE56A36DD1B9EAD50DCF444406BFDFD851DE76E51B + """ + ) + } + + private var appleRemoteDesktopUnsafePrime: Data { + hexadecimalData( + """ + F1EEAEF06F42BDFEF9524C7A03A6B26F074DC39F74F8C160BD15BA3869F54450 + CE55FD8DA6415AF88CEF7FFE7768BB1A061B7A3C0BCE0023B2C15C0A095D416B + E103EB8EE3BE0EE5874ADFE2BF7270B8719CC8F99B38BFFC126D6005DBEABAB + EE0037C10BAFB4D9CC864259DA28E1F5ECB949DCAC308512F9FA3E911F1E36061 + """ + ) + } + + private func hexadecimalData(_ value: String) -> Data { + let hex = value.filter(\.isHexDigit) + return Data( + stride(from: 0, to: hex.count, by: 2).compactMap { offset in + let start = hex.index(hex.startIndex, offsetBy: offset) + let end = hex.index(start, offsetBy: 2) + return UInt8(hex[start.. Data { + withUnsafeBytes(of: value.bigEndian) { Data($0) } + } +} diff --git a/macos/CrabfleetMac/scripts/build-app.sh b/macos/CrabfleetMac/scripts/build-app.sh index 9347e28c..b6456f24 100644 --- a/macos/CrabfleetMac/scripts/build-app.sh +++ b/macos/CrabfleetMac/scripts/build-app.sh @@ -35,6 +35,7 @@ contents_dir="$app_dir/Contents" macos_dir="$contents_dir/MacOS" resources_dir="$contents_dir/Resources" +rm -rf "$app_dir" mkdir -p "$macos_dir" "$resources_dir" install -m 755 "$build_dir/CrabfleetMac" "$macos_dir/CrabfleetMac" install -m 755 "$build_dir/libRoyalVNCKit.dylib" "$macos_dir/libRoyalVNCKit.dylib" diff --git a/migrations/0033_desktop_host_ownership.sql b/migrations/0033_desktop_host_ownership.sql new file mode 100644 index 00000000..92d18468 --- /dev/null +++ b/migrations/0033_desktop_host_ownership.sql @@ -0,0 +1,30 @@ +ALTER TABLE desktop_hosts + ADD COLUMN ownership_token TEXT NOT NULL DEFAULT ''; + +CREATE TRIGGER IF NOT EXISTS protect_token_owned_desktop_host_update +BEFORE UPDATE ON desktop_hosts +WHEN OLD.ownership_token <> '' + AND NEW.ownership_token = OLD.ownership_token + AND ( + NEW.owner_subject IS NOT OLD.owner_subject + OR NEW.id IS NOT OLD.id + OR NEW.owner IS NOT OLD.owner + OR NEW.name IS NOT OLD.name + OR NEW.address IS NOT OLD.address + OR NEW.port IS NOT OLD.port + OR NEW.created_at IS NOT OLD.created_at + OR NEW.updated_at IS NOT OLD.updated_at + ) +BEGIN + SELECT RAISE(IGNORE); +END; + +-- Token-aware workers replace the exact token with this transient marker and +-- delete it in the same atomic batch. Legacy workers can never create it. +CREATE TRIGGER IF NOT EXISTS protect_token_owned_desktop_host_delete +BEFORE DELETE ON desktop_hosts +WHEN OLD.ownership_token <> '' + AND OLD.ownership_token NOT GLOB 'delete-authorized:*' +BEGIN + SELECT RAISE(IGNORE); +END; diff --git a/migrations/0034_credential_policy_registration_staging.sql b/migrations/0034_credential_policy_registration_staging.sql new file mode 100644 index 00000000..53d62cc6 --- /dev/null +++ b/migrations/0034_credential_policy_registration_staging.sql @@ -0,0 +1,138 @@ +CREATE TABLE IF NOT EXISTS interactive_session_credential_policy_registrations ( + session_id TEXT NOT NULL, + sandbox_id TEXT NOT NULL, + state TEXT NOT NULL CHECK (state IN ('registering', 'cleanup_pending')), + registration_generation TEXT NOT NULL, + registration_claim TEXT, + registration_claim_expires_at INTEGER, + attempt_count INTEGER NOT NULL DEFAULT 0, + last_attempt_at INTEGER, + last_error TEXT, + cleanup_claim TEXT, + cleanup_claim_expires_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (session_id, sandbox_id), + CHECK ( + ( + state = 'registering' + AND registration_claim IS NOT NULL + AND registration_claim_expires_at IS NOT NULL + ) + OR ( + state = 'cleanup_pending' + AND registration_claim IS NULL + AND registration_claim_expires_at IS NULL + ) + ) +); + +CREATE INDEX IF NOT EXISTS idx_credential_policy_registration_cleanup + ON interactive_session_credential_policy_registrations( + state, + cleanup_claim_expires_at, + last_attempt_at + ); + +CREATE INDEX IF NOT EXISTS idx_credential_policy_registration_expiry + ON interactive_session_credential_policy_registrations( + state, + registration_claim_expires_at, + updated_at + ); + +-- Once a new worker stages a rotation, legacy workers must not claim, promote, +-- or remove the policy rows underneath its rollback snapshot. Expired claims +-- and abandoned cleanup rows eventually release the compatibility fence so a +-- rollback to legacy worker code cannot wedge this policy group forever. +CREATE TRIGGER IF NOT EXISTS fence_staged_credential_policy_insert +BEFORE INSERT ON interactive_session_credential_policies +WHEN NEW.state != 'cleanup_pending' + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = NEW.session_id + AND staged.sandbox_id = NEW.sandbox_id + AND staged.registration_generation != NEW.registration_generation + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + ) +BEGIN + SELECT RAISE(IGNORE); +END; + +CREATE TRIGGER IF NOT EXISTS fence_staged_credential_policy_update +BEFORE UPDATE ON interactive_session_credential_policies +WHEN NOT (OLD.state != 'cleanup_pending' AND NEW.state = 'cleanup_pending') + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = NEW.session_id + AND staged.sandbox_id = NEW.sandbox_id + AND staged.registration_generation != NEW.registration_generation + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + ) +BEGIN + SELECT RAISE(IGNORE); +END; + +CREATE TRIGGER IF NOT EXISTS fence_staged_credential_policy_delete +BEFORE DELETE ON interactive_session_credential_policies +WHEN EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = OLD.session_id + AND staged.sandbox_id = OLD.sandbox_id + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) +) +BEGIN + SELECT RAISE(IGNORE); +END; + +-- Legacy workers renew registration claims only in the policy table. Leaving +-- those claims unstaged prevents the new scanner from recovering a stale +-- snapshot while the legacy registration is still live. diff --git a/migrations/0035_credential_policy_registration_rollback.sql b/migrations/0035_credential_policy_registration_rollback.sql new file mode 100644 index 00000000..8af5f195 --- /dev/null +++ b/migrations/0035_credential_policy_registration_rollback.sql @@ -0,0 +1,2 @@ +ALTER TABLE interactive_session_credential_policy_registrations + ADD COLUMN rollback_policies_json TEXT; diff --git a/migrations/0036_credential_policy_lookup_repair.sql b/migrations/0036_credential_policy_lookup_repair.sql new file mode 100644 index 00000000..50325c7c --- /dev/null +++ b/migrations/0036_credential_policy_lookup_repair.sql @@ -0,0 +1,82 @@ +ALTER TABLE interactive_session_credential_policy_registrations + ADD COLUMN repair_generation TEXT; + +DROP TRIGGER IF EXISTS fence_staged_credential_policy_insert; +DROP TRIGGER IF EXISTS fence_staged_credential_policy_update; + +CREATE TRIGGER fence_staged_credential_policy_insert +BEFORE INSERT ON interactive_session_credential_policies +WHEN NEW.state != 'cleanup_pending' + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = NEW.session_id + AND staged.sandbox_id = NEW.sandbox_id + AND staged.registration_generation != NEW.registration_generation + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + AND NOT ( + staged.state = 'registering' + AND staged.repair_generation = NEW.registration_generation + AND staged.registration_claim = NEW.registration_claim + AND staged.registration_claim_expires_at = NEW.registration_claim_expires_at + AND NEW.state = 'registering' + ) + ) +BEGIN + SELECT RAISE(IGNORE); +END; + +CREATE TRIGGER fence_staged_credential_policy_update +BEFORE UPDATE ON interactive_session_credential_policies +WHEN NOT (OLD.state != 'cleanup_pending' AND NEW.state = 'cleanup_pending') + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = NEW.session_id + AND staged.sandbox_id = NEW.sandbox_id + AND staged.registration_generation != NEW.registration_generation + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + AND NOT ( + staged.state = 'registering' + AND staged.repair_generation = NEW.registration_generation + AND staged.registration_claim = OLD.registration_claim + AND staged.registration_claim_expires_at = OLD.registration_claim_expires_at + AND OLD.state = 'registering' + AND NEW.state = 'active' + AND NEW.registration_claim IS NULL + AND NEW.registration_claim_expires_at IS NULL + ) + ) +BEGIN + SELECT RAISE(IGNORE); +END; diff --git a/migrations/0037_credential_policy_registration_lookup_ids.sql b/migrations/0037_credential_policy_registration_lookup_ids.sql new file mode 100644 index 00000000..ddd15f17 --- /dev/null +++ b/migrations/0037_credential_policy_registration_lookup_ids.sql @@ -0,0 +1,50 @@ +ALTER TABLE interactive_session_credential_policy_registrations + ADD COLUMN lookup_ids_json TEXT; + +-- Pre-migration staging rows may already have registered a new-generation +-- Durable Object lookup that D1 cannot reconstruct. Keep those rows nullable +-- so recovery derives the current compatibility lookup set at runtime. + +DROP TRIGGER IF EXISTS fence_staged_credential_policy_delete; + +CREATE TRIGGER fence_staged_credential_policy_delete +BEFORE DELETE ON interactive_session_credential_policies +WHEN EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = OLD.session_id + AND staged.sandbox_id = OLD.sandbox_id + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + AND NOT ( + staged.state = 'registering' + AND staged.repair_generation = OLD.registration_generation + AND staged.registration_claim = OLD.cleanup_claim + AND staged.registration_claim_expires_at = OLD.cleanup_claim_expires_at + AND OLD.state = 'cleanup_pending' + AND json_valid(staged.lookup_ids_json) + AND NOT EXISTS ( + SELECT 1 + FROM json_each(staged.lookup_ids_json) AS current_lookup + WHERE current_lookup.type = 'text' + AND current_lookup.value = OLD.lookup_id + ) + ) +) +BEGIN + SELECT RAISE(IGNORE); +END; diff --git a/migrations/0037_runtime_adapter_workspace_cleanup.sql b/migrations/0037_runtime_adapter_workspace_cleanup.sql new file mode 100644 index 00000000..866cb03a --- /dev/null +++ b/migrations/0037_runtime_adapter_workspace_cleanup.sql @@ -0,0 +1,30 @@ +CREATE TABLE IF NOT EXISTS runtime_adapter_workspace_cleanups ( + session_id TEXT NOT NULL, + adapter_workspace_id TEXT NOT NULL, + profile TEXT, + control_plane TEXT, + create_pending INTEGER NOT NULL CHECK (create_pending IN (0, 1)), + message TEXT NOT NULL, + reconcile_error TEXT, + attempt_count INTEGER NOT NULL DEFAULT 0, + last_attempt_at INTEGER, + next_attempt_at INTEGER NOT NULL, + cleanup_claim TEXT, + cleanup_claim_expires_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (session_id, adapter_workspace_id), + CHECK ( + (cleanup_claim IS NULL AND cleanup_claim_expires_at IS NULL) + OR (cleanup_claim IS NOT NULL AND cleanup_claim_expires_at IS NOT NULL) + ) +); + +CREATE INDEX IF NOT EXISTS idx_runtime_adapter_workspace_cleanup_due + ON runtime_adapter_workspace_cleanups( + next_attempt_at, + cleanup_claim_expires_at, + updated_at, + session_id, + adapter_workspace_id + ); diff --git a/migrations/0038_desktop_host_publication_identity.sql b/migrations/0038_desktop_host_publication_identity.sql new file mode 100644 index 00000000..52f1b1fd --- /dev/null +++ b/migrations/0038_desktop_host_publication_identity.sql @@ -0,0 +1,25 @@ +ALTER TABLE desktop_hosts + ADD COLUMN publication_id TEXT NOT NULL DEFAULT ''; + +ALTER TABLE desktop_hosts + ADD COLUMN publication_write_token TEXT NOT NULL DEFAULT ''; + +-- New workers rotate publication_write_token with ownership_token. Older +-- token-aware workers leave it unchanged, which identifies a token-only write +-- that must invalidate recovery authority for the previous publisher. +CREATE TRIGGER IF NOT EXISTS clear_stale_desktop_host_publication_identity +AFTER UPDATE ON desktop_hosts +WHEN OLD.publication_id <> '' + AND NEW.ownership_token <> OLD.ownership_token + AND NEW.publication_id = OLD.publication_id + AND NEW.publication_write_token = OLD.publication_write_token +BEGIN + UPDATE desktop_hosts + SET publication_id = '', + publication_write_token = '' + WHERE owner_subject = NEW.owner_subject + AND id = NEW.id + AND ownership_token = NEW.ownership_token + AND publication_id = NEW.publication_id + AND publication_write_token = NEW.publication_write_token; +END; diff --git a/migrations/0039_runtime_adapter_cleanup_deletion_observed.sql b/migrations/0039_runtime_adapter_cleanup_deletion_observed.sql new file mode 100644 index 00000000..e93c3b52 --- /dev/null +++ b/migrations/0039_runtime_adapter_cleanup_deletion_observed.sql @@ -0,0 +1,3 @@ +ALTER TABLE runtime_adapter_workspace_cleanups + ADD COLUMN deletion_observed INTEGER NOT NULL DEFAULT 0 + CHECK (deletion_observed IN (0, 1)); diff --git a/migrations/0040_credential_policy_registration_write_fence.sql b/migrations/0040_credential_policy_registration_write_fence.sql new file mode 100644 index 00000000..694384f1 --- /dev/null +++ b/migrations/0040_credential_policy_registration_write_fence.sql @@ -0,0 +1,174 @@ +ALTER TABLE interactive_session_credential_policy_registrations + ADD COLUMN registration_write_started INTEGER NOT NULL DEFAULT 0 + CHECK (registration_write_started IN (0, 1)); + +-- Existing staged rows may already have written a replacement generation to +-- the Durable Object. Conservatively retain their legacy-writer fence until +-- current recovery either completes or removes the staged registration. +UPDATE interactive_session_credential_policy_registrations +SET registration_write_started = 1; + +DROP TRIGGER IF EXISTS fence_staged_credential_policy_insert; +DROP TRIGGER IF EXISTS fence_staged_credential_policy_update; +DROP TRIGGER IF EXISTS fence_staged_credential_policy_delete; + +CREATE TRIGGER fence_staged_credential_policy_insert +BEFORE INSERT ON interactive_session_credential_policies +WHEN EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = NEW.session_id + AND staged.sandbox_id = NEW.sandbox_id + AND staged.registration_generation != NEW.registration_generation + AND ( + ( + staged.registration_write_started = 1 + AND NOT ( + staged.state = 'cleanup_pending' + AND NEW.state = 'cleanup_pending' + ) + ) + OR ( + NEW.state != 'cleanup_pending' + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + ) + ) + AND NOT ( + staged.state = 'registering' + AND staged.repair_generation = NEW.registration_generation + AND staged.registration_claim = NEW.registration_claim + AND staged.registration_claim_expires_at = NEW.registration_claim_expires_at + AND NEW.state = 'registering' + ) +) +BEGIN + SELECT RAISE(IGNORE); +END; + +CREATE TRIGGER fence_staged_credential_policy_update +BEFORE UPDATE ON interactive_session_credential_policies +WHEN EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = NEW.session_id + AND staged.sandbox_id = NEW.sandbox_id + AND staged.registration_generation != NEW.registration_generation + AND ( + ( + staged.registration_write_started = 1 + AND NOT ( + ( + staged.state = 'cleanup_pending' + AND OLD.state != 'cleanup_pending' + AND NEW.state = 'cleanup_pending' + ) + OR ( + staged.state = 'registering' + AND staged.repair_generation = OLD.registration_generation + AND NEW.registration_generation = OLD.registration_generation + AND staged.registration_claim = NEW.cleanup_claim + AND staged.registration_claim_expires_at = NEW.cleanup_claim_expires_at + AND OLD.state = 'active' + AND NEW.state = 'cleanup_pending' + AND json_valid(staged.lookup_ids_json) + AND NOT EXISTS ( + SELECT 1 + FROM json_each(staged.lookup_ids_json) AS current_lookup + WHERE current_lookup.type = 'text' + AND current_lookup.value = OLD.lookup_id + ) + ) + ) + ) + OR ( + NOT (OLD.state != 'cleanup_pending' AND NEW.state = 'cleanup_pending') + AND ( + ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + ) + ) + AND NOT ( + staged.state = 'registering' + AND staged.repair_generation = NEW.registration_generation + AND staged.registration_claim = OLD.registration_claim + AND staged.registration_claim_expires_at = OLD.registration_claim_expires_at + AND OLD.state = 'registering' + AND NEW.state = 'active' + AND NEW.registration_claim IS NULL + AND NEW.registration_claim_expires_at IS NULL + ) +) +BEGIN + SELECT RAISE(IGNORE); +END; + +CREATE TRIGGER fence_staged_credential_policy_delete +BEFORE DELETE ON interactive_session_credential_policies +WHEN EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = OLD.session_id + AND staged.sandbox_id = OLD.sandbox_id + AND ( + staged.registration_write_started = 1 + OR ( + staged.state = 'registering' + AND staged.registration_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + ) + OR ( + staged.state = 'cleanup_pending' + AND ( + staged.cleanup_claim_expires_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 + OR staged.updated_at > + CAST(strftime('%s', 'now') AS INTEGER) * 1000 - 300000 + ) + ) + ) + AND NOT ( + staged.state = 'registering' + AND staged.repair_generation = OLD.registration_generation + AND staged.registration_claim = OLD.cleanup_claim + AND staged.registration_claim_expires_at = OLD.cleanup_claim_expires_at + AND OLD.state = 'cleanup_pending' + AND json_valid(staged.lookup_ids_json) + AND NOT EXISTS ( + SELECT 1 + FROM json_each(staged.lookup_ids_json) AS current_lookup + WHERE current_lookup.type = 'text' + AND current_lookup.value = OLD.lookup_id + ) + ) +) +BEGIN + SELECT RAISE(IGNORE); +END; diff --git a/migrations/0041_desktop_host_ownership_errors.sql b/migrations/0041_desktop_host_ownership_errors.sql new file mode 100644 index 00000000..590da19e --- /dev/null +++ b/migrations/0041_desktop_host_ownership_errors.sql @@ -0,0 +1,28 @@ +DROP TRIGGER IF EXISTS protect_token_owned_desktop_host_update; +DROP TRIGGER IF EXISTS protect_token_owned_desktop_host_delete; + +CREATE TRIGGER protect_token_owned_desktop_host_update +BEFORE UPDATE ON desktop_hosts +WHEN OLD.ownership_token <> '' + AND NEW.ownership_token = OLD.ownership_token + AND ( + NEW.owner_subject IS NOT OLD.owner_subject + OR NEW.id IS NOT OLD.id + OR NEW.owner IS NOT OLD.owner + OR NEW.name IS NOT OLD.name + OR NEW.address IS NOT OLD.address + OR NEW.port IS NOT OLD.port + OR NEW.created_at IS NOT OLD.created_at + OR NEW.updated_at IS NOT OLD.updated_at + ) +BEGIN + SELECT RAISE(ABORT, 'token-owned desktop host update requires ownership token'); +END; + +CREATE TRIGGER protect_token_owned_desktop_host_delete +BEFORE DELETE ON desktop_hosts +WHEN OLD.ownership_token <> '' + AND OLD.ownership_token NOT GLOB 'delete-authorized:*' +BEGIN + SELECT RAISE(ABORT, 'token-owned desktop host delete requires ownership token'); +END; diff --git a/src/app/app-navigation.js b/src/app/app-navigation.js index b64f73de..403345d8 100644 --- a/src/app/app-navigation.js +++ b/src/app/app-navigation.js @@ -1,5 +1,5 @@ import { useEffect, useRef, useState } from "preact/hooks"; -import { appViewUrl, initialAppView, sessionRouteUrl } from "./routing.js"; +import { appViewUrl, initialAppView, parseSessionLink, sessionRouteUrl } from "./routing.js"; import { loadSessionLayout, saveSessionLayout } from "./session-layout.js"; import { disposeAllTerminals, warmGhosttyModule } from "./terminal.js"; @@ -27,6 +27,21 @@ export function sessionOpenTarget(id, currentId, sessionItemById, options = {}) }; } +export function appNavigationLocationState(locationLike = location) { + const sessionLink = parseSessionLink(locationLike); + return { + appView: initialAppView(locationLike), + drawers: sessionLink.route ? { sessions: true } : {}, + focusedSessionId: sessionLink.id, + sharedSessionId: sessionLink.id, + sharedToken: sessionLink.token, + }; +} + +export function shouldDisposeTerminalsForNavigation(state) { + return !state.focusedSessionId && !state.drawers.sessions; +} + export function useAppNavigation({ initialSessionLink, sessionItemByIdRef }) { const [appView, setAppViewState] = useState(initialAppView); const [drawers, setDrawers] = useState(initialSessionLink.route ? { sessions: true } : {}); @@ -44,7 +59,18 @@ export function useAppNavigation({ initialSessionLink, sessionItemByIdRef }) { focusedSessionIdRef.current = focusedSessionId; useEffect(() => { - const onPopState = () => setAppViewState(initialAppView()); + const onPopState = () => { + const next = appNavigationLocationState(); + setAppViewState(next.appView); + setDrawers(next.drawers); + setActiveRunId(null); + setFocusedSessionId(next.focusedSessionId); + focusedSessionIdRef.current = next.focusedSessionId; + setSharedSessionId(next.sharedSessionId); + setSharedToken(next.sharedToken); + if (next.focusedSessionId) warmGhosttyModule(); + else if (shouldDisposeTerminalsForNavigation(next)) disposeAllTerminals(); + }; window.addEventListener("popstate", onPopState); return () => window.removeEventListener("popstate", onPopState); }, []); diff --git a/src/app/routing.js b/src/app/routing.js index a501b22c..bf47edfe 100644 --- a/src/app/routing.js +++ b/src/app/routing.js @@ -2,9 +2,17 @@ export const loginReturnKey = "crabbox-login-return"; export function parseSessionLink(locationLike = location) { const match = locationLike.pathname.match(/^\/(?:app\/)?sessions(?:\/([^/]+))?\/?$/); + let id = null; + if (match?.[1]) { + try { + id = decodeURIComponent(match[1]); + } catch { + return { route: false, id: null, token: null }; + } + } return { route: Boolean(match), - id: match?.[1] ? decodeURIComponent(match[1]) : null, + id, token: new URLSearchParams(locationLike.search).get("token"), }; } diff --git a/src/credential-policy-fence.ts b/src/credential-policy-fence.ts index 2697e8f4..f70ee05c 100644 --- a/src/credential-policy-fence.ts +++ b/src/credential-policy-fence.ts @@ -11,6 +11,11 @@ export type CredentialPolicyGenerationTombstone = { tombstonedAt: number; }; +export type CredentialPolicyRollbackRecord = { + generation: string; + policy: T; +}; + export function isCurrentCredentialPolicyGeneration(value: unknown): value is string { return ( typeof value === "string" && @@ -20,6 +25,31 @@ export function isCurrentCredentialPolicyGeneration(value: unknown): value is st ); } +export function credentialPolicyRollbackRecord( + value: unknown, +): CredentialPolicyRollbackRecord | undefined { + if (!value || typeof value !== "object") return undefined; + const record = value as Partial>; + if ( + !isCurrentCredentialPolicyGeneration(record.generation) || + !record.policy || + typeof record.policy !== "object" || + typeof record.policy.sessionId !== "string" || + !record.policy.sessionId + ) { + return undefined; + } + return record as CredentialPolicyRollbackRecord; +} + +export function credentialPolicyRollbackExpiresAt( + replacedRegistrationExpiresAt: number, + now: number, + claimTtlMs: number, +): number { + return Math.max(replacedRegistrationExpiresAt + 1, now + claimTtlMs); +} + export function credentialPolicyRegistrationAccepted( current: CredentialPolicyGenerationRecord | undefined, tombstone: CredentialPolicyGenerationTombstone | undefined, @@ -31,12 +61,12 @@ export function credentialPolicyRegistrationAccepted current.registrationExpiresAt; + } if (current.registrationClaim === incoming.registrationClaim) { return incoming.registrationExpiresAt >= current.registrationExpiresAt; } diff --git a/src/fleet-state.ts b/src/fleet-state.ts index 543db5a4..16e4dd6b 100644 --- a/src/fleet-state.ts +++ b/src/fleet-state.ts @@ -301,13 +301,13 @@ export function fleetSessionSummary( terminalCapable && (session.ptyAvailable === true || (session.canControl !== false && - (session.runtime === "github_actions" || - (session.ptyAvailable ?? - Boolean( - ptyRouteKind(session, { - sandboxAvailable: options.sandboxAvailable, - }), - ))))) && + session.runtime !== "github_actions" && + (session.ptyAvailable ?? + Boolean( + ptyRouteKind(session, { + sandboxAvailable: options.sandboxAvailable, + }), + )))) && ptyReadyStatuses.has(session.status), vnc: !inactiveStatuses.has(session.status) && diff --git a/src/github-actions-runner.ts b/src/github-actions-runner.ts new file mode 100644 index 00000000..48dca9dd --- /dev/null +++ b/src/github-actions-runner.ts @@ -0,0 +1,194 @@ +import { + encodeGitHubActionsRelayOutput, + parseGitHubActionsRelayInput, + sendGitHubActionsRelayInputAcknowledgement, + type GitHubActionsRelaySocket, +} from "./github-actions-runtime.ts"; + +const runnerInputQueueMaxBytes = 16 * 1024 * 1024; +const runnerInputQueueMaxFrames = 32; +const runnerInputQueueMaxAgeMs = 5_000; +const runnerInputWriteTimeoutMs = 5_000; +const runnerInputBacklogError = "GitHub Actions runner input backlog exceeded"; +const runnerInputExpiredError = "GitHub Actions runner input expired"; +const runnerInputGenerationError = "GitHub Actions runner generation changed"; +const runnerInputWriteTimeoutError = "GitHub Actions runner input write timed out"; + +type RunnerInputQueue = { + bytes: number; + frames: number; + generation: string | undefined; + retired: boolean; + tail: Promise; +}; + +const runnerInputQueues = new WeakMap(); + +export function sendGitHubActionsRunnerOutput( + socket: GitHubActionsRelaySocket, + output: string | ArrayBuffer | ArrayBufferView, +): void { + socket.send(encodeGitHubActionsRelayOutput(output)); +} + +export function acceptGitHubActionsRunnerInput( + socket: GitHubActionsRelaySocket, + message: string | ArrayBuffer, + writeToPty: (payload: ArrayBuffer) => void | Promise, + now: () => number = Date.now, + writeTimeoutMs: number = runnerInputWriteTimeoutMs, +): Promise { + const input = parseGitHubActionsRelayInput(message); + if (!input) return Promise.resolve(false); + + const existingQueue = runnerInputQueues.get(socket); + if (existingQueue?.retired || socket.readyState !== WebSocket.OPEN) { + return Promise.resolve(true); + } + if (existingQueue && existingQueue.generation !== input.generation) { + existingQueue.retired = true; + sendRunnerInputAcknowledgement( + socket, + input.inputId, + input.generation, + false, + runnerInputGenerationError, + ); + closeRunnerInputSocket(socket, runnerInputGenerationError); + return Promise.resolve(true); + } + + const queue = + existingQueue ?? + ({ + bytes: 0, + frames: 0, + generation: input.generation, + retired: false, + tail: Promise.resolve(), + } satisfies RunnerInputQueue); + if (!existingQueue) runnerInputQueues.set(socket, queue); + + if ( + queue.frames >= runnerInputQueueMaxFrames || + queue.bytes + input.payload.byteLength > runnerInputQueueMaxBytes + ) { + sendRunnerInputAcknowledgement( + socket, + input.inputId, + input.generation, + false, + runnerInputBacklogError, + ); + return Promise.resolve(true); + } + + const queuedAt = now(); + queue.frames += 1; + queue.bytes += input.payload.byteLength; + const queued = queue.tail + .catch(() => undefined) + .then(async () => { + if (!isActiveRunnerInputQueue(socket, queue)) return; + if (now() - queuedAt >= runnerInputQueueMaxAgeMs) { + sendRunnerInputAcknowledgement( + socket, + input.inputId, + input.generation, + false, + runnerInputExpiredError, + ); + return; + } + const writeResult = await writeRunnerInputWithTimeout( + () => writeToPty(input.payload), + writeTimeoutMs, + ); + if (writeResult === "timed-out") { + queue.retired = true; + closeRunnerInputSocket(socket, runnerInputWriteTimeoutError); + return; + } + if (writeResult === "accepted") { + if (isActiveRunnerInputQueue(socket, queue)) { + sendRunnerInputAcknowledgement(socket, input.inputId, input.generation, true); + } + } else { + if (isActiveRunnerInputQueue(socket, queue)) { + sendRunnerInputAcknowledgement(socket, input.inputId, input.generation, false); + } + } + }) + .finally(() => { + queue.frames -= 1; + queue.bytes -= input.payload.byteLength; + }); + queue.tail = queued; + return queued.then(() => true); +} + +function writeRunnerInputWithTimeout( + write: () => void | Promise, + timeoutMs: number, +): Promise<"accepted" | "rejected" | "timed-out"> { + return new Promise((resolve) => { + let settled = false; + const finish = (result: "accepted" | "rejected" | "timed-out") => { + if (settled) return; + settled = true; + clearTimeout(timeout); + resolve(result); + }; + const timeout = setTimeout(() => finish("timed-out"), timeoutMs); + let result: void | Promise; + try { + result = write(); + } catch { + finish("rejected"); + return; + } + Promise.resolve(result).then( + () => finish("accepted"), + () => finish("rejected"), + ); + }); +} + +function sendRunnerInputAcknowledgement( + socket: GitHubActionsRelaySocket, + inputId: string, + generation: string | undefined, + accepted: boolean, + error?: string, +): void { + sendGitHubActionsRelayInputAcknowledgement(socket, { + inputId, + accepted, + ...(error ? { error } : {}), + ...(generation ? { generation } : {}), + }); +} + +function closeRunnerInputSocket(socket: GitHubActionsRelaySocket, reason: string): void { + if (socket.readyState !== WebSocket.OPEN) return; + try { + socket.close(1012, reason); + } catch { + // Queue retirement still prevents later writes when the socket cannot be closed cleanly. + } +} + +function isActiveRunnerInputQueue( + socket: GitHubActionsRelaySocket, + queue: RunnerInputQueue, +): boolean { + if ( + runnerInputQueues.get(socket) !== queue || + queue.retired || + socket.readyState !== WebSocket.OPEN + ) { + queue.retired = true; + return false; + } + return true; +} diff --git a/src/github-actions-runtime.ts b/src/github-actions-runtime.ts index 454f722d..fdcaa8bf 100644 --- a/src/github-actions-runtime.ts +++ b/src/github-actions-runtime.ts @@ -14,8 +14,36 @@ export type GitHubActionsRelaySocket = { readyState: number; send(message: string | ArrayBuffer): void; close(code?: number, reason?: string): void; + serializeAttachment?(attachment: unknown): void; + deserializeAttachment?(): unknown; }; +export type GitHubActionsRelayInputAcknowledgement = { + accepted: boolean; + error?: string; + generation?: string; + inputId: string; +}; + +export type GitHubActionsRelayInput = { + generation?: string; + inputId: string; + payload: ArrayBuffer; +}; + +export const githubActionsFramedRunnerCapability = "cfr1-framed-io-v1"; +export const githubActionsGenerationFencedCapability = "cfr1-framed-io-v2"; +export const githubActionsRunnerProtocolHeader = "sec-websocket-protocol"; +export const githubActionsRunnerProtocolQuery = "runnerProtocol"; +export const githubActionsViewerProtocolQuery = "viewerProtocol"; +export const githubActionsViewerProtocolHeader = "x-crabfleet-viewer-protocol"; +export const githubActionsViewerGenerationHeader = "x-crabfleet-runner-generation"; +export type GitHubActionsRelayProtocol = + | typeof githubActionsFramedRunnerCapability + | typeof githubActionsGenerationFencedCapability; +export type GitHubActionsRunnerProtocol = GitHubActionsRelayProtocol; +export type GitHubActionsViewerProtocol = GitHubActionsRelayProtocol; + export const githubActionsCapabilities = { terminal: true, takeover: true, @@ -42,6 +70,39 @@ const terminalWorkStates = new Set([ ]); const webSocketOpen = 1; +const relayInputRejectedError = "GitHub Actions runner did not accept terminal input"; +const relayFrameMagic = new Uint8Array([0x43, 0x46, 0x52, 0x31]); +const relayFrameHeaderBytes = relayFrameMagic.byteLength + 2; +const relayInputFrameType = 1; +const relayInputAcknowledgementFrameType = 2; +const relayEventFrameType = 3; +const relayOutputFrameType = 4; +const relayGenerationInputFrameType = 5; +const relayGenerationInputAcknowledgementFrameType = 6; +const relayGenerationEventFrameType = 7; +const relayInputIdMaximumBytes = 80; +const relayInputIdPattern = /^[A-Za-z0-9_-]+$/; +const relayGenerationMaximumBytes = 80; +const relayGenerationPattern = /^[A-Za-z0-9_-]+$/; +export const githubActionsLegacyRelayGeneration = "legacy"; +const relayEventCodes = { + runner_connected: 1, + runner_disconnected: 2, + runner_waiting: 3, +} as const; +const relayEvents = new Map( + Object.entries(relayEventCodes).map(([event, code]) => [ + code, + event as keyof typeof relayEventCodes, + ]), +); +const encoder = new TextEncoder(); +const decoder = new TextDecoder(); + +type GitHubActionsRelayAttachment = { + generation?: string; + protocol?: GitHubActionsRelayProtocol; +}; export function githubActionsRuntimeLabel(runtime: unknown): string { return runtime === githubActionsRuntime ? "GitHub Actions" : ""; @@ -61,6 +122,28 @@ export function buildGitHubActionsRunnerPtyUrl( return url.toString(); } +export function buildGitHubActionsViewerRelayUrl(): string { + const url = new URL("https://crabfleet.internal/api/session-control/github-actions/viewer"); + url.searchParams.set(githubActionsViewerProtocolQuery, githubActionsGenerationFencedCapability); + return url.toString(); +} + +export function gitHubActionsViewerResponseUsesFramedProtocol(response: Response): boolean { + return isGitHubActionsRelayProtocol(response.headers.get(githubActionsViewerProtocolHeader)); +} + +export function gitHubActionsViewerResponseUsesGenerations(response: Response): boolean { + return ( + response.headers.get(githubActionsViewerProtocolHeader) === + githubActionsGenerationFencedCapability + ); +} + +export function gitHubActionsViewerResponseGeneration(response: Response): string | null { + const generation = response.headers.get(githubActionsViewerGenerationHeader); + return isGitHubActionsRelayGeneration(generation) ? generation : null; +} + export function parseGitHubActionsWorkState(value: unknown): GitHubActionsWorkState | null { const state = String(value ?? "").trim() as GitHubActionsWorkState; return workStates.has(state) ? state : null; @@ -110,17 +193,326 @@ export function forwardGitHubActionsRelayMessage( runners: readonly GitHubActionsRelaySocket[], viewers: readonly GitHubActionsRelaySocket[], ): number { - if (sender === "viewer" && isGitHubActionsViewerControlMessage(message)) return 0; - const targets = sender === "runner" ? viewers : runners.slice(0, 1); + const targets = + sender === "runner" + ? viewers + : runners.filter((socket) => socket.readyState === webSocketOpen).slice(0, 1); let forwarded = 0; for (const socket of targets) { if (socket.readyState !== webSocketOpen) continue; - socket.send(message); - forwarded += 1; + try { + socket.send(message); + forwarded += 1; + } catch { + // The caller uses the forwarded count to reject undelivered viewer input. + } + } + return forwarded; +} + +export function relayGitHubActionsWebSocketMessage( + sender: GitHubActionsRelayRole, + senderSocket: GitHubActionsRelaySocket, + message: string | ArrayBuffer, + runners: readonly GitHubActionsRelaySocket[], + viewers: readonly GitHubActionsRelaySocket[], +): number { + if (sender === "viewer") { + if (isGitHubActionsViewerControlMessage(message)) return 0; + const framedViewer = gitHubActionsViewerUsesFramedProtocol(senderSocket); + const input = framedViewer ? parseGitHubActionsRelayInput(message) : null; + if (framedViewer && !input) return 0; + const runner = runners.find((socket) => socket.readyState === webSocketOpen); + if (!runner) { + sendGitHubActionsViewerInputAcknowledgement( + senderSocket, + input?.inputId ?? null, + false, + input?.generation, + ); + return 0; + } + const generation = gitHubActionsRelayGeneration(runner) ?? githubActionsLegacyRelayGeneration; + if ( + gitHubActionsRelayUsesGenerations(senderSocket) && + (!input?.generation || input.generation !== generation) + ) { + sendGitHubActionsViewerInputAcknowledgement( + senderSocket, + input?.inputId ?? null, + false, + input?.generation, + ); + return 0; + } + const framed = gitHubActionsRunnerUsesFramedProtocol(runner); + try { + if (framed) { + runner.send( + encodeGitHubActionsRelayInput( + input?.inputId ?? createGitHubActionsRelayInputId(), + input?.payload ?? message, + gitHubActionsRelayUsesGenerations(runner) ? generation : undefined, + ), + ); + } else { + runner.send(framedViewer ? input!.payload : message); + } + if (!framedViewer || !framed) { + sendGitHubActionsViewerInputAcknowledgement( + senderSocket, + input?.inputId ?? null, + true, + input?.generation, + ); + } + return 1; + } catch { + sendGitHubActionsViewerInputAcknowledgement( + senderSocket, + input?.inputId ?? null, + false, + input?.generation, + ); + return 0; + } + } + + if (runners.find((socket) => socket.readyState === webSocketOpen) !== senderSocket) { + return 0; + } + + if (gitHubActionsRunnerUsesFramedProtocol(senderSocket)) { + const acknowledgement = parseGitHubActionsRelayInputAcknowledgement(message); + const output = parseGitHubActionsRelayOutput(message); + if (!acknowledgement && !output) return 0; + const generation = gitHubActionsRelayGeneration(senderSocket); + if ( + acknowledgement && + gitHubActionsRelayUsesGenerations(senderSocket) && + acknowledgement.generation !== generation + ) { + return 0; + } + let forwarded = 0; + for (const viewer of viewers) { + if (viewer.readyState !== webSocketOpen) continue; + if (acknowledgement && !gitHubActionsViewerUsesFramedProtocol(viewer)) continue; + try { + viewer.send( + acknowledgement + ? encodeGitHubActionsRelayInputAcknowledgement({ + inputId: acknowledgement.inputId, + accepted: acknowledgement.accepted, + ...(acknowledgement.error ? { error: acknowledgement.error } : {}), + ...(gitHubActionsRelayUsesGenerations(viewer) && generation ? { generation } : {}), + }) + : gitHubActionsViewerUsesFramedProtocol(viewer) + ? message + : output!, + ); + forwarded += 1; + } catch { + // A failed viewer does not prevent delivery to the remaining viewers. + } + } + return forwarded; + } + + let forwarded = 0; + for (const viewer of viewers) { + if (viewer.readyState !== webSocketOpen) continue; + try { + viewer.send( + gitHubActionsViewerUsesFramedProtocol(viewer) + ? encodeGitHubActionsRelayOutput(message) + : message, + ); + forwarded += 1; + } catch { + // A failed viewer does not prevent delivery to the remaining viewers. + } } return forwarded; } +export function sendGitHubActionsRelayInputAcknowledgement( + viewer: GitHubActionsRelaySocket, + acknowledgement: GitHubActionsRelayInputAcknowledgement, +): boolean { + if (viewer.readyState !== webSocketOpen) return false; + try { + viewer.send(encodeGitHubActionsRelayInputAcknowledgement(acknowledgement)); + return true; + } catch { + return false; + } +} + +export function parseGitHubActionsRelayInputAcknowledgement( + message: string | ArrayBuffer, +): GitHubActionsRelayInputAcknowledgement | null { + const legacyFrame = decodeGitHubActionsRelayFrame(message, relayInputAcknowledgementFrameType); + const generatedFrame = decodeGitHubActionsRelayFrame( + message, + relayGenerationInputAcknowledgementFrameType, + ); + const decoded = generatedFrame + ? decodeGitHubActionsRelayGeneration(generatedFrame.payload) + : null; + const frame = legacyFrame ?? (decoded ? { ...generatedFrame!, payload: decoded.payload } : null); + if (!frame?.inputId || frame.payload.byteLength < 1) return null; + const acceptedByte = frame.payload[0]; + if (acceptedByte !== 0 && acceptedByte !== 1) return null; + const accepted = acceptedByte === 1; + if (accepted) { + return { + inputId: frame.inputId, + accepted: true, + ...(decoded ? { generation: decoded.generation } : {}), + }; + } + const error = decoder.decode(frame.payload.subarray(1)).trim(); + return { + inputId: frame.inputId, + accepted: false, + error: error || relayInputRejectedError, + ...(decoded ? { generation: decoded.generation } : {}), + }; +} + +export function createGitHubActionsRelayInputId(): string { + return crypto.randomUUID().replaceAll("-", ""); +} + +export function encodeGitHubActionsRelayInput( + inputId: string, + payload: string | ArrayBuffer | ArrayBufferView, + generation?: string, +): ArrayBuffer { + requireGitHubActionsRelayInputId(inputId); + return encodeGitHubActionsRelayFrame( + generation ? relayGenerationInputFrameType : relayInputFrameType, + inputId, + generation + ? encodeGitHubActionsRelayGeneration(generation, messageBytes(payload)) + : messageBytes(payload), + ); +} + +export function encodeGitHubActionsRelayOutput( + payload: string | ArrayBuffer | ArrayBufferView, +): ArrayBuffer { + return encodeGitHubActionsRelayFrame(relayOutputFrameType, "", messageBytes(payload)); +} + +export function parseGitHubActionsRelayOutput(message: string | ArrayBuffer): ArrayBuffer | null { + const frame = decodeGitHubActionsRelayFrame(message, relayOutputFrameType); + if (!frame || frame.inputId) return null; + return Uint8Array.from(frame.payload).buffer; +} + +export function parseGitHubActionsRelayInput( + message: string | ArrayBuffer, +): GitHubActionsRelayInput | null { + const legacyFrame = decodeGitHubActionsRelayFrame(message, relayInputFrameType); + const generatedFrame = decodeGitHubActionsRelayFrame(message, relayGenerationInputFrameType); + const decoded = generatedFrame + ? decodeGitHubActionsRelayGeneration(generatedFrame.payload) + : null; + const frame = legacyFrame ?? (decoded ? { ...generatedFrame!, payload: decoded.payload } : null); + if (!frame?.inputId) return null; + return { + inputId: frame.inputId, + payload: Uint8Array.from(frame.payload).buffer, + ...(decoded ? { generation: decoded.generation } : {}), + }; +} + +export function parseGitHubActionsRunnerProtocol( + value: string | null, +): GitHubActionsRunnerProtocol | null { + return isGitHubActionsRelayProtocol(value) ? value : null; +} + +export function parseGitHubActionsRunnerProtocolOffer( + value: string | null, +): typeof githubActionsGenerationFencedCapability | null { + if (!value) return null; + return value + .split(",") + .map((protocol) => protocol.trim()) + .includes(githubActionsGenerationFencedCapability) + ? githubActionsGenerationFencedCapability + : null; +} + +export function parseGitHubActionsViewerProtocol( + value: string | null, +): GitHubActionsViewerProtocol | null { + return isGitHubActionsRelayProtocol(value) ? value : null; +} + +export function attachGitHubActionsRunnerProtocol( + socket: GitHubActionsRelaySocket, + protocol: GitHubActionsRunnerProtocol | null, + generation?: string, +): void { + socket.serializeAttachment?.( + protocol || generation + ? ({ + ...(protocol ? { protocol } : {}), + ...(generation ? { generation } : {}), + } satisfies GitHubActionsRelayAttachment) + : {}, + ); +} + +export function attachGitHubActionsViewerProtocol( + socket: GitHubActionsRelaySocket, + protocol: GitHubActionsViewerProtocol | null, +): void { + socket.serializeAttachment?.( + protocol ? ({ protocol } satisfies GitHubActionsRelayAttachment) : {}, + ); +} + +export function encodeGitHubActionsRelayInputAcknowledgement( + acknowledgement: GitHubActionsRelayInputAcknowledgement, +): ArrayBuffer { + requireGitHubActionsRelayInputId(acknowledgement.inputId); + const error = + acknowledgement.accepted || !acknowledgement.error + ? new Uint8Array() + : encoder.encode(acknowledgement.error.trim()); + const payload = new Uint8Array(1 + error.byteLength); + payload[0] = acknowledgement.accepted ? 1 : 0; + payload.set(error, 1); + return encodeGitHubActionsRelayFrame( + acknowledgement.generation + ? relayGenerationInputAcknowledgementFrameType + : relayInputAcknowledgementFrameType, + acknowledgement.inputId, + acknowledgement.generation + ? encodeGitHubActionsRelayGeneration(acknowledgement.generation, payload) + : payload, + ); +} + +export function parseGitHubActionsRelayEvent( + message: string | ArrayBuffer, +): { generation?: string; type: keyof typeof relayEventCodes } | null { + const legacyFrame = decodeGitHubActionsRelayFrame(message, relayEventFrameType); + const generatedFrame = decodeGitHubActionsRelayFrame(message, relayGenerationEventFrameType); + const decoded = generatedFrame + ? decodeGitHubActionsRelayGeneration(generatedFrame.payload) + : null; + const frame = legacyFrame ?? (decoded ? { ...generatedFrame!, payload: decoded.payload } : null); + if (!frame || frame.inputId || frame.payload.byteLength !== 1) return null; + const type = relayEvents.get(frame.payload[0] ?? 0); + return type ? { type, ...(decoded ? { generation: decoded.generation } : {}) } : null; +} + export function isGitHubActionsViewerControlMessage(message: string | ArrayBuffer): boolean { if (typeof message !== "string") return false; try { @@ -136,13 +528,198 @@ export function isGitHubActionsViewerControlMessage(message: string | ArrayBuffe export function notifyGitHubActionsViewers( viewers: readonly GitHubActionsRelaySocket[], type: "runner_connected" | "runner_disconnected" | "runner_waiting", + generation?: string, ): number { - const payload = JSON.stringify({ type }); let notified = 0; for (const socket of viewers) { if (socket.readyState !== webSocketOpen) continue; - socket.send(payload); + const usesGenerations = gitHubActionsRelayUsesGenerations(socket); + const framedPayload = encodeGitHubActionsRelayFrame( + usesGenerations ? relayGenerationEventFrameType : relayEventFrameType, + "", + usesGenerations + ? encodeGitHubActionsRelayGeneration( + generation ?? "none", + new Uint8Array([relayEventCodes[type]]), + ) + : new Uint8Array([relayEventCodes[type]]), + ); + socket.send( + gitHubActionsViewerUsesFramedProtocol(socket) ? framedPayload : JSON.stringify({ type }), + ); notified += 1; } return notified; } + +function encodeGitHubActionsRelayFrame( + type: number, + inputId: string, + payload: Uint8Array, +): ArrayBuffer { + const inputIdBytes = encoder.encode(inputId); + if ( + inputIdBytes.byteLength > relayInputIdMaximumBytes || + (inputId && !relayInputIdPattern.test(inputId)) + ) { + throw new Error("invalid GitHub Actions relay input id"); + } + const frame = new Uint8Array( + relayFrameHeaderBytes + inputIdBytes.byteLength + payload.byteLength, + ); + frame.set(relayFrameMagic, 0); + frame[relayFrameMagic.byteLength] = type; + frame[relayFrameMagic.byteLength + 1] = inputIdBytes.byteLength; + frame.set(inputIdBytes, relayFrameHeaderBytes); + frame.set(payload, relayFrameHeaderBytes + inputIdBytes.byteLength); + return frame.buffer; +} + +function decodeGitHubActionsRelayFrame( + message: string | ArrayBuffer, + expectedType: number, +): { inputId: string; payload: Uint8Array } | null { + if (typeof message === "string") return null; + const frame = new Uint8Array(message); + if (frame.byteLength < relayFrameHeaderBytes) return null; + for (const [index, value] of relayFrameMagic.entries()) { + if (frame[index] !== value) return null; + } + if (frame[relayFrameMagic.byteLength] !== expectedType) return null; + const inputIdBytes = frame[relayFrameMagic.byteLength + 1] ?? 0; + if ( + inputIdBytes > relayInputIdMaximumBytes || + relayFrameHeaderBytes + inputIdBytes > frame.byteLength + ) { + return null; + } + const inputId = decoder.decode( + frame.subarray(relayFrameHeaderBytes, relayFrameHeaderBytes + inputIdBytes), + ); + if (inputId && !relayInputIdPattern.test(inputId)) return null; + return { + inputId, + payload: frame.subarray(relayFrameHeaderBytes + inputIdBytes), + }; +} + +function messageBytes(message: string | ArrayBuffer | ArrayBufferView): Uint8Array { + if (typeof message === "string") return encoder.encode(message); + if (ArrayBuffer.isView(message)) { + return new Uint8Array(message.buffer, message.byteOffset, message.byteLength); + } + return new Uint8Array(message); +} + +function requireGitHubActionsRelayInputId(inputId: string): void { + if (!inputId) throw new Error("invalid GitHub Actions relay input id"); +} + +export function gitHubActionsRunnerUsesFramedProtocol(socket: GitHubActionsRelaySocket): boolean { + return gitHubActionsRelayUsesFramedProtocol(socket); +} + +export function gitHubActionsViewerUsesFramedProtocol(socket: GitHubActionsRelaySocket): boolean { + return gitHubActionsRelayUsesFramedProtocol(socket); +} + +export function gitHubActionsRelayGeneration(socket: GitHubActionsRelaySocket): string | undefined { + const attachment = socket.deserializeAttachment?.(); + if (!attachment || typeof attachment !== "object") return undefined; + const generation = (attachment as GitHubActionsRelayAttachment).generation; + return isGitHubActionsRelayGeneration(generation) ? generation : undefined; +} + +export function gitHubActionsRelayUsesGenerations(socket: GitHubActionsRelaySocket): boolean { + const attachment = socket.deserializeAttachment?.(); + if (!attachment || typeof attachment !== "object") return false; + return ( + (attachment as GitHubActionsRelayAttachment).protocol === + githubActionsGenerationFencedCapability + ); +} + +export function createGitHubActionsRelayGeneration(): string { + return crypto.randomUUID().replaceAll("-", ""); +} + +function gitHubActionsRelayUsesFramedProtocol(socket: GitHubActionsRelaySocket): boolean { + const attachment = socket.deserializeAttachment?.(); + if (!attachment || typeof attachment !== "object") return false; + return isGitHubActionsRelayProtocol((attachment as GitHubActionsRelayAttachment).protocol); +} + +function sendGitHubActionsViewerInputAcknowledgement( + viewer: GitHubActionsRelaySocket, + inputId: string | null, + accepted: boolean, + generation?: string, +): boolean { + if (inputId) { + return sendGitHubActionsRelayInputAcknowledgement(viewer, { + inputId, + accepted, + ...(gitHubActionsRelayUsesGenerations(viewer) ? { generation: generation ?? "none" } : {}), + }); + } + if (viewer.readyState !== webSocketOpen) return false; + try { + viewer.send( + JSON.stringify({ + type: "github_actions_input_ack", + accepted, + ...(accepted ? {} : { error: relayInputRejectedError }), + }), + ); + return true; + } catch { + return false; + } +} + +function isGitHubActionsRelayProtocol(value: unknown): value is GitHubActionsRelayProtocol { + return ( + value === githubActionsFramedRunnerCapability || + value === githubActionsGenerationFencedCapability + ); +} + +function isGitHubActionsRelayGeneration(value: unknown): value is string { + return ( + typeof value === "string" && + value.length > 0 && + encoder.encode(value).byteLength <= relayGenerationMaximumBytes && + relayGenerationPattern.test(value) + ); +} + +function encodeGitHubActionsRelayGeneration(generation: string, payload: Uint8Array): Uint8Array { + if (!isGitHubActionsRelayGeneration(generation)) { + throw new Error("invalid GitHub Actions relay generation"); + } + const generationBytes = encoder.encode(generation); + const generatedPayload = new Uint8Array(1 + generationBytes.byteLength + payload.byteLength); + generatedPayload[0] = generationBytes.byteLength; + generatedPayload.set(generationBytes, 1); + generatedPayload.set(payload, 1 + generationBytes.byteLength); + return generatedPayload; +} + +function decodeGitHubActionsRelayGeneration( + payload: Uint8Array, +): { generation: string; payload: Uint8Array } | null { + const generationBytes = payload[0] ?? 0; + if ( + generationBytes === 0 || + generationBytes > relayGenerationMaximumBytes || + 1 + generationBytes > payload.byteLength + ) { + return null; + } + const generation = decoder.decode(payload.subarray(1, 1 + generationBytes)); + if (!isGitHubActionsRelayGeneration(generation)) return null; + return { + generation, + payload: payload.subarray(1 + generationBytes), + }; +} diff --git a/src/terminal-multiplayer.ts b/src/terminal-multiplayer.ts index 3fb8a860..05346150 100644 --- a/src/terminal-multiplayer.ts +++ b/src/terminal-multiplayer.ts @@ -69,8 +69,7 @@ export function attributedTerminalInputPayloads( ): Uint8Array[] { const sender = terminalSenderTag(user); const attributed = `${sender} ${terminalSingleLineInput(submitted.text)}${submitted.eol}`; - const chunks = submitted.replaceCurrentLine ? ["\x15", ...attributed] : [...attributed]; - return chunks.map((chunk) => encoder.encode(chunk)); + return [encoder.encode(`${submitted.replaceCurrentLine ? "\x15" : ""}${attributed}`)]; } function updateTerminalInputLine(state: TerminalInputState, text: string): void { diff --git a/src/worker/card-repository.ts b/src/worker/card-repository.ts index bf06472c..5b1fde5f 100644 --- a/src/worker/card-repository.ts +++ b/src/worker/card-repository.ts @@ -337,7 +337,7 @@ export class CardRepository implements CardLifecycleStore { async claimRun(input: CardRunClaimInput): Promise<"claimed" | "capacity" | "active"> { const db = database(this.env); - const transition = await sql` + const transition = sql` UPDATE cards SET lane = 'Running', active_run_id = ${input.runId}, @@ -351,8 +351,66 @@ export class CardRepository implements CardLifecycleStore { AND (lane = 'Running' OR ( SELECT count(*) FROM cards WHERE lane = 'Running' AND id <> ${input.card.id} ) < ${input.cap}) - `.execute(db); - if ((transition.numAffectedRows ?? 0n) === 0n) { + AND NOT EXISTS ( + SELECT 1 + FROM run_attempts + WHERE id = ${input.runId} + OR (card_id = ${input.card.id} AND attempt = ${input.attempt}) + ) + `; + const insert = sql` + INSERT INTO run_attempts ( + id, card_id, attempt, runtime, status, control_intent, lease_id, attach_url, vnc_url, + selection_reason, capabilities_json, operator, last_heartbeat_at, started_at, ended_at, + created_at, updated_at, error + ) + SELECT + ${input.runId}, ${input.card.id}, ${input.attempt}, ${input.descriptor.runtime}, 'queued', + NULL, NULL, NULL, NULL, ${input.descriptor.reason}, + ${JSON.stringify(input.descriptor.capabilities)}, NULL, ${input.now}, ${input.now}, NULL, + ${input.now}, ${input.now}, NULL + FROM cards + WHERE id = ${input.card.id} + AND active_run_id = ${input.runId} + AND updated_at = ${input.now} + AND NOT EXISTS ( + SELECT 1 + FROM run_attempts + WHERE id = ${input.runId} + OR (card_id = ${input.card.id} AND attempt = ${input.attempt}) + ) + `; + const results = await this.env.DB.batch( + [transition, insert].map((query) => { + const compiled = query.compile(db); + return this.env.DB.prepare(compiled.sql).bind(...compiled.parameters); + }), + ); + if ((results[0]?.meta.changes ?? 0) === 0) { + const duplicateAttempt = await db + .selectFrom("run_attempts") + .select("id") + .where((expressions) => + expressions.or([ + expressions("id", "=", input.runId), + expressions.and([ + expressions("card_id", "=", input.card.id), + expressions("attempt", "=", input.attempt), + ]), + ]), + ) + .executeTakeFirst(); + if (duplicateAttempt) return "active"; + + const activeCardRun = await db + .selectFrom("cards") + .innerJoin("run_attempts", "run_attempts.id", "cards.active_run_id") + .select("cards.id") + .where("cards.id", "=", input.card.id) + .where("run_attempts.status", "in", activeRunStatuses) + .executeTakeFirst(); + if (activeCardRun) return "active"; + const activeCount = await db .selectFrom("cards") .select(sql`count(*)`.as("count")) @@ -360,30 +418,9 @@ export class CardRepository implements CardLifecycleStore { .executeTakeFirst(); return Number(activeCount?.count ?? 0) >= input.cap ? "capacity" : "active"; } - await db - .insertInto("run_attempts") - .values({ - id: input.runId, - card_id: input.card.id, - attempt: input.attempt, - runtime: input.descriptor.runtime, - status: "queued", - control_intent: null, - lease_id: null, - attach_url: null, - vnc_url: null, - selection_reason: input.descriptor.reason, - capabilities_json: JSON.stringify(input.descriptor.capabilities), - operator: null, - last_heartbeat_at: input.now, - started_at: input.now, - ended_at: null, - created_at: input.now, - updated_at: input.now, - error: null, - }) - .onConflict((conflict) => conflict.doNothing()) - .execute(); + if ((results[1]?.meta.changes ?? 0) !== 1) { + throw new Error("card run claim did not persist its run attempt"); + } return "claimed"; } diff --git a/src/worker/database.ts b/src/worker/database.ts index 2bdd96e9..971807e5 100644 --- a/src/worker/database.ts +++ b/src/worker/database.ts @@ -69,6 +69,9 @@ export type DesktopHostTable = { name: string; address: string; port: number; + ownership_token: string; + publication_id: string; + publication_write_token: Generated; created_at: number; updated_at: number; }; @@ -220,6 +223,26 @@ export type InteractiveSessionTable = { export type InteractiveSessionRow = Selectable; +export type RuntimeAdapterWorkspaceCleanupTable = { + session_id: string; + adapter_workspace_id: string; + profile: string | null; + control_plane: string | null; + create_pending: number; + deletion_observed: Generated; + message: string; + reconcile_error: string | null; + attempt_count: Generated; + last_attempt_at: number | null; + next_attempt_at: number; + cleanup_claim: string | null; + cleanup_claim_expires_at: number | null; + created_at: number; + updated_at: number; +}; + +export type RuntimeAdapterWorkspaceCleanupRow = Selectable; + export type InteractiveSessionGrantTable = { session_id: string; subject: string; @@ -300,6 +323,25 @@ export type InteractiveSessionCredentialPolicyTable = { updated_at: number; }; +export type InteractiveSessionCredentialPolicyRegistrationTable = { + session_id: string; + sandbox_id: string; + state: "registering" | "cleanup_pending"; + registration_generation: string; + registration_claim: string | null; + registration_claim_expires_at: number | null; + attempt_count: Generated; + last_attempt_at: number | null; + last_error: string | null; + cleanup_claim: string | null; + cleanup_claim_expires_at: number | null; + rollback_policies_json: Generated; + lookup_ids_json: Generated; + repair_generation: Generated; + created_at: number; + updated_at: number; +}; + export type CredentialPolicyReconcileStateTable = { id: number; last_rowid: number; @@ -375,11 +417,13 @@ export type Database = { cards: CardTable; run_attempts: RunAttemptTable; interactive_sessions: InteractiveSessionTable; + runtime_adapter_workspace_cleanups: RuntimeAdapterWorkspaceCleanupTable; interactive_session_grants: InteractiveSessionGrantTable; openclaw_request_replays: OpenClawRequestReplayTable; interactive_session_events: InteractiveSessionEventTable; interactive_session_log_archives: InteractiveSessionLogArchiveTable; interactive_session_credential_policies: InteractiveSessionCredentialPolicyTable; + interactive_session_credential_policy_registrations: InteractiveSessionCredentialPolicyRegistrationTable; credential_policy_reconcile_state: CredentialPolicyReconcileStateTable; standalone_sandbox_provisions: StandaloneSandboxProvisionTable; repo_workflows: RepoWorkflowTable; diff --git a/src/worker/deployment.ts b/src/worker/deployment.ts index 4ca7710b..f94ab9f2 100644 --- a/src/worker/deployment.ts +++ b/src/worker/deployment.ts @@ -9,6 +9,7 @@ import { runtimeProfileByID, type RuntimeProfileDescriptor, } from "../runtime-profiles.ts"; +import { runtimeAdapterControlPlaneForProfile } from "../runtime-adapter.ts"; import { trustedProxyPublicOrigin, type TrustedProxyEnv } from "../trusted-proxy-auth.ts"; import { configuredHttpOrigin } from "../url-security.ts"; import { badRequest } from "./http.ts"; @@ -43,6 +44,8 @@ export type DeploymentEnv = TrustedProxyEnv & { CRABFLEET_INTERACTIVE_RUNTIMES?: string; CRABFLEET_DEFAULT_PROFILE?: string; CRABFLEET_RUNTIME_PROFILES_JSON?: string; + CRABBOX_RUNTIME_ADAPTER_URL?: string; + CRABBOX_RUNTIME_ADAPTER_URL_TEMPLATE?: string; }; export function deploymentConfig(env: DeploymentEnv): DeploymentConfig { @@ -52,6 +55,7 @@ export function deploymentConfig(env: DeploymentEnv): DeploymentConfig { if (runtimeProfiles.length > 0 && !runtimeProfileByID(runtimeProfiles, defaultProfile)) { throw new TypeError("CRABFLEET_DEFAULT_PROFILE must name a configured runtime profile"); } + validateRuntimeProfileRoutes(env, runtimeProfiles, defaultProfile); return { label: clean(env.CRABFLEET_LABEL, 80) || "Crabfleet", canonicalUrl: configuredHttpOrigin(env.CRABFLEET_CANONICAL_URL, appCanonicalOrigin), @@ -65,6 +69,28 @@ export function deploymentConfig(env: DeploymentEnv): DeploymentConfig { }; } +function validateRuntimeProfileRoutes( + env: DeploymentEnv, + runtimeProfiles: RuntimeProfileDescriptor[], + defaultProfile: string, +): void { + const direct = env.CRABBOX_RUNTIME_ADAPTER_URL; + const template = env.CRABBOX_RUNTIME_ADAPTER_URL_TEMPLATE; + const hasDirect = typeof direct === "string" && direct.length > 0; + const hasTemplate = typeof template === "string" && template.length > 0; + if (hasDirect || !hasTemplate) return; + const profileIDs = + runtimeProfiles.length > 0 ? runtimeProfiles.map((profile) => profile.id) : [defaultProfile]; + const unroutable = profileIDs.find( + (profile) => !runtimeAdapterControlPlaneForProfile(undefined, template, profile), + ); + if (unroutable) { + throw new TypeError( + `runtime profile ${unroutable} cannot be routed by CRABBOX_RUNTIME_ADAPTER_URL_TEMPLATE`, + ); + } +} + export function selectedRuntimeProfile( deployment: DeploymentConfig, value: unknown, diff --git a/src/worker/desktop-host-repository.ts b/src/worker/desktop-host-repository.ts index 63147aae..8746e4d9 100644 --- a/src/worker/desktop-host-repository.ts +++ b/src/worker/desktop-host-repository.ts @@ -1,5 +1,6 @@ -import { database } from "./database.ts"; +import { database, executeBatch } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; +import { conflict } from "./http.ts"; export type DesktopHostRow = { ownerSubject: string; @@ -8,6 +9,8 @@ export type DesktopHostRow = { name: string; address: string; port: number; + ownershipToken: string; + publicationID: string; createdAt: number; updatedAt: number; }; @@ -17,7 +20,12 @@ export type DesktopHostWrite = DesktopHostRow; export interface DesktopHostStore { list(ownerSubject: string): Promise; upsert(host: DesktopHostWrite): Promise; - remove(ownerSubject: string, id: string): Promise; + ownershipTokenForPublication( + ownerSubject: string, + id: string, + publicationID: string, + ): Promise; + remove(ownerSubject: string, id: string, ownershipToken: string | null): Promise; } export class DesktopHostRepository implements DesktopHostStore { @@ -42,13 +50,15 @@ export class DesktopHostRepository implements DesktopHostStore { name: row.name, address: row.address, port: row.port, + ownershipToken: row.ownership_token, + publicationID: row.publication_id, createdAt: row.created_at, updatedAt: row.updated_at, })); } async upsert(host: DesktopHostWrite): Promise { - await database(this.env) + const row = await database(this.env) .insertInto("desktop_hosts") .values({ owner_subject: host.ownerSubject, @@ -57,25 +67,38 @@ export class DesktopHostRepository implements DesktopHostStore { name: host.name, address: host.address, port: host.port, + ownership_token: host.ownershipToken, + publication_id: host.publicationID, + publication_write_token: host.ownershipToken, created_at: host.createdAt, updated_at: host.updatedAt, }) - .onConflict((conflict) => - conflict.columns(["owner_subject", "id"]).doUpdateSet({ - owner: host.owner, - name: host.name, - address: host.address, - port: host.port, - updated_at: host.updatedAt, - }), - ) - .execute(); - const row = await database(this.env) - .selectFrom("desktop_hosts") - .selectAll() - .where("owner_subject", "=", host.ownerSubject) - .where("id", "=", host.id) - .executeTakeFirstOrThrow(); + .onConflict((conflict) => { + const update = conflict.columns(["owner_subject", "id"]); + return host.ownershipToken + ? update.doUpdateSet({ + owner: host.owner, + name: host.name, + address: host.address, + port: host.port, + ownership_token: host.ownershipToken, + publication_id: host.publicationID, + publication_write_token: host.ownershipToken, + updated_at: host.updatedAt, + }) + : update + .doUpdateSet({ + owner: host.owner, + name: host.name, + address: host.address, + port: host.port, + updated_at: host.updatedAt, + }) + .where("desktop_hosts.ownership_token", "=", ""); + }) + .returningAll() + .executeTakeFirst(); + if (!row) throw desktopHostOwnershipConflict(); return { ownerSubject: row.owner_subject, id: row.id, @@ -83,16 +106,67 @@ export class DesktopHostRepository implements DesktopHostStore { name: row.name, address: row.address, port: row.port, + ownershipToken: row.ownership_token, + publicationID: row.publication_id, createdAt: row.created_at, updatedAt: row.updated_at, }; } - async remove(ownerSubject: string, id: string): Promise { - await database(this.env) - .deleteFrom("desktop_hosts") + async ownershipTokenForPublication( + ownerSubject: string, + id: string, + publicationID: string, + ): Promise { + const row = await database(this.env) + .selectFrom("desktop_hosts") + .select("ownership_token") .where("owner_subject", "=", ownerSubject) .where("id", "=", id) - .execute(); + .where("publication_id", "=", publicationID) + .where("ownership_token", "<>", "") + .executeTakeFirst(); + return row?.ownership_token ?? null; + } + + async remove(ownerSubject: string, id: string, ownershipToken: string | null): Promise { + const db = database(this.env); + if (!ownershipToken) { + const deleted = await db + .deleteFrom("desktop_hosts") + .where("owner_subject", "=", ownerSubject) + .where("id", "=", id) + .where("ownership_token", "=", "") + .executeTakeFirst(); + if ((deleted.numDeletedRows ?? 0n) > 0n) return; + const existing = await db + .selectFrom("desktop_hosts") + .select("ownership_token") + .where("owner_subject", "=", ownerSubject) + .where("id", "=", id) + .executeTakeFirst(); + if (existing?.ownership_token) throw desktopHostOwnershipConflict(); + return; + } + const deleteMarker = `delete-authorized:${crypto.randomUUID()}`; + // The migration trigger permits this marker only; the atomic batch keeps it + // invisible to legacy workers between authorization and deletion. + await executeBatch(this.env, [ + db + .updateTable("desktop_hosts") + .set({ ownership_token: deleteMarker }) + .where("owner_subject", "=", ownerSubject) + .where("id", "=", id) + .where("ownership_token", "=", ownershipToken), + db + .deleteFrom("desktop_hosts") + .where("owner_subject", "=", ownerSubject) + .where("id", "=", id) + .where("ownership_token", "=", deleteMarker), + ]); } } + +function desktopHostOwnershipConflict(): ReturnType { + return conflict("desktop host is owned by a token-aware registration"); +} diff --git a/src/worker/desktop-host-service.ts b/src/worker/desktop-host-service.ts index f3931419..fbc7d322 100644 --- a/src/worker/desktop-host-service.ts +++ b/src/worker/desktop-host-service.ts @@ -19,13 +19,30 @@ export type DesktopHost = { updatedAt: number; }; +export type DesktopHostRegistration = { + host: DesktopHost; + ownershipToken?: string; +}; + +export const desktopHostOwnershipHeader = "x-crabfleet-ownership-token"; +export const desktopHostOwnershipModeHeader = "x-crabfleet-ownership-mode"; +export const desktopHostPublicationHeader = "x-crabfleet-publication-id"; +export const desktopHostTokenOwnershipMode = "token-v1"; +export type DesktopHostOwnershipMode = "legacy" | typeof desktopHostTokenOwnershipMode; + export class DesktopHostService { private readonly store: DesktopHostStore; private readonly now: () => number; + private readonly createOwnershipToken: () => string; - constructor(store: DesktopHostStore, now: () => number = Date.now) { + constructor( + store: DesktopHostStore, + now: () => number = Date.now, + createOwnershipToken: () => string = randomOwnershipToken, + ) { this.store = store; this.now = now; + this.createOwnershipToken = createOwnershipToken; } async list(user: User): Promise { @@ -33,12 +50,24 @@ export class DesktopHostService { return rows.map(presentDesktopHost); } - async register(user: User, rawID: string, input: DesktopHostInput): Promise { + async register( + user: User, + rawID: string, + input: DesktopHostInput, + ownershipMode: DesktopHostOwnershipMode = "legacy", + rawPublicationID: unknown = null, + ): Promise { const id = desktopHostID(rawID); const name = boundedText(input.name, "name", 100); const address = tailscaleIPv4(input.address); const port = desktopHostPort(input.port); const now = this.now(); + const publicationID = + ownershipMode === desktopHostTokenOwnershipMode + ? desktopHostPublicationID(rawPublicationID) + : ""; + const ownershipToken = + ownershipMode === desktopHostTokenOwnershipMode ? this.createOwnershipToken() : ""; const host: DesktopHostRow = { ownerSubject: tenantSubject(user), id, @@ -46,17 +75,44 @@ export class DesktopHostService { name, address, port, + ownershipToken, + publicationID, createdAt: now, updatedAt: now, }; - return presentDesktopHost(await this.store.upsert(host)); + const registration: DesktopHostRegistration = { + host: presentDesktopHost(await this.store.upsert(host)), + }; + if (ownershipToken) registration.ownershipToken = ownershipToken; + return registration; + } + + async recover( + user: User, + rawID: string, + rawPublicationID: unknown, + ): Promise<{ ownershipToken: string | null }> { + const ownershipToken = await this.store.ownershipTokenForPublication( + tenantSubject(user), + desktopHostID(rawID), + desktopHostPublicationID(rawPublicationID), + ); + return { ownershipToken }; } - async remove(user: User, rawID: string): Promise { - await this.store.remove(tenantSubject(user), desktopHostID(rawID)); + async remove(user: User, rawID: string, rawOwnershipToken: unknown): Promise { + await this.store.remove( + tenantSubject(user), + desktopHostID(rawID), + desktopHostOwnershipToken(rawOwnershipToken), + ); } } +function randomOwnershipToken(): string { + return crypto.randomUUID() + crypto.randomUUID(); +} + function presentDesktopHost(row: DesktopHostRow): DesktopHost { return { id: row.id, @@ -121,3 +177,34 @@ function desktopHostPort(value: unknown): number { } return value; } + +function desktopHostOwnershipToken(value: unknown): string | null { + if (value === null || value === undefined) return null; + if ( + typeof value !== "string" || + value.length === 0 || + new TextEncoder().encode(value).byteLength > 200 || + [...value].some((character) => { + const codePoint = character.codePointAt(0) ?? 0; + return codePoint <= 0x20 || codePoint === 0x7f; + }) + ) { + throw badRequest("desktop host ownership token is required"); + } + return value; +} + +function desktopHostPublicationID(value: unknown): string { + if ( + typeof value !== "string" || + value.length === 0 || + new TextEncoder().encode(value).byteLength > 200 || + [...value].some((character) => { + const codePoint = character.codePointAt(0) ?? 0; + return codePoint <= 0x20 || codePoint === 0x7f; + }) + ) { + throw badRequest("desktop host publication id is required"); + } + return value; +} diff --git a/src/worker/github-actions-application.ts b/src/worker/github-actions-application.ts index 066682d4..ac8e27ce 100644 --- a/src/worker/github-actions-application.ts +++ b/src/worker/github-actions-application.ts @@ -1,4 +1,10 @@ import type { GitHubActionsSessionRegistrationInput } from "./github-actions-session-registration.ts"; +import { + githubActionsRunnerProtocolHeader, + githubActionsRunnerProtocolQuery, + parseGitHubActionsRunnerProtocol, + parseGitHubActionsRunnerProtocolOffer, +} from "../github-actions-runtime.ts"; import { AdminRepository } from "./admin-repository.ts"; import { GitHubActionsSessionRegistrationService, @@ -84,7 +90,11 @@ export class GitHubActionsApplication { nextSessionId: () => nextInteractiveSessionId(this.env), insertSession: (values) => repository.insertSession(values), readById: (id) => repository.readById(id), - updateSession: (id, values) => repository.updateSession(id, values), + updateSession: (id, values, expected) => + repository.updateSession(id, values, { + kind: "registration", + registration: expected, + }), isConstraintError, disconnectRunner: (id) => this.disconnectRunner(id), appendEvent: (id, message, now) => this.appendMessageEvent(id, user, message, now), @@ -104,7 +114,12 @@ export class GitHubActionsApplication { const store: GitHubActionsWorkStateStore = { now: () => Date.now(), readRow: (sessionId) => repository.readById(sessionId), - persist: (sessionId, values) => repository.updateSession(sessionId, values), + persist: (sessionId, values, expectedRevision, expectedTerminalStatus) => + repository.updateSession(sessionId, values, { + kind: "authenticated", + revision: expectedRevision, + ...(expectedTerminalStatus ? { terminalStatus: expectedTerminalStatus } : {}), + }), appendEvent: (sessionId, message, now) => this.appendMessageEvent(sessionId, user, message, now), disconnectRunner: (sessionId) => this.disconnectRunner(sessionId), @@ -151,13 +166,17 @@ export class GitHubActionsApplication { const repository = this.repository(); const store: GitHubActionsRunnerConnectionStore = { now: () => Date.now(), - persist: (sessionId, values) => repository.updateSession(sessionId, values), + persist: (sessionId, values, expectedRevision) => + repository.updateSession(sessionId, values, { + kind: "authenticated", + revision: expectedRevision, + }), appendEvent: (sessionId, message, now) => this.appendMessageEvent(sessionId, user, message, now), }; await new GitHubActionsRunnerConnectionService(store).connect(session); - return stub.fetch("https://crabfleet.internal/api/session-control/github-actions/runner", { - headers: { upgrade: "websocket" }, + return stub.fetch(gitHubActionsRelayRunnerUrl(request), { + headers: gitHubActionsRelayRunnerHeaders(request), }); } @@ -200,6 +219,24 @@ export class GitHubActionsApplication { } } +export function gitHubActionsRelayRunnerUrl(request: Request): string { + const relayUrl = new URL("https://crabfleet.internal/api/session-control/github-actions/runner"); + const protocol = parseGitHubActionsRunnerProtocol( + new URL(request.url).searchParams.get(githubActionsRunnerProtocolQuery), + ); + if (protocol) relayUrl.searchParams.set(githubActionsRunnerProtocolQuery, protocol); + return relayUrl.toString(); +} + +export function gitHubActionsRelayRunnerHeaders(request: Request): Headers { + const headers = new Headers({ upgrade: "websocket" }); + const offeredProtocol = parseGitHubActionsRunnerProtocolOffer( + request.headers.get(githubActionsRunnerProtocolHeader), + ); + if (offeredProtocol) headers.set(githubActionsRunnerProtocolHeader, offeredProtocol); + return headers; +} + function isConstraintError(error: unknown): boolean { return error instanceof Error && /constraint|unique/i.test(error.message); } diff --git a/src/worker/github-actions-repository.ts b/src/worker/github-actions-repository.ts index 931c69d7..c349726b 100644 --- a/src/worker/github-actions-repository.ts +++ b/src/worker/github-actions-repository.ts @@ -4,14 +4,33 @@ import type { InteractiveSessionRow, InteractiveSessionTable } from "./database. import { database } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import type { GitHubActionsRunnerConnectionUpdate } from "./github-actions-runner-connection.ts"; -import type { GitHubActionsSessionRegistrationUpdate } from "./github-actions-session-registration.ts"; +import type { + GitHubActionsSessionRegistrationExpectation, + GitHubActionsSessionRegistrationUpdate, +} from "./github-actions-session-registration.ts"; import type { GitHubActionsWorkStateUpdate } from "./github-actions-session-work-state.ts"; +import { conflict } from "./http.ts"; +import type { InteractiveSessionStatus } from "./models.ts"; type GitHubActionsSessionUpdate = | GitHubActionsSessionRegistrationUpdate | GitHubActionsWorkStateUpdate | GitHubActionsRunnerConnectionUpdate; +export type GitHubActionsSessionUpdateExpectation = + | { + kind: "registration"; + registration: GitHubActionsSessionRegistrationExpectation; + } + | { + kind: "authenticated"; + revision: number; + terminalStatus?: InteractiveSessionStatus; + }; + +const terminalWorkStates = ["completed", "failed", "canceled", "blocked"]; +const terminalSessionStatuses = ["stopped", "expired", "failed"] as const; + export class GitHubActionsRepository { private readonly env: RuntimeEnv; @@ -43,11 +62,71 @@ export class GitHubActionsRepository { await database(this.env).insertInto("interactive_sessions").values(values).execute(); } - async updateSession(id: string, values: GitHubActionsSessionUpdate): Promise { - await database(this.env) + async updateSession( + id: string, + values: GitHubActionsSessionUpdate, + expectation: GitHubActionsSessionUpdateExpectation, + ): Promise { + let update = database(this.env) .updateTable("interactive_sessions") .set(values) .where("id", "=", id) - .execute(); + .where("runtime", "=", "github_actions"); + + if (isRegistrationUpdate(values)) { + if (expectation.kind !== "registration") { + throw new Error("GitHub Actions registration update requires expected state"); + } + const expectedRegistration = expectation.registration; + update = update + .where("updated_at", "=", expectedRegistration.updated_at) + .where("status", "=", expectedRegistration.status) + .where("work_state", "=", expectedRegistration.work_state) + .where("work_phase", "=", expectedRegistration.work_phase) + .where("owner_subject", "=", values.owner_subject); + update = + expectedRegistration.agent_token_hash === null + ? update.where("agent_token_hash", "is", null) + : update.where("agent_token_hash", "=", expectedRegistration.agent_token_hash); + } else if (isWorkStateUpdate(values) && terminalWorkStates.includes(values.work_state)) { + if (expectation.kind !== "authenticated" || !expectation.terminalStatus) { + throw new Error("terminal GitHub Actions update requires expected session status"); + } + update = update + .where("updated_at", "=", expectation.revision) + .where("status", "=", expectation.terminalStatus) + .where("status", "not in", terminalSessionStatuses) + .where((expressions) => + expressions.or([ + expressions("work_state", "not in", terminalWorkStates), + expressions("work_state", "=", values.work_state), + ]), + ); + } else { + if (expectation.kind !== "authenticated") { + throw new Error("GitHub Actions update requires authenticated revision"); + } + update = update + .where("updated_at", "=", expectation.revision) + .where("work_state", "not in", terminalWorkStates) + .where("status", "not in", terminalSessionStatuses); + } + + const result = await update.executeTakeFirst(); + if ((result.numUpdatedRows ?? 0n) !== 1n) { + throw conflict("GitHub Actions session changed; retry"); + } } } + +function isRegistrationUpdate( + values: GitHubActionsSessionUpdate, +): values is GitHubActionsSessionRegistrationUpdate { + return "agent_token_hash" in values; +} + +function isWorkStateUpdate( + values: GitHubActionsSessionUpdate, +): values is GitHubActionsWorkStateUpdate { + return "stopped_at" in values && "codex_thread_id" in values; +} diff --git a/src/worker/github-actions-runner-connection.ts b/src/worker/github-actions-runner-connection.ts index ca9cdf53..5eb37e40 100644 --- a/src/worker/github-actions-runner-connection.ts +++ b/src/worker/github-actions-runner-connection.ts @@ -17,7 +17,11 @@ export type GitHubActionsRunnerConnectionUpdate = { export type GitHubActionsRunnerConnectionStore = { now(): number; - persist(id: string, values: GitHubActionsRunnerConnectionUpdate): Promise; + persist( + id: string, + values: GitHubActionsRunnerConnectionUpdate, + expectedRevision: number, + ): Promise; appendEvent(id: string, message: string, now: number): Promise; }; @@ -33,6 +37,7 @@ export class GitHubActionsRunnerConnectionService { throw badRequest("session is not a GitHub Actions work session"); } const now = this.store.now(); + const revision = Math.max(session.updatedAt + 1, now); const state = session.workState === "registered" || !session.workState ? "running" : session.workState; const phase = @@ -41,15 +46,19 @@ export class GitHubActionsRunnerConnectionService { : session.workPhase; const status = session.status === "attached" || session.status === "detached" ? session.status : "ready"; - await this.store.persist(session.id, { - status, - work_state: state, - work_phase: phase, - last_heartbeat_at: now, - last_seen_at: now, - updated_at: now, - last_event: githubActionsRunnerConnectedEvent, - }); + await this.store.persist( + session.id, + { + status, + work_state: state, + work_phase: phase, + last_heartbeat_at: now, + last_seen_at: now, + updated_at: revision, + last_event: githubActionsRunnerConnectedEvent, + }, + session.updatedAt, + ); await this.store.appendEvent(session.id, githubActionsRunnerConnectedEvent, now); } } diff --git a/src/worker/github-actions-session-registration.ts b/src/worker/github-actions-session-registration.ts index cf6b640d..b3e0bf28 100644 --- a/src/worker/github-actions-session-registration.ts +++ b/src/worker/github-actions-session-registration.ts @@ -47,6 +47,11 @@ export type GitHubActionsSessionRegistrationUpdate = { completion_reason: null; }; +export type GitHubActionsSessionRegistrationExpectation = Pick< + InteractiveSessionRow, + "agent_token_hash" | "updated_at" | "status" | "work_state" | "work_phase" +>; + export type GitHubActionsSessionRegistrationStore = { now(): number; newAgentToken(): string; @@ -57,7 +62,11 @@ export type GitHubActionsSessionRegistrationStore = { nextSessionId(): Promise; insertSession(values: Insertable): Promise; readById(id: string): Promise; - updateSession(id: string, values: GitHubActionsSessionRegistrationUpdate): Promise; + updateSession( + id: string, + values: GitHubActionsSessionRegistrationUpdate, + expected: GitHubActionsSessionRegistrationExpectation, + ): Promise; isConstraintError(error: unknown): boolean; disconnectRunner(id: string): Promise; appendEvent(id: string, message: string, now: number): Promise; @@ -148,33 +157,44 @@ export class GitHubActionsSessionRegistrationService { const resumed = existing.work_state !== "registered" || existing.status !== "ready"; const message = resumed ? "GitHub Actions work resumed" : "GitHub Actions work registered"; - await this.store.updateSession(existing.id, { - owner, - owner_subject: ownerSubject, - repo, - branch, - purpose, - summary, - prompt: purpose, - status: "ready", - lease_id: null, - stopped_at: null, - terminal_status: null, - terminal_failure_reason: null, - terminal_finalize_pending: 0, - credential_cleanup_terminal_status: null, - updated_at: now, - last_seen_at: now, - last_event: message, - agent_token_hash: agentTokenHash, - work_kind: workKind, - work_state: "registered", - work_phase: "waiting_for_runner", - source_url: input.sourceUrl === undefined ? existing.source_url : sourceUrl, - github_run_url: input.runUrl === undefined ? existing.github_run_url : runUrl, - last_heartbeat_at: null, - completion_reason: null, - }); + const registrationRevision = Math.max(existing.updated_at + 1, now); + await this.store.updateSession( + existing.id, + { + owner, + owner_subject: ownerSubject, + repo, + branch, + purpose, + summary, + prompt: purpose, + status: "ready", + lease_id: null, + stopped_at: null, + terminal_status: null, + terminal_failure_reason: null, + terminal_finalize_pending: 0, + credential_cleanup_terminal_status: null, + updated_at: registrationRevision, + last_seen_at: now, + last_event: message, + agent_token_hash: agentTokenHash, + work_kind: workKind, + work_state: "registered", + work_phase: "waiting_for_runner", + source_url: input.sourceUrl === undefined ? existing.source_url : sourceUrl, + github_run_url: input.runUrl === undefined ? existing.github_run_url : runUrl, + last_heartbeat_at: null, + completion_reason: null, + }, + { + agent_token_hash: existing.agent_token_hash, + updated_at: existing.updated_at, + status: existing.status, + work_state: existing.work_state, + work_phase: existing.work_phase, + }, + ); await this.store.disconnectRunner(existing.id).catch(() => undefined); await this.store.appendEvent(existing.id, message, now); await this.store.audit( diff --git a/src/worker/github-actions-session-work-state.ts b/src/worker/github-actions-session-work-state.ts index a19f0a10..a02e89e1 100644 --- a/src/worker/github-actions-session-work-state.ts +++ b/src/worker/github-actions-session-work-state.ts @@ -38,7 +38,12 @@ export type GitHubActionsWorkStateUpdate = { export type GitHubActionsWorkStateStore = { now(): number; readRow(id: string): Promise; - persist(id: string, values: GitHubActionsWorkStateUpdate): Promise; + persist( + id: string, + values: GitHubActionsWorkStateUpdate, + expectedRevision: number, + expectedTerminalStatus?: InteractiveSessionStatus, + ): Promise; appendEvent(id: string, message: string, now: number): Promise; disconnectRunner(id: string): Promise; readSession(id: string): Promise; @@ -90,21 +95,27 @@ export class GitHubActionsWorkStateService { row.codex_turn_id !== codexTurnId || row.completion_reason !== completionReason; const now = this.store.now(); + const revision = Math.max(session.updatedAt + 1, now); - await this.store.persist(session.id, { - status, - summary, - work_state: state, - work_phase: phase, - codex_thread_id: codexThreadId, - codex_turn_id: codexTurnId, - last_heartbeat_at: now, - completion_reason: completionReason, - last_event: lastEvent, - last_seen_at: now, - updated_at: now, - stopped_at: terminal ? now : null, - }); + await this.store.persist( + session.id, + { + status, + summary, + work_state: state, + work_phase: phase, + codex_thread_id: codexThreadId, + codex_turn_id: codexTurnId, + last_heartbeat_at: now, + completion_reason: completionReason, + last_event: lastEvent, + last_seen_at: now, + updated_at: revision, + stopped_at: terminal ? now : null, + }, + session.updatedAt, + terminal ? row.status : undefined, + ); if (changed) { await this.store.appendEvent(session.id, lastEvent, now); } diff --git a/src/worker/http.ts b/src/worker/http.ts index 724097b4..a14258c2 100644 --- a/src/worker/http.ts +++ b/src/worker/http.ts @@ -54,11 +54,21 @@ export function wantsMarkdown(request: Request): boolean { } export async function readJson(request: Request): Promise { + let source: string; try { - return (await request.json()) as T; + source = await request.text(); } catch { throw badRequest("invalid json"); } + let parsed: unknown; + try { + parsed = JSON.parse(source) as unknown; + } catch { + throw badRequest("invalid json"); + } + assertRoundTrippableJsonIntegerLexemes(source); + assertRoundTrippableJsonIntegers(parsed); + return parsed as T; } export async function readBoundedJson(request: Request, maximumBytes: number): Promise { @@ -97,11 +107,16 @@ export async function readBoundedJson(request: Request, maximumBytes: number) bytes.set(chunk, offset); offset += chunk.byteLength; } + const source = new TextDecoder().decode(bytes); + let parsed: unknown; try { - return JSON.parse(new TextDecoder().decode(bytes)) as T; + parsed = JSON.parse(source) as unknown; } catch { throw badRequest("invalid json"); } + assertRoundTrippableJsonIntegerLexemes(source); + assertRoundTrippableJsonIntegers(parsed); + return parsed as T; } export function bearerToken(request: Request): string { @@ -170,3 +185,92 @@ function clean(value: unknown, maximum: number): string { .trim() .slice(0, maximum); } + +function assertRoundTrippableJsonIntegers(value: unknown): void { + const pending = [value]; + while (pending.length > 0) { + const current = pending.pop(); + if (typeof current === "number") { + if ( + !Number.isFinite(current) || + (Number.isInteger(current) && (!Number.isSafeInteger(current) || Object.is(current, -0))) + ) { + throw badRequest("json integers must be safe and round-trippable"); + } + continue; + } + if (!current || typeof current !== "object") continue; + if (Array.isArray(current)) { + for (const item of current) pending.push(item); + continue; + } + for (const item of Object.values(current)) pending.push(item); + } +} + +function assertRoundTrippableJsonIntegerLexemes(source: string): void { + const numberPattern = /-?(?:0|[1-9]\d*)(?:\.\d+)?(?:[eE][+-]?\d+)?/y; + let inString = false; + let escaped = false; + for (let index = 0; index < source.length; index += 1) { + const character = source[index]!; + if (inString) { + if (escaped) { + escaped = false; + } else if (character === "\\") { + escaped = true; + } else if (character === '"') { + inString = false; + } + continue; + } + if (character === '"') { + inString = true; + continue; + } + if (character !== "-" && (character < "0" || character > "9")) continue; + numberPattern.lastIndex = index; + const match = numberPattern.exec(source); + if (!match) continue; + const token = match[0]; + const value = Number(token); + if ( + Number.isInteger(value) && + (!Number.isSafeInteger(value) || + Object.is(value, -0) || + exactJsonInteger(token) !== String(value)) + ) { + throw badRequest("json integers must be safe and round-trippable"); + } + index = numberPattern.lastIndex - 1; + } +} + +function exactJsonInteger(token: string): string | null { + const negative = token.startsWith("-"); + const unsigned = negative ? token.slice(1) : token; + const exponentIndex = unsigned.search(/[eE]/u); + const mantissa = exponentIndex === -1 ? unsigned : unsigned.slice(0, exponentIndex); + const exponentText = exponentIndex === -1 ? "" : unsigned.slice(exponentIndex + 1); + const decimalIndex = mantissa.indexOf("."); + const integerDigits = decimalIndex === -1 ? mantissa.length : decimalIndex; + const digits = + decimalIndex === -1 + ? mantissa + : mantissa.slice(0, decimalIndex) + mantissa.slice(decimalIndex + 1); + if (/^0+$/u.test(digits)) return negative ? "-0" : "0"; + + const exponent = exponentText ? Number(exponentText) : 0; + if (!Number.isSafeInteger(exponent)) return null; + const decimalPosition = integerDigits + exponent; + if (decimalPosition <= 0) return null; + if (decimalPosition < digits.length && !/^0+$/u.test(digits.slice(decimalPosition))) { + return null; + } + const exactDigits = + decimalPosition >= digits.length + ? digits + "0".repeat(decimalPosition - digits.length) + : digits.slice(0, decimalPosition); + const canonicalDigits = exactDigits.replace(/^0+/u, "") || "0"; + return negative ? `-${canonicalDigits}` : canonicalDigits; +} diff --git a/src/worker/interactive-session-application.ts b/src/worker/interactive-session-application.ts index 4739ae7a..c2708924 100644 --- a/src/worker/interactive-session-application.ts +++ b/src/worker/interactive-session-application.ts @@ -24,7 +24,10 @@ import { sandboxLeasePrefix, } from "./sandbox-lease.ts"; import { ServiceRegistry } from "./service-registry.ts"; -import { appendInteractiveSessionEventRecord } from "./session-events.ts"; +import { + appendInteractiveSessionEventRecord, + persistInteractiveSessionEventRecord, +} from "./session-events.ts"; import { canChangeInteractiveSessionMultiplayer, canControlInteractiveSession, @@ -480,7 +483,7 @@ export class InteractiveSessionApplication { activateReservation: (insertedSessionId, insertedAt, adapterWorkspaceId) => supervision.requireReservationActivation(insertedSessionId, insertedAt, adapterWorkspaceId), recordRequest: (insertedSessionId, insertedAt) => - appendInteractiveSessionEventRecord(this.env, { + persistInteractiveSessionEventRecord(this.env, { sessionId: insertedSessionId, actor: actor(user), message: "interactive workspace requested", @@ -504,10 +507,11 @@ export class InteractiveSessionApplication { finalizeTerminal: (sessionId, status, now) => finalizeTerminalInteractiveSession(this.env, sessionId, status, now), readSession: (sessionId) => readInteractiveSessionRecord(this.env, sessionId), - stopSupersededAdapter: (sessionId, adapterWorkspaceId, createPending, now) => + stopSupersededAdapter: (sessionId, adapterWorkspaceId, registration, createPending, now) => this.runtime.release().stopSuperseded({ sessionId, adapterWorkspaceId, + registration, createPending, now, }), diff --git a/src/worker/interactive-terminal-service.ts b/src/worker/interactive-terminal-service.ts index 266a2a58..bea36e7e 100644 --- a/src/worker/interactive-terminal-service.ts +++ b/src/worker/interactive-terminal-service.ts @@ -1,4 +1,3 @@ -import { getSandbox } from "@cloudflare/sandbox"; import { terminalOutputAcknowledgements } from "@openclaw/libterminal/worker"; import { @@ -7,7 +6,13 @@ import { terminalSubmittedLine, type TerminalInputState, } from "../terminal-multiplayer.ts"; -import { githubActionsRuntime } from "../github-actions-runtime.ts"; +import { + buildGitHubActionsViewerRelayUrl, + gitHubActionsViewerResponseGeneration, + gitHubActionsViewerResponseUsesFramedProtocol, + gitHubActionsViewerResponseUsesGenerations, + githubActionsRuntime, +} from "../github-actions-runtime.ts"; import { terminalFailureStatusForAdapter } from "../runtime-adapter.ts"; import { cachedBooleanGrant } from "../terminal-authorization.ts"; import { actor, requireRole } from "./auth.ts"; @@ -19,7 +24,6 @@ import { isOpenClawEmbedSessionToken, terminalInputAuthorization, } from "./openclaw-embed-access.ts"; -import { SandboxLifecycleService } from "./provisioning/sandbox-lifecycle.ts"; import { isSandboxLeaseOwnerReconnectError } from "./provisioning/sandbox.ts"; import { readTerminalClipboardBytes, @@ -28,7 +32,6 @@ import { terminalClipboardMaxBytes, } from "./interactive-terminal.ts"; import { InteractiveTerminalRepository } from "./interactive-terminal-repository.ts"; -import { reconcileSandboxCredentialPolicyCleanupBatch } from "./sandbox-credential-policy-cleanup-service.ts"; import { stageTerminalCredentialPolicyCleanupById } from "./sandbox-credential-policy-cleanup.ts"; import { isSandboxInteractiveSession, sandboxLeaseInfo } from "./sandbox-lease.ts"; import { openSandboxTerminalResponse, sandboxWorkdir } from "./sandbox-runtime.ts"; @@ -56,7 +59,38 @@ import { import { interactiveTerminalFetch } from "./runtime-adapter-transport.ts"; import { tenancyMode, tenantSubject } from "./tenancy.ts"; -const terminalInputStates = new Map(); +export class TerminalInputStateRegistry { + private readonly entries = new Map(); + + retain(sessionId: string): void { + const entry = this.entry(sessionId); + entry.subscribers += 1; + } + + release(sessionId: string): void { + const entry = this.entries.get(sessionId); + if (!entry || entry.subscribers <= 1) { + this.entries.delete(sessionId); + return; + } + entry.subscribers -= 1; + } + + state(sessionId: string): TerminalInputState { + return this.entry(sessionId).state; + } + + private entry(sessionId: string): { state: TerminalInputState; subscribers: number } { + let entry = this.entries.get(sessionId); + if (!entry) { + entry = { state: newTerminalInputState(), subscribers: 0 }; + this.entries.set(sessionId, entry); + } + return entry; + } +} + +const terminalInputStates = new TerminalInputStateRegistry(); export type InteractiveTerminalServiceDependencies = { readSession(sessionId: string): Promise; @@ -98,10 +132,9 @@ export class InteractiveTerminalService { if (session.runtime === githubActionsRuntime) { const stub = githubActionsRelayStub(this.env, session.id); if (!stub) throw serviceUnavailable("SESSION_CONTROL Durable Object is not configured"); - const upstreamResponse = await stub.fetch( - "https://crabfleet.internal/api/session-control/github-actions/viewer", - { headers: { upgrade: "websocket" } }, - ); + const upstreamResponse = await stub.fetch(buildGitHubActionsViewerRelayUrl(), { + headers: { upgrade: "websocket" }, + }); const upstream = upstreamResponse.webSocket; if (!upstream || upstreamResponse.status !== 101) { throw serviceUnavailable(`GitHub Actions relay HTTP ${upstreamResponse.status}`); @@ -109,6 +142,9 @@ export class InteractiveTerminalService { upstream.accept(); return { socket: upstream, + inputAcknowledgements: gitHubActionsViewerResponseUsesFramedProtocol(upstreamResponse), + inputGenerations: gitHubActionsViewerResponseUsesGenerations(upstreamResponse), + initialRunnerGeneration: gitHubActionsViewerResponseGeneration(upstreamResponse), outputAcknowledgements: false, markConnected: () => markInteractiveTerminalConnected( @@ -125,12 +161,14 @@ export class InteractiveTerminalService { const routeKind = interactivePtyRouteKind(this.env, session); if (routeKind === "sandbox" && this.env.SANDBOX) { const runtimeSession = await this.dependencies.resolveSandboxSession(request, user, session); + const { SandboxLifecycleService } = await import("./provisioning/sandbox-lifecycle.ts"); const sandboxSession = await new SandboxLifecycleService(this.env).ensureCurrentLease( request, user, runtimeSession, ); const lease = sandboxLeaseInfo(sandboxSession); + const { getSandbox } = await import("@cloudflare/sandbox"); const sandbox = getSandbox(this.env.SANDBOX, lease.sandboxId); const upstreamResponse = await openSandboxTerminalResponse( request, @@ -212,6 +250,7 @@ export class InteractiveTerminalService { } private terminalHub(): TerminalHub { + const retainedInputSessions = new Set(); return new TerminalHub({ createSocketPair: () => { const pair = new WebSocketPair(); @@ -241,11 +280,25 @@ export class InteractiveTerminalService { terminalViewGrant(request, this.env, this.repository, user, session), reconcileSubscription: (sessionId) => terminalSubscriptionReconciler(this.dependencies, sessionId), - openUpstream: (request, user, session, cols, rows) => - this.openUpstream(request, user, session, cols, rows), + openUpstream: async (request, user, session, cols, rows) => { + const upstream = await this.openUpstream(request, user, session, cols, rows); + return { + ...upstream, + markConnected: async () => { + if (!retainedInputSessions.has(session.id)) { + retainedInputSessions.add(session.id); + terminalInputStates.retain(session.id); + } + await upstream.markConnected(); + }, + }; + }, inputPayloads: (subscription, user, payload) => multiplayerTerminalInputPayloads(this.repository, subscription, user, payload), - releaseInputState: releaseTerminalInputState, + releaseInputState: (sessionId) => { + if (!retainedInputSessions.delete(sessionId)) return; + terminalInputStates.release(sessionId); + }, markConnectionFailure: async (user, session, message, error) => { if (isSandboxLeaseOwnerReconnectError(error)) return; const markTerminal = @@ -285,6 +338,7 @@ async function writeTerminalClipboardFile( const mediaType = clean(rawMediaType || "application/octet-stream", 120); const name = terminalClipboardFilename(rawName, mediaType); const lease = sandboxLeaseInfo(session); + const { getSandbox } = await import("@cloudflare/sandbox"); const sandbox = getSandbox(env.SANDBOX, lease.sandboxId); const directory = `${sandboxWorkdir(session.id)}/.crabbox/clipboard`; const path = `${directory}/${Date.now()}-${crypto.randomUUID().slice(0, 8)}-${name}`; @@ -357,6 +411,8 @@ async function markInteractiveTerminalUnavailable( ); if (!staged) return; await appendTerminalLog(env, sessionId, user, message, now); + const { reconcileSandboxCredentialPolicyCleanupBatch } = + await import("./sandbox-credential-policy-cleanup-service.ts"); await reconcileSandboxCredentialPolicyCleanupBatch(env, now, sessionId); return; } @@ -544,16 +600,7 @@ async function multiplayerTerminalInputPayloads( } function terminalInputState(sessionId: string): TerminalInputState { - let state = terminalInputStates.get(sessionId); - if (!state) { - state = newTerminalInputState(); - terminalInputStates.set(sessionId, state); - } - return state; -} - -function releaseTerminalInputState(sessionId: string): void { - terminalInputStates.delete(sessionId); + return terminalInputStates.state(sessionId); } async function readInteractiveSessionMultiplayerMode( diff --git a/src/worker/openclaw-repository.ts b/src/worker/openclaw-repository.ts index 11214318..035df816 100644 --- a/src/worker/openclaw-repository.ts +++ b/src/worker/openclaw-repository.ts @@ -3,6 +3,7 @@ import { sql } from "kysely"; import { database, executeBatch, type InteractiveSessionRow } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import { interactiveSession, type InteractiveSession } from "./session-model.ts"; +import { cleanupSessionLogArchiveObjects } from "./session-log-archive.ts"; export type OpenClawRoomSessions = { sessions: InteractiveSession[]; @@ -172,35 +173,70 @@ export async function removeInteractiveSessionReservation( insertedAt: number, ): Promise { const db = database(env); - const ownsReservation = sql`EXISTS ( + const rollbackClaim = `reservation-rollback:${insertedAt}`; + const rollbackClaimedAt = insertedAt + 1; + await db + .updateTable("interactive_sessions") + .set({ + reconcile_error: rollbackClaim, + updated_at: rollbackClaimedAt, + }) + .where("id", "=", insertedSessionId) + .where("status", "=", "provisioning") + .where("preparation_pending", "=", 1) + .where("created_at", "=", insertedAt) + .where("updated_at", "=", insertedAt) + .execute(); + const claimed = await db + .selectFrom("interactive_sessions") + .select("id") + .where("id", "=", insertedSessionId) + .where("status", "=", "provisioning") + .where("preparation_pending", "=", 1) + .where("created_at", "=", insertedAt) + .where("updated_at", "=", rollbackClaimedAt) + .where("reconcile_error", "=", rollbackClaim) + .executeTakeFirst(); + if (!claimed) return false; + + const archive = await db + .selectFrom("interactive_session_log_archives") + .select(["events_key", "transcript_key", "summary_key"]) + .where("session_id", "=", insertedSessionId) + .executeTakeFirst(); + await cleanupSessionLogArchiveObjects(env, archive); + + const ownsRollbackClaim = sql`EXISTS ( SELECT 1 FROM interactive_sessions WHERE id = ${insertedSessionId} AND status = 'provisioning' AND preparation_pending = 1 AND created_at = ${insertedAt} - AND updated_at = ${insertedAt} + AND updated_at = ${rollbackClaimedAt} + AND reconcile_error = ${rollbackClaim} )`; await executeBatch(env, [ db .deleteFrom("openclaw_request_replays") .where("session_id", "=", insertedSessionId) - .where(ownsReservation), + .where(ownsRollbackClaim), db .deleteFrom("interactive_session_events") .where("session_id", "=", insertedSessionId) - .where(ownsReservation), + .where(ownsRollbackClaim), db .deleteFrom("interactive_session_log_archives") .where("session_id", "=", insertedSessionId) - .where(ownsReservation), + .where(ownsRollbackClaim), db .deleteFrom("interactive_sessions") .where("id", "=", insertedSessionId) .where("status", "=", "provisioning") .where("preparation_pending", "=", 1) .where("created_at", "=", insertedAt) - .where("updated_at", "=", insertedAt), + .where("updated_at", "=", rollbackClaimedAt) + .where("reconcile_error", "=", rollbackClaim), ]); const current = await db .selectFrom("interactive_sessions") diff --git a/src/worker/provisioning/runtime-adapter-release-repository.ts b/src/worker/provisioning/runtime-adapter-release-repository.ts new file mode 100644 index 00000000..2441eb00 --- /dev/null +++ b/src/worker/provisioning/runtime-adapter-release-repository.ts @@ -0,0 +1,184 @@ +import { sql } from "kysely"; + +import { database, type RuntimeAdapterWorkspaceCleanupRow } from "../database.ts"; +import type { RuntimeEnv } from "../env.ts"; +import type { + RuntimeAdapterWorkspaceCleanup, + RuntimeAdapterWorkspaceRegistration, +} from "./runtime-adapter-release-service.ts"; + +const cleanupClaimTtlMs = 60_000; +const cleanupRetryDelayMs = 15_000; + +export async function stageRuntimeAdapterWorkspaceCleanup( + env: RuntimeEnv, + input: { + sessionId: string; + adapterWorkspaceId: string; + registration: RuntimeAdapterWorkspaceRegistration | null; + createPending: boolean; + now: number; + }, +): Promise { + await sql` + INSERT INTO runtime_adapter_workspace_cleanups ( + session_id, + adapter_workspace_id, + profile, + control_plane, + create_pending, + message, + reconcile_error, + next_attempt_at, + created_at, + updated_at + ) VALUES ( + ${input.sessionId}, + ${input.adapterWorkspaceId}, + ${input.registration?.profile ?? null}, + ${input.registration?.controlPlane ?? null}, + ${input.createPending ? 1 : 0}, + 'superseded runtime adapter cleanup pending', + NULL, + ${input.now}, + ${input.now}, + ${input.now} + ) + ON CONFLICT(session_id, adapter_workspace_id) DO NOTHING + `.execute(database(env)); +} + +export async function claimRuntimeAdapterWorkspaceCleanup( + env: RuntimeEnv, + sessionId: string, + adapterWorkspaceId: string, + now: number, +): Promise { + const claim = `runtime-cleanup:${crypto.randomUUID()}`; + const row = await database(env) + .updateTable("runtime_adapter_workspace_cleanups") + .set({ + cleanup_claim: claim, + cleanup_claim_expires_at: now + cleanupClaimTtlMs, + attempt_count: sql`attempt_count + 1`, + last_attempt_at: now, + updated_at: sql`MAX(updated_at + 1, ${now})`, + }) + .where("session_id", "=", sessionId) + .where("adapter_workspace_id", "=", adapterWorkspaceId) + .where("next_attempt_at", "<=", now) + .where((expression) => + expression.or([ + expression("cleanup_claim", "is", null), + expression("cleanup_claim_expires_at", "<=", now), + ]), + ) + .returningAll() + .executeTakeFirst(); + return row ? cleanupClaim(row) : null; +} + +export async function claimRuntimeAdapterWorkspaceCleanupBatch( + env: RuntimeEnv, + now: number, + limit: number, +): Promise { + const candidates = await database(env) + .selectFrom("runtime_adapter_workspace_cleanups") + .select(["session_id", "adapter_workspace_id"]) + .where("next_attempt_at", "<=", now) + .where((expression) => + expression.or([ + expression("cleanup_claim", "is", null), + expression("cleanup_claim_expires_at", "<=", now), + ]), + ) + .orderBy("next_attempt_at", "asc") + .orderBy("updated_at", "asc") + .orderBy("session_id", "asc") + .orderBy("adapter_workspace_id", "asc") + .limit(limit) + .execute(); + const claims: RuntimeAdapterWorkspaceCleanup[] = []; + for (const candidate of candidates) { + const claimed = await claimRuntimeAdapterWorkspaceCleanup( + env, + candidate.session_id, + candidate.adapter_workspace_id, + now, + ); + if (claimed) claims.push(claimed); + } + return claims; +} + +export async function persistRuntimeAdapterWorkspaceCleanupEvidence( + env: RuntimeEnv, + cleanup: RuntimeAdapterWorkspaceCleanup, + message: string, + now: number, + reconcileError: string | null, +): Promise { + await database(env) + .updateTable("runtime_adapter_workspace_cleanups") + .set({ + message, + reconcile_error: reconcileError, + next_attempt_at: now + cleanupRetryDelayMs, + cleanup_claim: null, + cleanup_claim_expires_at: null, + updated_at: sql`MAX(updated_at + 1, ${now})`, + }) + .where("session_id", "=", cleanup.sessionId) + .where("adapter_workspace_id", "=", cleanup.adapterWorkspaceId) + .where("cleanup_claim", "=", cleanup.claim) + .execute(); +} + +export async function markRuntimeAdapterWorkspaceCleanupDeletionObserved( + env: RuntimeEnv, + cleanup: RuntimeAdapterWorkspaceCleanup, + now: number, +): Promise { + const row = await database(env) + .updateTable("runtime_adapter_workspace_cleanups") + .set({ + deletion_observed: 1, + updated_at: sql`MAX(updated_at + 1, ${now})`, + }) + .where("session_id", "=", cleanup.sessionId) + .where("adapter_workspace_id", "=", cleanup.adapterWorkspaceId) + .where("cleanup_claim", "=", cleanup.claim) + .returning("deletion_observed") + .executeTakeFirst(); + if (!row) throw new Error("runtime adapter cleanup ownership changed"); +} + +export async function completeRuntimeAdapterWorkspaceCleanup( + env: RuntimeEnv, + cleanup: RuntimeAdapterWorkspaceCleanup, +): Promise { + await database(env) + .deleteFrom("runtime_adapter_workspace_cleanups") + .where("session_id", "=", cleanup.sessionId) + .where("adapter_workspace_id", "=", cleanup.adapterWorkspaceId) + .where("cleanup_claim", "=", cleanup.claim) + .execute(); +} + +function cleanupClaim(row: RuntimeAdapterWorkspaceCleanupRow): RuntimeAdapterWorkspaceCleanup { + return { + sessionId: row.session_id, + adapterWorkspaceId: row.adapter_workspace_id, + registration: + row.profile && row.control_plane + ? { + profile: row.profile, + controlPlane: row.control_plane, + } + : null, + createPending: row.create_pending === 1, + deletionObserved: row.deletion_observed === 1, + claim: row.cleanup_claim ?? "", + }; +} diff --git a/src/worker/provisioning/runtime-adapter-release-service.ts b/src/worker/provisioning/runtime-adapter-release-service.ts index cba571be..5bd2add3 100644 --- a/src/worker/provisioning/runtime-adapter-release-service.ts +++ b/src/worker/provisioning/runtime-adapter-release-service.ts @@ -1,10 +1,47 @@ import type { RuntimeAdapterWorkspaceStopResult } from "../session-runtime-adapter-stop.ts"; +export type RuntimeAdapterWorkspaceRegistration = { + profile: string; + controlPlane: string; +}; + +export type RuntimeAdapterWorkspaceCleanup = { + sessionId: string; + adapterWorkspaceId: string; + registration: RuntimeAdapterWorkspaceRegistration | null; + createPending: boolean; + deletionObserved: boolean; + claim: string; +}; + export type RuntimeAdapterReleaseServiceDependencies = { + stageCleanup(input: { + sessionId: string; + adapterWorkspaceId: string; + registration: RuntimeAdapterWorkspaceRegistration | null; + createPending: boolean; + now: number; + }): Promise; + claimCleanup( + sessionId: string, + adapterWorkspaceId: string, + now: number, + ): Promise; + claimPendingCleanups(now: number): Promise; + persistCleanupEvidence( + cleanup: RuntimeAdapterWorkspaceCleanup, + message: string, + now: number, + reconcileError: string | null, + ): Promise; + markCleanupDeletionObserved(cleanup: RuntimeAdapterWorkspaceCleanup, now: number): Promise; + completeCleanup(cleanup: RuntimeAdapterWorkspaceCleanup): Promise; clearCreatePending(sessionId: string, adapterWorkspaceId: string): Promise; stopWorkspace( sessionId: string, adapterWorkspaceId: string, + registration: RuntimeAdapterWorkspaceRegistration | null, + retryMissing: boolean, ): Promise; confirmRelease( sessionId: string, @@ -32,35 +69,76 @@ export class RuntimeAdapterReleaseService { async stopSuperseded(input: { sessionId: string; adapterWorkspaceId: string; + registration: RuntimeAdapterWorkspaceRegistration | null; createPending: boolean; now: number; }): Promise { - const { sessionId, adapterWorkspaceId, createPending, now } = input; - if (!createPending) { - await this.dependencies.clearCreatePending(sessionId, adapterWorkspaceId); + await this.dependencies.stageCleanup(input); + const cleanup = await this.dependencies.claimCleanup( + input.sessionId, + input.adapterWorkspaceId, + input.now, + ); + if (!cleanup) return; + await this.releaseCleanup(cleanup, input.now); + } + + async retryPending(now: number): Promise { + const cleanups = await this.dependencies.claimPendingCleanups(now); + for (const cleanup of cleanups) { + await this.releaseCleanup(cleanup, now); } + } + + private async releaseCleanup( + cleanup: RuntimeAdapterWorkspaceCleanup, + now: number, + ): Promise { + const { sessionId, adapterWorkspaceId, registration, createPending, deletionObserved } = + cleanup; try { - const release = await this.dependencies.stopWorkspace(sessionId, adapterWorkspaceId); - if (release.status === "stopped") { - await this.dependencies.confirmRelease(sessionId, adapterWorkspaceId, now, release.message); - return; + if (!createPending) { + await this.dependencies.clearCreatePending(sessionId, adapterWorkspaceId); } - await this.dependencies.persistStopEvidence( + const release = await this.dependencies.stopWorkspace( sessionId, adapterWorkspaceId, - release.message, - now, - null, + registration, + createPending && !deletionObserved, ); + if (release.status === "stopped") { + await this.dependencies.markCleanupDeletionObserved(cleanup, now); + await this.dependencies.confirmRelease(sessionId, adapterWorkspaceId, now, release.message); + await this.dependencies.completeCleanup(cleanup); + return; + } + await this.persistEvidence(cleanup, release.message, now, null); } catch (error) { const message = this.dependencies.providerError(error, adapterWorkspaceId); - await this.dependencies.persistStopEvidence( - sessionId, - adapterWorkspaceId, + await this.persistEvidence( + cleanup, `superseded runtime adapter stop pending: ${message}`, now, message, ); } } + + private async persistEvidence( + cleanup: RuntimeAdapterWorkspaceCleanup, + message: string, + now: number, + reconcileError: string | null, + ): Promise { + await this.dependencies.persistCleanupEvidence(cleanup, message, now, reconcileError); + await this.dependencies + .persistStopEvidence( + cleanup.sessionId, + cleanup.adapterWorkspaceId, + message, + now, + reconcileError, + ) + .catch(() => undefined); + } } diff --git a/src/worker/provisioning/runtime-adapter-repository.ts b/src/worker/provisioning/runtime-adapter-repository.ts index b32ad377..6774b711 100644 --- a/src/worker/provisioning/runtime-adapter-repository.ts +++ b/src/worker/provisioning/runtime-adapter-repository.ts @@ -362,13 +362,18 @@ export async function clearRuntimeAdapterCreatePending( sessionId: string, adapterWorkspaceId: string, ): Promise { + const now = Date.now(); await database(env) .updateTable("interactive_sessions") - .set({ adapter_create_pending: 0 }) + .set({ + adapter_create_pending: 0, + updated_at: sql`MAX(updated_at + 1, ${now})`, + }) .where("id", "=", sessionId) .where("adapter", "=", runtimeAdapterName) .where("adapter_workspace_id", "=", adapterWorkspaceId) .where("status", "=", "stopping") + .where("adapter_create_pending", "=", 1) .execute(); } diff --git a/src/worker/provisioning/sandbox-repository.ts b/src/worker/provisioning/sandbox-repository.ts index 6761afd8..b4f62eeb 100644 --- a/src/worker/provisioning/sandbox-repository.ts +++ b/src/worker/provisioning/sandbox-repository.ts @@ -237,7 +237,11 @@ export async function commitManagedSandboxLeaseRefresh( db .updateTable("interactive_sessions") .set({ - status: provisioned.status, + status: sql`CASE + WHEN ${provisioned.status} = 'ready' AND status IN ('attached', 'detached') + THEN status + ELSE ${provisioned.status} + END`, lease_id: provisioned.leaseId, attach_url: provisioned.attachUrl, vnc_url: provisioned.vncUrl, diff --git a/src/worker/routes/control-plane.ts b/src/worker/routes/control-plane.ts index 59eaa8ea..e8cc15d8 100644 --- a/src/worker/routes/control-plane.ts +++ b/src/worker/routes/control-plane.ts @@ -5,15 +5,34 @@ import type { AdminRepoInput, AdminWorkflowInput, } from "../admin-service.ts"; -import type { DesktopHost, DesktopHostInput } from "../desktop-host-service.ts"; -import { json, notFound, readJson } from "../http.ts"; +import { + desktopHostOwnershipHeader, + desktopHostOwnershipModeHeader, + desktopHostPublicationHeader, + desktopHostTokenOwnershipMode, + type DesktopHostInput, + type DesktopHostOwnershipMode, + type DesktopHostRegistration, +} from "../desktop-host-service.ts"; +import { badRequest, json, notFound, readJson } from "../http.ts"; import type { User } from "../models.ts"; export type ControlPlaneRouteDependencies = { readState(request: Request, user: User): Promise; readFleet(user: User): Promise; - registerDesktopHost(user: User, id: string, input: DesktopHostInput): Promise; - removeDesktopHost(user: User, id: string): Promise; + registerDesktopHost( + user: User, + id: string, + input: DesktopHostInput, + ownershipMode: DesktopHostOwnershipMode, + publicationID: string | null, + ): Promise; + recoverDesktopHost( + user: User, + id: string, + publicationID: unknown, + ): Promise<{ ownershipToken: string | null }>; + removeDesktopHost(user: User, id: string, ownershipToken: string | null): Promise; searchGitHubRefs(number: unknown): Promise; createCard(request: Request, user: User): Promise; readCardRuns(user: User, cardId: string): Promise; @@ -44,16 +63,28 @@ export async function handleControlPlaneRoute( const desktopHostMatch = url.pathname.match(/^\/api\/desktop-hosts\/([^/]+)$/); if (request.method === "PUT" && desktopHostMatch) { requireRole(user, "viewer"); - const host = await dependencies.registerDesktopHost( + const registration = await dependencies.registerDesktopHost( user, decoded(desktopHostMatch[1]), await readJson(request), + request.headers.get(desktopHostOwnershipModeHeader) === desktopHostTokenOwnershipMode + ? desktopHostTokenOwnershipMode + : "legacy", + request.headers.get(desktopHostPublicationHeader), + ); + return json(registration); + } + if (request.method === "POST" && desktopHostMatch && url.searchParams.get("recover") === "1") { + requireRole(user, "viewer"); + const body = await readJson<{ publicationID?: unknown }>(request); + return json( + await dependencies.recoverDesktopHost(user, decoded(desktopHostMatch[1]), body.publicationID), ); - return json({ host }); } if (request.method === "DELETE" && desktopHostMatch) { requireRole(user, "viewer"); - await dependencies.removeDesktopHost(user, decoded(desktopHostMatch[1])); + const ownershipToken = request.headers.get(desktopHostOwnershipHeader); + await dependencies.removeDesktopHost(user, decoded(desktopHostMatch[1]), ownershipToken); return json({ ok: true }); } if (request.method === "GET" && url.pathname === "/api/github/refs") { @@ -123,5 +154,9 @@ export async function handleControlPlaneRoute( } function decoded(value: string | undefined): string { - return decodeURIComponent(value ?? ""); + try { + return decodeURIComponent(value ?? ""); + } catch { + throw badRequest("invalid path identifier"); + } } diff --git a/src/worker/routes/service-sessions.ts b/src/worker/routes/service-sessions.ts index 9271fdd6..db69f995 100644 --- a/src/worker/routes/service-sessions.ts +++ b/src/worker/routes/service-sessions.ts @@ -1,4 +1,4 @@ -import { json } from "../http.ts"; +import { badRequest, json } from "../http.ts"; import type { User } from "../models.ts"; import type { InteractiveSession } from "../session-model.ts"; import type { InteractiveSessionSummaryInput } from "../session-metadata.ts"; @@ -92,5 +92,9 @@ export async function handleServiceSessionRoute( } function decoded(value: string | undefined): string { - return decodeURIComponent(value ?? ""); + try { + return decodeURIComponent(value ?? ""); + } catch { + throw badRequest("invalid session id"); + } } diff --git a/src/worker/runtime-adapter-workspaces.ts b/src/worker/runtime-adapter-workspaces.ts index 1e911275..123c64a8 100644 --- a/src/worker/runtime-adapter-workspaces.ts +++ b/src/worker/runtime-adapter-workspaces.ts @@ -26,6 +26,7 @@ import { type RuntimeAdapterCreateAttemptFence, type RuntimeAdapterWorkspaceConflictInput, } from "./provisioning/runtime-adapter.ts"; +import type { RuntimeAdapterWorkspaceRegistration } from "./provisioning/runtime-adapter-release-service.ts"; import { safeProviderError } from "./provisioning/result.ts"; import type { InteractiveProvisionResult } from "./provisioning/types.ts"; import { @@ -35,6 +36,9 @@ import { } from "./runtime-adapter-preflight.ts"; import type { RuntimeAdapterWorkspaceStopResult } from "./session-runtime-adapter-stop.ts"; +export const runtimeAdapterCapabilitiesHeader = "x-crabfleet-runtime-adapter-capabilities"; +export const runtimeAdapterDeleteTombstoneCapability = "delete-tombstone-v1"; + export type RuntimeAdapterWorkspaceLifecycleDependencies = { now(): number; fetch(input: string, init: RequestInit): Promise; @@ -186,26 +190,40 @@ export class RuntimeAdapterWorkspaceLifecycle { async stopForSession( sessionId: string, adapterWorkspaceId: string, + retainedRegistration?: RuntimeAdapterWorkspaceRegistration | null, + retryMissing?: boolean, ): Promise { - const registration = await database(this.env) - .selectFrom("interactive_sessions") - .select(["adapter_control_plane", "adapter_create_pending", "profile"]) - .where("id", "=", sessionId) - .where("adapter", "=", runtimeAdapterName) - .where("adapter_workspace_id", "=", adapterWorkspaceId) - .executeTakeFirst(); + const supersededCleanup = retainedRegistration !== undefined; + const registration = retainedRegistration + ? { + adapter_control_plane: retainedRegistration.controlPlane, + adapter_create_pending: retryMissing ? 1 : 0, + profile: retainedRegistration.profile, + } + : await database(this.env) + .selectFrom("interactive_sessions") + .select(["adapter_control_plane", "adapter_create_pending", "profile"]) + .where("id", "=", sessionId) + .where("adapter", "=", runtimeAdapterName) + .where("adapter_workspace_id", "=", adapterWorkspaceId) + .executeTakeFirst(); const controlPlane = requireRegisteredRuntimeAdapterControlPlane( this.env, registration?.profile ?? "", registration?.adapter_control_plane, ); - if (registration?.adapter_create_pending !== 0) { + if (registration?.adapter_create_pending !== 0 && !supersededCleanup) { return { status: "stopping", message: "runtime adapter stop waiting for create resolution", }; } - return this.stopWorkspace(registration?.profile ?? "", controlPlane, adapterWorkspaceId); + return this.stopWorkspace( + registration?.profile ?? "", + controlPlane, + adapterWorkspaceId, + supersededCleanup && registration?.adapter_create_pending !== 0, + ); } private async reconcileStopping( @@ -514,6 +532,7 @@ export class RuntimeAdapterWorkspaceLifecycle { profile: string, registeredControlPlane: string, adapterWorkspaceId: string, + retryMissing = false, ): Promise { const controlPlane = requireRegisteredRuntimeAdapterControlPlane( this.env, @@ -522,7 +541,12 @@ export class RuntimeAdapterWorkspaceLifecycle { ); const response = await this.dependencies.fetch( runtimeAdapterWorkspaceUrl(controlPlane, adapterWorkspaceId), - { method: "DELETE" }, + { + method: "DELETE", + headers: { + [runtimeAdapterCapabilitiesHeader]: runtimeAdapterDeleteTombstoneCapability, + }, + }, ); const body = response.status === 204 ? null : await this.dependencies.readResponseBody(response); @@ -537,6 +561,18 @@ export class RuntimeAdapterWorkspaceLifecycle { const message = parsed?.message ?? redactedAdapterResponseMessage(body, fallbackMessage, [adapterWorkspaceId]); + if ( + response.status === 404 && + retryMissing && + adapterAdvertisesCapability(response, runtimeAdapterDeleteTombstoneCapability) + ) { + // An ambiguous create may still appear. Accepted DELETE retries must replay + // the adapter's retained stopping or terminal tombstone instead. + return { + status: "stopping", + message: "runtime adapter workspace not yet visible; cleanup retry pending", + }; + } if (response.status === 404 || response.status === 204) { return { status: "stopped", message }; } @@ -549,6 +585,12 @@ export class RuntimeAdapterWorkspaceLifecycle { } } +function adapterAdvertisesCapability(response: Response, capability: string): boolean { + return (response.headers.get(runtimeAdapterCapabilitiesHeader) ?? "") + .split(/[\s,]+/u) + .includes(capability); +} + export function runtimeAdapterProviderConfigured(env: RuntimeEnv): boolean { return Boolean( configuredRuntimeAdapterControlPlane(env, "profile-route") && runtimeAdapterToken(env), diff --git a/src/worker/runtime-application.ts b/src/worker/runtime-application.ts index 5551dc3a..0bad2787 100644 --- a/src/worker/runtime-application.ts +++ b/src/worker/runtime-application.ts @@ -3,6 +3,14 @@ import { mapWithConcurrency } from "./concurrency.ts"; import type { InteractiveSessionRow } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import { readAbandonedInteractiveSessionReservations } from "./openclaw-repository.ts"; +import { + claimRuntimeAdapterWorkspaceCleanup, + claimRuntimeAdapterWorkspaceCleanupBatch, + completeRuntimeAdapterWorkspaceCleanup, + markRuntimeAdapterWorkspaceCleanupDeletionObserved, + persistRuntimeAdapterWorkspaceCleanupEvidence, + stageRuntimeAdapterWorkspaceCleanup, +} from "./provisioning/runtime-adapter-release-repository.ts"; import { safeProviderError } from "./provisioning/result.ts"; import { RuntimeAdapterReleaseService } from "./provisioning/runtime-adapter-release-service.ts"; import { @@ -186,10 +194,31 @@ export class RuntimeApplication { services.release, () => new RuntimeAdapterReleaseService({ + stageCleanup: (input) => stageRuntimeAdapterWorkspaceCleanup(this.env, input), + claimCleanup: (sessionId, adapterWorkspaceId, now) => + claimRuntimeAdapterWorkspaceCleanup(this.env, sessionId, adapterWorkspaceId, now), + claimPendingCleanups: (now) => + claimRuntimeAdapterWorkspaceCleanupBatch(this.env, now, runtimeAdapterReconcileLimit), + persistCleanupEvidence: (cleanup, message, now, reconcileError) => + persistRuntimeAdapterWorkspaceCleanupEvidence( + this.env, + cleanup, + message, + now, + reconcileError, + ), + markCleanupDeletionObserved: (cleanup, now) => + markRuntimeAdapterWorkspaceCleanupDeletionObserved(this.env, cleanup, now), + completeCleanup: (cleanup) => completeRuntimeAdapterWorkspaceCleanup(this.env, cleanup), clearCreatePending: (sessionId, adapterWorkspaceId) => clearRuntimeAdapterCreatePending(this.env, sessionId, adapterWorkspaceId), - stopWorkspace: (sessionId, adapterWorkspaceId) => - this.workspaceLifecycle().stopForSession(sessionId, adapterWorkspaceId), + stopWorkspace: (sessionId, adapterWorkspaceId, registration, retryMissing) => + this.workspaceLifecycle().stopForSession( + sessionId, + adapterWorkspaceId, + registration, + retryMissing, + ), confirmRelease: (sessionId, adapterWorkspaceId, now, message) => confirmRuntimeAdapterRelease(this.env, sessionId, adapterWorkspaceId, now, message), persistStopEvidence: (sessionId, adapterWorkspaceId, message, now, reconcileError) => @@ -263,10 +292,11 @@ export class RuntimeApplication { runtimeAdapterName, ), readSession: (sessionId) => readInteractiveSessionRecord(this.env, sessionId), - stopSuperseded: (sessionId, adapterWorkspaceId, createPending, now) => + stopSuperseded: (sessionId, adapterWorkspaceId, registration, createPending, now) => this.release().stopSuperseded({ sessionId, adapterWorkspaceId, + registration, createPending, now, }), @@ -369,6 +399,7 @@ export class RuntimeApplication { } private async cleanupAbandonedPreparations(now: number): Promise { + await this.release().retryPending(now); const rows = await readAbandonedInteractiveSessionReservations( this.env, now - interactiveSessionPreparationStaleMs, diff --git a/src/worker/sandbox-credential-policy-cleanup-service.ts b/src/worker/sandbox-credential-policy-cleanup-service.ts index a2eeea41..00834637 100644 --- a/src/worker/sandbox-credential-policy-cleanup-service.ts +++ b/src/worker/sandbox-credential-policy-cleanup-service.ts @@ -7,6 +7,7 @@ import { database, type Database, type InteractiveSessionCredentialPolicyTable, + type InteractiveSessionCredentialPolicyRegistrationTable, } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import { serviceUnavailable } from "./http.ts"; @@ -15,12 +16,18 @@ import { safeProviderError } from "./provisioning/result.ts"; import { queueSandboxCredentialPolicyCleanup, sandboxCredentialPolicyCleanupAuthorizedCondition, + sandboxCredentialPolicyPersistedLookupIds, sandboxLookupIds, } from "./sandbox-credential-policy-repository.ts"; +import { restoreSandboxCredentialPolicyRollback } from "./sandbox-credential-policy-rollback.ts"; import { scanCredentialPolicyCleanupPage } from "./sandbox-credential-policy-scanner.ts"; import { isCurrentSandboxLease, sandboxLeaseInfo } from "./sandbox-lease.ts"; import { isSandboxSessionAlreadyGone } from "./sandbox-session-errors.ts"; import { sandboxControlStub } from "./session-control-do.ts"; +import { + sandboxCredentialPolicyRegistrationLookupIds, + sandboxCredentialPolicyRollbackLookupIds, +} from "./session-control-policy.ts"; import { finalizeTerminalInteractiveSession } from "./session-terminal-finalization.ts"; const credentialPolicyCleanupLimit = 8; @@ -30,11 +37,12 @@ export async function sandboxCredentialPolicyExists( env: RuntimeEnv, sandboxId: string, generation: string, + lookupIds: readonly string[] = sandboxLookupIds(env, sandboxId), ): Promise { const stub = sandboxControlStub(env); if (!stub) return false; const responses = await Promise.all( - sandboxLookupIds(env, sandboxId).map((lookupId) => + lookupIds.map((lookupId) => stub.fetch( `https://crabfleet.internal/api/session-control/egress/${encodeURIComponent(lookupId)}`, ), @@ -75,6 +83,119 @@ async function unregisterSandboxCredentialPolicyLookup( } } +async function reconcileStagedCredentialPolicyCleanup( + env: RuntimeEnv, + now: number, + sessionId?: string, +): Promise { + let query = database(env) + .selectFrom("interactive_session_credential_policy_registrations") + .selectAll() + .where("state", "=", "cleanup_pending") + .where((expression) => + expression.or([ + expression("cleanup_claim", "is", null), + expression("cleanup_claim_expires_at", "<", now), + ]), + ) + .orderBy(sql`COALESCE(last_attempt_at, created_at)`, "asc") + .limit(credentialPolicyCleanupLimit); + if (sessionId) query = query.where("session_id", "=", sessionId); + await mapWithConcurrency(await query.execute(), 3, async (registration) => { + await reconcileStagedCredentialPolicyRegistration(env, registration, now); + }); +} + +async function reconcileStagedCredentialPolicyRegistration( + env: RuntimeEnv, + registration: Selectable, + now: number, +): Promise { + const claim = crypto.randomUUID(); + const claimed = await database(env) + .updateTable("interactive_session_credential_policy_registrations") + .set({ + cleanup_claim: claim, + cleanup_claim_expires_at: now + credentialPolicyCleanupClaimMs, + attempt_count: sql`attempt_count + 1`, + last_attempt_at: now, + updated_at: now, + }) + .where("session_id", "=", registration.session_id) + .where("sandbox_id", "=", registration.sandbox_id) + .where("state", "=", "cleanup_pending") + .where("registration_generation", "=", registration.registration_generation) + .where((expression) => + expression.or([ + expression("cleanup_claim", "is", null), + expression("cleanup_claim_expires_at", "<", now), + ]), + ) + .executeTakeFirst(); + if ((claimed.numUpdatedRows ?? 0n) === 0n) return; + try { + const persistedLookupIds = await sandboxCredentialPolicyPersistedLookupIds( + env, + registration.session_id, + registration.sandbox_id, + ); + let rollbackLookupIds: string[] = []; + if (registration.rollback_policies_json !== null) { + try { + rollbackLookupIds = sandboxCredentialPolicyRollbackLookupIds( + registration.rollback_policies_json, + registration.session_id, + ); + } catch { + // Malformed rollback state cannot authorize additional cleanup identities. + } + } + await Promise.all( + sandboxCredentialPolicyRegistrationLookupIds( + registration.lookup_ids_json, + registration.sandbox_id, + sandboxLookupIds(env, registration.sandbox_id), + [...persistedLookupIds, ...rollbackLookupIds], + ).map((lookupId) => + unregisterSandboxCredentialPolicyLookup( + env, + lookupId, + registration.registration_generation, + registration.session_id, + ), + ), + ); + } catch (error) { + await database(env) + .updateTable("interactive_session_credential_policy_registrations") + .set({ + last_error: clean(error instanceof Error ? error.message : String(error), 500), + cleanup_claim: null, + cleanup_claim_expires_at: null, + updated_at: Date.now(), + }) + .where("session_id", "=", registration.session_id) + .where("sandbox_id", "=", registration.sandbox_id) + .where("registration_generation", "=", registration.registration_generation) + .where("cleanup_claim", "=", claim) + .execute(); + return; + } + await database(env) + .deleteFrom("interactive_session_credential_policy_registrations") + .where("session_id", "=", registration.session_id) + .where("sandbox_id", "=", registration.sandbox_id) + .where("registration_generation", "=", registration.registration_generation) + .where("cleanup_claim", "=", claim) + .execute(); + await completeCredentialPolicyCleanupSession(env, registration.session_id, Date.now()); + await completeStandaloneSandboxProvisionCleanupSafely( + env, + registration.session_id, + registration.sandbox_id, + ); +} + async function normalizeCredentialPolicyCleanupGroups( env: RuntimeEnv, now: number, @@ -197,11 +318,28 @@ export async function reconcileSandboxCredentialPolicyCleanupBatch( await expireStandaloneSandboxProvisions(env, now, sessionId).catch((error) => { console.error("standalone Sandbox expiry failed", error); }); - await scanCredentialPolicyCleanupPage(env, now, sandboxCredentialPolicyExists, sessionId).catch( - (error) => { - console.error("credential policy cleanup scan failed", error); + await scanCredentialPolicyCleanupPage( + env, + now, + sandboxCredentialPolicyExists, + sessionId, + async ({ registration, registrationExpiresAt, rollbackJson, sessionId: rollbackSessionId }) => { + const stub = sandboxControlStub(env); + if (!stub) throw serviceUnavailable("sandbox credential policy rollback is unavailable"); + await restoreSandboxCredentialPolicyRollback( + stub, + registration, + registrationExpiresAt, + rollbackJson, + rollbackSessionId, + ); }, - ); + ).catch((error) => { + console.error("credential policy cleanup scan failed", error); + }); + await reconcileStagedCredentialPolicyCleanup(env, now, sessionId).catch((error) => { + console.error("staged credential policy cleanup failed", error); + }); await normalizeCredentialPolicyCleanupGroups(env, now, sessionId).catch((error) => { console.error("credential policy cleanup group normalization failed", error); }); @@ -225,6 +363,16 @@ export async function reconcileSandboxCredentialPolicyCleanupBatch( AND registration.registration_claim_expires_at > ${now} ) `) + .where(sql` + NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS registration + WHERE registration.session_id = interactive_session_credential_policies.session_id + AND registration.sandbox_id = interactive_session_credential_policies.sandbox_id + AND registration.state = 'registering' + AND registration.registration_claim_expires_at > ${now} + ) + `) .orderBy(sql`COALESCE(last_attempt_at, created_at)`, "asc") .orderBy("session_id", "asc") .orderBy("sandbox_id", "asc") @@ -246,6 +394,11 @@ export async function reconcileSandboxCredentialPolicyCleanupBatch( FROM interactive_session_credential_policies AS policy WHERE policy.session_id = interactive_sessions.id ) + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS registration + WHERE registration.session_id = interactive_sessions.id + ) `) .orderBy("stopped_at", "asc") .orderBy("id", "asc") @@ -298,6 +451,14 @@ async function reconcileCredentialPolicyCleanup( AND registration.registration_claim IS NOT NULL AND registration.registration_claim_expires_at > ${now} ) + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS registration + WHERE registration.session_id = interactive_session_credential_policies.session_id + AND registration.sandbox_id = interactive_session_credential_policies.sandbox_id + AND registration.state = 'registering' + AND registration.registration_claim_expires_at > ${now} + ) `.execute(database(env)); if ((claimed.numAffectedRows ?? 0n) === 0n) return; try { @@ -411,6 +572,12 @@ async function completeStandaloneSandboxProvisionCleanup( WHERE policy.session_id = ${provisionId} AND policy.sandbox_id = ${sandboxId} ) + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS registration + WHERE registration.session_id = ${provisionId} + AND registration.sandbox_id = ${sandboxId} + ) `) .execute(); } @@ -421,12 +588,19 @@ async function completeCredentialPolicyCleanupSession( now: number, ): Promise { const db = database(env); - const remaining = await db - .selectFrom("interactive_session_credential_policies") - .select(({ fn }) => fn.countAll().as("count")) - .where("session_id", "=", sessionId) - .executeTakeFirst(); - if (Number(remaining?.count ?? 0) > 0) return; + const remaining = await sql<{ count: number }>` + SELECT + ( + SELECT count(*) + FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + ) + ( + SELECT count(*) + FROM interactive_session_credential_policy_registrations + WHERE session_id = ${sessionId} + ) AS count + `.execute(db); + if (Number(remaining.rows[0]?.count ?? 0) > 0) return; const session = await db .selectFrom("interactive_sessions") .select([ @@ -479,6 +653,11 @@ async function completeCredentialPolicyCleanupSession( FROM interactive_session_credential_policies WHERE session_id = ${sessionId} ) + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations + WHERE session_id = ${sessionId} + ) `) .executeTakeFirst(); if ((updated.numUpdatedRows ?? 0n) === 0n) return; diff --git a/src/worker/sandbox-credential-policy-cleanup.ts b/src/worker/sandbox-credential-policy-cleanup.ts index 31823781..62e4368b 100644 --- a/src/worker/sandbox-credential-policy-cleanup.ts +++ b/src/worker/sandbox-credential-policy-cleanup.ts @@ -187,6 +187,21 @@ export async function stageTerminalCredentialPolicyCleanup( ]), ) .where(sandboxManagedStoredOwnershipCondition(ownership.fence)); + const registrationTransitions = generations.map(({ sandboxId }) => + db + .updateTable("interactive_session_credential_policy_registrations") + .set({ + state: "cleanup_pending", + registration_claim: null, + registration_claim_expires_at: null, + updated_at: stageRevision, + }) + .where("session_id", "=", session.id) + .where("sandbox_id", "=", sandboxId) + .where( + sandboxCredentialPolicyCleanupAuthorizedCondition(session.id, sandboxId, stageRevision), + ), + ); const policyTransitions = generations.flatMap(({ generation, sandboxId }) => [ ...sandboxCredentialPolicyRefQueries( env, @@ -209,7 +224,7 @@ export async function stageTerminalCredentialPolicyCleanup( sandboxCredentialPolicyCleanupAuthorizedCondition(session.id, sandboxId, stageRevision), ), ]); - await executeBatch(env, [sessionTransition, ...policyTransitions]); + await executeBatch(env, [sessionTransition, ...registrationTransitions, ...policyTransitions]); const staged = await db .selectFrom("interactive_sessions") .select([ diff --git a/src/worker/sandbox-credential-policy-registration-service.ts b/src/worker/sandbox-credential-policy-registration-service.ts index 44b2c5ed..ff5c6060 100644 --- a/src/worker/sandbox-credential-policy-registration-service.ts +++ b/src/worker/sandbox-credential-policy-registration-service.ts @@ -2,15 +2,29 @@ import { fetchGithubRepoNodeId } from "./github.ts"; import { sealSecret } from "./crypto.ts"; import type { RuntimeEnv } from "./env.ts"; import { + activeSandboxCredentialPolicyGeneration, abandonSandboxCredentialPolicyRegistration, beginSandboxCredentialPolicyRegistration, + claimObsoleteSandboxCredentialPolicyReferences, + deferSandboxCredentialPolicyRollback, existingSandboxCredentialPolicyGeneration, finishSandboxCredentialPolicyRegistration, + incompleteSandboxCredentialPolicyGeneration, + markSandboxCredentialPolicyRegistrationWriteStarted, recordSandboxCredentialPolicyRefs, + recordSandboxCredentialPolicyRollback, + repairSandboxCredentialPolicyReferences, + retireObsoleteSandboxCredentialPolicyReference, renewSandboxCredentialPolicyRegistration, + sandboxCredentialPolicyLookupIdsForGeneration, + stageSandboxCredentialPolicyReferenceRepair, standaloneSandboxPolicyExpiresAt, type SandboxCredentialPolicyOwnershipFence, } from "./sandbox-credential-policy-repository.ts"; +import { + captureSandboxCredentialPolicyRollback, + restoreSandboxCredentialPolicyRollback, +} from "./sandbox-credential-policy-rollback.ts"; import { sandboxCredentialPolicyExists } from "./sandbox-credential-policy-cleanup-service.ts"; import { sandboxLeaseInfo, @@ -21,10 +35,228 @@ import type { SandboxRuntimeSession } from "./sandbox-runtime.ts"; import { sandboxControlStub } from "./session-control-do.ts"; import type { SandboxCredentialPolicy, + SandboxCredentialPolicyRegistration, StoredSandboxCredentialPolicy, } from "./session-control-policy.ts"; import type { InteractiveSession } from "./session-model.ts"; +type RestoreSandboxCredentialPolicyRollback = typeof restoreSandboxCredentialPolicyRollback; + +async function repairIncompleteSandboxCredentialPolicyLookupSet( + env: RuntimeEnv, + stub: Pick, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + ownershipFence: SandboxCredentialPolicyOwnershipFence, +): Promise { + const repairGeneration = await incompleteSandboxCredentialPolicyGeneration( + env, + sessionId, + sandboxId, + ); + if (!repairGeneration) return null; + const historicalLookupIds = await sandboxCredentialPolicyLookupIdsForGeneration( + env, + sessionId, + sandboxId, + repairGeneration, + ); + const repairLookupIds = [...new Set([...historicalLookupIds, ...registration.lookupIds])]; + const records = new Map( + await Promise.all( + repairLookupIds.map(async (lookupId) => { + const response = await stub.fetch( + `https://crabfleet.internal/api/session-control/egress/${encodeURIComponent(lookupId)}`, + ); + if (response.status === 404) return [lookupId, null] as const; + if (!response.ok) throw new Error("sandbox credential policy repair snapshot failed"); + const generation = response.headers.get("x-crabfleet-policy-generation"); + const policy = (await response.json()) as SandboxCredentialPolicy; + if ( + generation !== repairGeneration || + policy.sessionId !== sessionId || + policy.sandboxId !== lookupId + ) { + throw new Error("sandbox credential policy repair snapshot is inconsistent"); + } + return [lookupId, { generation, policy }] as const; + }), + ), + ); + const surviving = [...records.values()].filter((record) => record !== null); + if (surviving.length === 0) return null; + const source = surviving[0]!.policy; + if ( + surviving.some( + (record) => + JSON.stringify({ ...record.policy, sandboxId: "" }) !== + JSON.stringify({ ...source, sandboxId: "" }), + ) + ) { + throw new Error("sandbox credential policy repair snapshot is inconsistent"); + } + let registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + sessionId, + sandboxId, + registration, + ownershipFence, + ); + if ( + !registrationExpiresAt || + !(await stageSandboxCredentialPolicyReferenceRepair( + env, + sessionId, + sandboxId, + registration, + repairGeneration, + ownershipFence, + )) + ) { + throw new Error("sandbox credential policy registration claim was revoked"); + } + const missingLookupIds = registration.lookupIds.filter((lookupId) => !records.get(lookupId)); + if (missingLookupIds.length > 0) { + registrationExpiresAt = await markSandboxCredentialPolicyRegistrationWriteStarted( + env, + sessionId, + sandboxId, + registration, + ownershipFence, + ); + if (!registrationExpiresAt) { + throw new Error("sandbox credential policy registration claim was revoked"); + } + } + const repairExpiresAt = registrationExpiresAt - 1; + for (const lookupId of missingLookupIds) { + const response = await stub.fetch("https://crabfleet.internal/api/session-control/register", { + method: "POST", + body: JSON.stringify({ + generation: repairGeneration, + registrationClaim: registration.claim, + registrationExpiresAt: repairExpiresAt, + policy: { ...source, sandboxId: lookupId }, + } satisfies StoredSandboxCredentialPolicy), + headers: { "content-type": "application/json" }, + }); + if (!response.ok) throw new Error("sandbox credential policy lookup repair failed"); + } + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + sessionId, + sandboxId, + registration, + ownershipFence, + ); + if ( + !registrationExpiresAt || + !(await repairSandboxCredentialPolicyReferences( + env, + sessionId, + sandboxId, + registration, + repairGeneration, + ownershipFence, + registrationExpiresAt, + )) + ) { + throw new Error("sandbox credential policy lookup references were not repaired"); + } + const currentLookupIds = new Set(registration.lookupIds); + const obsoleteLookupIds = historicalLookupIds.filter( + (lookupId) => !currentLookupIds.has(lookupId), + ); + if (obsoleteLookupIds.length > 0) { + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + sessionId, + sandboxId, + registration, + ownershipFence, + ); + if (!registrationExpiresAt) { + throw new Error("sandbox credential policy registration claim was revoked"); + } + const claimed = await claimObsoleteSandboxCredentialPolicyReferences( + env, + sessionId, + sandboxId, + registration, + repairGeneration, + obsoleteLookupIds, + ownershipFence, + registrationExpiresAt, + ); + if ( + claimed.length !== obsoleteLookupIds.length || + obsoleteLookupIds.some((lookupId) => !claimed.includes(lookupId)) + ) { + throw new Error("obsolete sandbox credential policy references were not claimed"); + } + for (const lookupId of obsoleteLookupIds) { + const response = await stub.fetch( + `https://crabfleet.internal/api/session-control/sandbox/${encodeURIComponent(lookupId)}`, + { + method: "DELETE", + body: JSON.stringify({ + generation: repairGeneration, + sessionId, + tombstonedAt: Date.now(), + }), + headers: { "content-type": "application/json" }, + }, + ); + if (!response.ok) { + throw new Error("obsolete sandbox credential policy retirement failed"); + } + if ( + !(await retireObsoleteSandboxCredentialPolicyReference( + env, + sessionId, + sandboxId, + registration, + repairGeneration, + lookupId, + ownershipFence, + registrationExpiresAt, + )) + ) { + throw new Error("obsolete sandbox credential policy reference was not retired"); + } + } + } + if ( + (await activeSandboxCredentialPolicyGeneration(env, sessionId, sandboxId)) !== repairGeneration + ) { + throw new Error("sandbox credential policy namespace repair did not converge"); + } + return repairGeneration; +} + +export async function restoreSandboxCredentialPolicyRollbackIfOwned( + env: RuntimeEnv, + stub: Pick, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + rollbackJson: string, + ownershipFence: SandboxCredentialPolicyOwnershipFence, + restoreRollback: RestoreSandboxCredentialPolicyRollback = restoreSandboxCredentialPolicyRollback, +): Promise { + const registrationExpiresAt = await markSandboxCredentialPolicyRegistrationWriteStarted( + env, + sessionId, + sandboxId, + registration, + ownershipFence, + ); + if (!registrationExpiresAt) return false; + await restoreRollback(stub, registration, registrationExpiresAt, rollbackJson, sessionId); + return true; +} + export async function registerSandboxCredentialPolicy( env: RuntimeEnv, session: SandboxRuntimeSession, @@ -46,7 +278,38 @@ export async function registerSandboxCredentialPolicy( sandboxId, ownershipFence, ); + let rollbackJson: string | null = null; + let registrationWriteStarted = false; try { + const activeGeneration = + (await activeSandboxCredentialPolicyGeneration(env, session.id, sandboxId)) ?? + (await repairIncompleteSandboxCredentialPolicyLookupSet( + env, + stub, + session.id, + sandboxId, + registration, + ownershipFence, + )); + const rollback = await captureSandboxCredentialPolicyRollback( + stub, + registration.lookupIds, + activeGeneration, + session.id, + ); + if ( + !(await recordSandboxCredentialPolicyRollback( + env, + session.id, + sandboxId, + registration, + rollback, + ownershipFence, + )) + ) { + throw new Error("sandbox credential policy rollback snapshot was not recorded"); + } + rollbackJson = JSON.stringify(rollback); const githubToken = "githubToken" in session ? session.githubToken : undefined; const githubTokenCiphertext = githubToken ? await sealSecret(env, githubToken) : null; if (githubToken && !githubTokenCiphertext) { @@ -77,7 +340,7 @@ export async function registerSandboxCredentialPolicy( ...(env.OPENAI_ORG_ID ? { openAIOrgId: env.OPENAI_ORG_ID } : {}), }; for (const lookupId of registration.lookupIds) { - const registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + const registrationExpiresAt = await markSandboxCredentialPolicyRegistrationWriteStarted( env, session.id, sandboxId, @@ -87,6 +350,7 @@ export async function registerSandboxCredentialPolicy( if (!registrationExpiresAt) { throw new Error("sandbox credential policy registration claim was revoked"); } + registrationWriteStarted = true; const response = await stub.fetch("https://crabfleet.internal/api/session-control/register", { method: "POST", body: JSON.stringify({ @@ -113,12 +377,42 @@ export async function registerSandboxCredentialPolicy( throw new Error("sandbox credential policy cleanup became pending during registration"); } } catch (error) { + const message = clean(error instanceof Error ? error.message : String(error), 500); + if (registrationWriteStarted && rollbackJson) { + try { + const restored = await restoreSandboxCredentialPolicyRollbackIfOwned( + env, + stub, + session.id, + sandboxId, + registration, + rollbackJson, + ownershipFence, + ); + if (!restored) { + throw new Error("sandbox credential policy registration claim was revoked"); + } + } catch (rollbackError) { + const rollbackMessage = clean( + rollbackError instanceof Error ? rollbackError.message : String(rollbackError), + 500, + ); + await deferSandboxCredentialPolicyRollback( + env, + session.id, + sandboxId, + registration, + `${message}; ${rollbackMessage}`, + ).catch(() => undefined); + throw new Error("sandbox credential policy rollback restore is pending", { cause: error }); + } + } await abandonSandboxCredentialPolicyRegistration( env, session.id, sandboxId, registration, - clean(error instanceof Error ? error.message : String(error), 500), + message, ).catch(() => undefined); throw error; } diff --git a/src/worker/sandbox-credential-policy-repository.ts b/src/worker/sandbox-credential-policy-repository.ts index 6378ed76..aad47e16 100644 --- a/src/worker/sandbox-credential-policy-repository.ts +++ b/src/worker/sandbox-credential-policy-repository.ts @@ -6,6 +6,7 @@ import { type SandboxCredentialPolicy, type SandboxCredentialPolicyRegistration, } from "./session-control-policy.ts"; +import type { SandboxCredentialPolicyRollbackRecord } from "./sandbox-credential-policy-rollback.ts"; import { database, executeBatch, type CompilableQuery } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import { @@ -64,6 +65,7 @@ export function activeSandboxCredentialPolicyCondition( state != 'active' OR registration_generation != ${generation} OR registration_claim IS NOT NULL + OR lookup_id NOT IN (${sql.join(lookupIds)}) OR NOT (${updatedAtCondition}) ) ) @@ -81,11 +83,12 @@ export async function activeSandboxCredentialPolicyGeneration( .where("session_id", "=", sessionId) .where("sandbox_id", "=", sandboxId) .execute(); - const expected = sandboxLookupIds(env, sandboxId); + const expected = new Set(sandboxLookupIds(env, sandboxId)); const generation = rows[0]?.registration_generation; if ( !isCurrentCredentialPolicyGeneration(generation) || - !expected.every((lookupId) => + rows.length !== expected.size || + ![...expected].every((lookupId) => rows.some( (row) => row.lookup_id === lookupId && @@ -96,6 +99,7 @@ export async function activeSandboxCredentialPolicyGeneration( ) || rows.some( (row) => + !expected.has(row.lookup_id) || row.state !== "active" || row.registration_generation !== generation || row.registration_claim !== null, @@ -106,6 +110,70 @@ export async function activeSandboxCredentialPolicyGeneration( return generation; } +export async function incompleteSandboxCredentialPolicyGeneration( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, +): Promise { + const rows = await database(env) + .selectFrom("interactive_session_credential_policies") + .select(["lookup_id", "state", "registration_generation", "registration_claim"]) + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .execute(); + const expected = new Set(sandboxLookupIds(env, sandboxId)); + const generation = rows[0]?.registration_generation; + if ( + rows.length === 0 || + !isCurrentCredentialPolicyGeneration(generation) || + !rows.some((row) => row.lookup_id === sandboxId) || + rows.some( + (row) => + row.state !== "active" || + row.registration_generation !== generation || + row.registration_claim !== null, + ) || + (rows.length === expected.size && rows.every((row) => expected.has(row.lookup_id))) + ) { + return null; + } + return generation; +} + +export async function sandboxCredentialPolicyLookupIdsForGeneration( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + generation: string, +): Promise { + const rows = await database(env) + .selectFrom("interactive_session_credential_policies") + .select("lookup_id") + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .where("state", "=", "active") + .where("registration_generation", "=", generation) + .where("registration_claim", "is", null) + .orderBy("lookup_id") + .execute(); + return rows.map((row) => row.lookup_id); +} + +export async function sandboxCredentialPolicyPersistedLookupIds( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, +): Promise { + const rows = await database(env) + .selectFrom("interactive_session_credential_policies") + .select("lookup_id") + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .orderBy("lookup_id") + .execute(); + return rows.map((row) => row.lookup_id); +} + export async function sandboxCredentialPolicyHasDurableOwner( env: RuntimeEnv, lookupId: string, @@ -187,6 +255,16 @@ export function sandboxCredentialPolicyRefQueries( now: number, authorizationCondition: RawBuilder, ): CompilableQuery[] { + const stagedWriteAllowed = + state === "cleanup_pending" + ? sql`1 = 1` + : sql`NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = ${sessionId} + AND staged.sandbox_id = ${sandboxId} + AND staged.registration_generation != ${generation} + )`; return sandboxLookupIds(env, sandboxId).map( (lookupId) => sql` INSERT INTO interactive_session_credential_policies ( @@ -220,6 +298,7 @@ export function sandboxCredentialPolicyRefQueries( ${now}, ${now} WHERE ${authorizationCondition} + AND ${stagedWriteAllowed} ON CONFLICT(session_id, sandbox_id, lookup_id) DO UPDATE SET state = CASE WHEN interactive_session_credential_policies.state = 'cleanup_pending' @@ -431,6 +510,22 @@ export function sandboxCredentialPolicyOwnerCondition( )`; } +function noLivePolicyTableRegistrationCondition( + sessionId: string, + sandboxId: string, + now: number, +): RawBuilder { + return sql`NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'registering' + AND registration_claim IS NOT NULL + AND registration_claim_expires_at > ${now} + )`; +} + export function sandboxCredentialPolicyRegistrationQueries( sessionId: string, sandboxId: string, @@ -439,16 +534,16 @@ export function sandboxCredentialPolicyRegistrationQueries( now: number, ownershipFence: SandboxCredentialPolicyOwnershipFence, ): CompilableQuery[] { - return registration.lookupIds.map( - (lookupId) => sql` - INSERT INTO interactive_session_credential_policies ( + return [ + sql` + INSERT INTO interactive_session_credential_policy_registrations ( session_id, sandbox_id, - lookup_id, state, registration_generation, registration_claim, registration_claim_expires_at, + lookup_ids_json, attempt_count, last_attempt_at, last_error, @@ -460,11 +555,11 @@ export function sandboxCredentialPolicyRegistrationQueries( SELECT ${sessionId}, ${sandboxId}, - ${lookupId}, 'registering', ${registration.generation}, ${registration.claim}, ${registrationExpiresAt}, + ${JSON.stringify(registration.lookupIds)}, 0, NULL, NULL, @@ -473,23 +568,43 @@ export function sandboxCredentialPolicyRegistrationQueries( ${now}, ${now} WHERE ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} - ON CONFLICT(session_id, sandbox_id, lookup_id) DO UPDATE SET + AND ${noLivePolicyTableRegistrationCondition(sessionId, sandboxId, now)} + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'cleanup_pending' + ) + ON CONFLICT(session_id, sandbox_id) DO UPDATE SET state = 'registering', registration_generation = excluded.registration_generation, registration_claim = excluded.registration_claim, registration_claim_expires_at = excluded.registration_claim_expires_at, + lookup_ids_json = excluded.lookup_ids_json, + repair_generation = NULL, + rollback_policies_json = NULL, last_error = NULL, cleanup_claim = NULL, cleanup_claim_expires_at = NULL, updated_at = excluded.updated_at - WHERE interactive_session_credential_policies.state != 'cleanup_pending' + WHERE interactive_session_credential_policy_registrations.state != 'cleanup_pending' + AND interactive_session_credential_policy_registrations.registration_write_started = 0 AND ( - interactive_session_credential_policies.registration_claim IS NULL - OR interactive_session_credential_policies.registration_claim_expires_at <= ${now} + interactive_session_credential_policy_registrations.registration_claim IS NULL + OR interactive_session_credential_policy_registrations.registration_claim_expires_at <= ${now} ) AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + AND ${noLivePolicyTableRegistrationCondition(sessionId, sandboxId, now)} + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'cleanup_pending' + ) `, - ); + ]; } export async function beginSandboxCredentialPolicyRegistration( @@ -500,18 +615,8 @@ export async function beginSandboxCredentialPolicyRegistration( ): Promise { const db = database(env); const lookupIds = sandboxLookupIds(env, sandboxId); - const existing = await db - .selectFrom("interactive_session_credential_policies") - .select("registration_generation") - .distinct() - .where("session_id", "=", sessionId) - .where("sandbox_id", "=", sandboxId) - .execute(); - const existingGeneration = currentSandboxCredentialPolicyGeneration( - existing.map((row) => row.registration_generation), - ); const registration = { - generation: existingGeneration ?? newSandboxCredentialPolicyGeneration(), + generation: newSandboxCredentialPolicyGeneration(), claim: `registration:${crypto.randomUUID()}`, lookupIds, }; @@ -529,27 +634,23 @@ export async function beginSandboxCredentialPolicyRegistration( ), ); const claimed = await db - .selectFrom("interactive_session_credential_policies") + .selectFrom("interactive_session_credential_policy_registrations") .select([ - "lookup_id", "state", "registration_generation", "registration_claim", "registration_claim_expires_at", + "lookup_ids_json", ]) .where("session_id", "=", sessionId) .where("sandbox_id", "=", sandboxId) - .where("lookup_id", "in", lookupIds) - .execute(); + .executeTakeFirst(); if ( - claimed.length !== lookupIds.length || - claimed.some( - (row) => - row.state !== "registering" || - row.registration_generation !== registration.generation || - row.registration_claim !== registration.claim || - row.registration_claim_expires_at !== registrationExpiresAt, - ) + claimed?.state !== "registering" || + claimed.registration_generation !== registration.generation || + claimed.registration_claim !== registration.claim || + claimed.registration_claim_expires_at !== registrationExpiresAt || + claimed.lookup_ids_json !== JSON.stringify(registration.lookupIds) ) { await abandonSandboxCredentialPolicyRegistration( env, @@ -573,49 +674,444 @@ export async function renewSandboxCredentialPolicyRegistration( const now = Date.now(); const registrationExpiresAt = now + credentialPolicyRegistrationClaimMs; const renewed = await database(env) - .updateTable("interactive_session_credential_policies") + .updateTable("interactive_session_credential_policy_registrations") .set({ registration_claim_expires_at: registrationExpiresAt, updated_at: now, }) .where("session_id", "=", sessionId) .where("sandbox_id", "=", sandboxId) - .where("lookup_id", "in", registration.lookupIds) .where("state", "=", "registering") .where("registration_generation", "=", registration.generation) .where("registration_claim", "=", registration.claim) + .where("registration_claim_expires_at", ">", now) .where(sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)) + .where(noLivePolicyTableRegistrationCondition(sessionId, sandboxId, now)) .executeTakeFirst(); - return Number(renewed.numUpdatedRows ?? 0n) === registration.lookupIds.length + return Number(renewed.numUpdatedRows ?? 0n) === 1 ? registrationExpiresAt : null; +} + +export async function markSandboxCredentialPolicyRegistrationWriteStarted( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + ownershipFence: SandboxCredentialPolicyOwnershipFence, +): Promise { + const now = Date.now(); + const registrationExpiresAt = now + credentialPolicyRegistrationClaimMs; + const marked = await sql<{ registration_claim_expires_at: number }>` + UPDATE interactive_session_credential_policy_registrations + SET + registration_claim_expires_at = ${registrationExpiresAt}, + registration_write_started = 1, + updated_at = ${now} + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'registering' + AND registration_generation = ${registration.generation} + AND registration_claim = ${registration.claim} + AND registration_claim_expires_at > ${now} + AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + AND ${noLivePolicyTableRegistrationCondition(sessionId, sandboxId, now)} + RETURNING registration_claim_expires_at + `.execute(database(env)); + return marked.rows[0]?.registration_claim_expires_at === registrationExpiresAt ? registrationExpiresAt : null; } -export async function finishSandboxCredentialPolicyRegistration( +export async function claimSandboxCredentialPolicyRegistrationRecovery( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + expiredRegistration: SandboxCredentialPolicyRegistration, + expiredRegistrationExpiresAt: number, + ownershipFence: SandboxCredentialPolicyOwnershipFence, +): Promise<{ + registration: SandboxCredentialPolicyRegistration; + registrationExpiresAt: number; +} | null> { + const now = Date.now(); + const registration = { + ...expiredRegistration, + claim: `registration:${crypto.randomUUID()}`, + }; + const registrationExpiresAt = now + credentialPolicyRegistrationClaimMs; + const claimed = await database(env) + .updateTable("interactive_session_credential_policy_registrations") + .set({ + registration_claim: registration.claim, + registration_claim_expires_at: registrationExpiresAt, + updated_at: now, + }) + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .where("state", "=", "registering") + .where("registration_generation", "=", expiredRegistration.generation) + .where("registration_claim", "=", expiredRegistration.claim) + .where("registration_claim_expires_at", "=", expiredRegistrationExpiresAt) + .where("registration_claim_expires_at", "<=", now) + .where(sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)) + .where(noLivePolicyTableRegistrationCondition(sessionId, sandboxId, now)) + .executeTakeFirst(); + return Number(claimed.numUpdatedRows ?? 0n) === 1 + ? { registration, registrationExpiresAt } + : null; +} + +export async function recordSandboxCredentialPolicyRollback( env: RuntimeEnv, sessionId: string, sandboxId: string, registration: SandboxCredentialPolicyRegistration, + rollback: readonly SandboxCredentialPolicyRollbackRecord[], ownershipFence: SandboxCredentialPolicyOwnershipFence, ): Promise { const now = Date.now(); - const db = database(env); - await db - .updateTable("interactive_session_credential_policies") + const recorded = await database(env) + .updateTable("interactive_session_credential_policy_registrations") .set({ - state: "active", - registration_claim: null, - registration_claim_expires_at: null, + rollback_policies_json: JSON.stringify(rollback), updated_at: now, }) .where("session_id", "=", sessionId) .where("sandbox_id", "=", sandboxId) - .where("lookup_id", "in", registration.lookupIds) .where("state", "=", "registering") .where("registration_generation", "=", registration.generation) .where("registration_claim", "=", registration.claim) + .where("registration_claim_expires_at", ">", now) .where(sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)) + .executeTakeFirst(); + return Number(recorded.numUpdatedRows ?? 0n) === 1; +} + +export async function stageSandboxCredentialPolicyReferenceRepair( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + repairGeneration: string, + ownershipFence: SandboxCredentialPolicyOwnershipFence, +): Promise { + const now = Date.now(); + const staged = await sql<{ repair_generation: string }>` + UPDATE interactive_session_credential_policy_registrations + SET repair_generation = ${repairGeneration}, updated_at = ${now} + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'registering' + AND registration_generation = ${registration.generation} + AND registration_claim = ${registration.claim} + AND registration_claim_expires_at > ${now} + AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + RETURNING repair_generation + `.execute(database(env)); + return staged.rows.length === 1 && staged.rows[0]?.repair_generation === repairGeneration; +} + +export async function repairSandboxCredentialPolicyReferences( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + repairGeneration: string, + ownershipFence: SandboxCredentialPolicyOwnershipFence, + registrationExpiresAt: number, +): Promise { + const now = Date.now(); + const repairAuthorized = sql` + EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = ${sessionId} + AND staged.sandbox_id = ${sandboxId} + AND staged.state = 'registering' + AND staged.registration_generation = ${registration.generation} + AND staged.registration_claim = ${registration.claim} + AND staged.registration_claim_expires_at = ${registrationExpiresAt} + AND staged.registration_claim_expires_at > ${now} + AND staged.repair_generation = ${repairGeneration} + ) + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policies AS surviving + WHERE surviving.session_id = ${sessionId} + AND surviving.sandbox_id = ${sandboxId} + AND surviving.state = 'active' + AND surviving.registration_generation = ${repairGeneration} + AND surviving.registration_claim IS NULL + ) + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policies AS conflicting + WHERE conflicting.session_id = ${sessionId} + AND conflicting.sandbox_id = ${sandboxId} + AND ( + conflicting.registration_generation != ${repairGeneration} + OR NOT ( + ( + conflicting.state = 'active' + AND conflicting.registration_claim IS NULL + ) + OR ( + conflicting.state = 'registering' + AND conflicting.registration_claim = ${registration.claim} + AND conflicting.registration_claim_expires_at = ${registrationExpiresAt} + ) + ) + ) + ) + AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + `; + const inserts = registration.lookupIds.map( + (lookupId) => sql` + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + attempt_count, + last_attempt_at, + last_error, + cleanup_claim, + cleanup_claim_expires_at, + created_at, + updated_at + ) + SELECT + ${sessionId}, + ${sandboxId}, + ${lookupId}, + 'registering', + ${repairGeneration}, + ${registration.claim}, + ${registrationExpiresAt}, + 0, + NULL, + NULL, + NULL, + NULL, + ${now}, + ${now} + WHERE ${repairAuthorized} + ON CONFLICT(session_id, sandbox_id, lookup_id) DO NOTHING + `, + ); + const promotions = registration.lookupIds.map( + (lookupId) => sql` + UPDATE interactive_session_credential_policies + SET + state = 'active', + registration_claim = NULL, + registration_claim_expires_at = NULL, + updated_at = ${now} + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND lookup_id = ${lookupId} + AND state = 'registering' + AND registration_generation = ${repairGeneration} + AND registration_claim = ${registration.claim} + AND registration_claim_expires_at = ${registrationExpiresAt} + AND ${repairAuthorized} + `, + ); + await executeBatch(env, [...inserts, ...promotions]); + const refs = await database(env) + .selectFrom("interactive_session_credential_policies") + .select(["lookup_id", "state", "registration_generation", "registration_claim"]) + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) .execute(); + return ( + registration.lookupIds.every((lookupId) => + refs.some( + (ref) => + ref.lookup_id === lookupId && + ref.state === "active" && + ref.registration_generation === repairGeneration && + ref.registration_claim === null, + ), + ) && + refs.every( + (ref) => + ref.state === "active" && + ref.registration_generation === repairGeneration && + ref.registration_claim === null, + ) + ); +} + +export async function claimObsoleteSandboxCredentialPolicyReferences( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + repairGeneration: string, + obsoleteLookupIds: readonly string[], + ownershipFence: SandboxCredentialPolicyOwnershipFence, + registrationExpiresAt: number, +): Promise { + if (obsoleteLookupIds.length === 0) return []; + const now = Date.now(); + const claimed = await sql<{ lookup_id: string }>` + UPDATE interactive_session_credential_policies + SET + state = 'cleanup_pending', + cleanup_claim = ${registration.claim}, + cleanup_claim_expires_at = ${registrationExpiresAt}, + last_error = 'obsolete sandbox credential policy namespace', + updated_at = ${now} + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND lookup_id IN (${sql.join(obsoleteLookupIds)}) + AND state = 'active' + AND registration_generation = ${repairGeneration} + AND registration_claim IS NULL + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = ${sessionId} + AND staged.sandbox_id = ${sandboxId} + AND staged.state = 'registering' + AND staged.registration_generation = ${registration.generation} + AND staged.registration_claim = ${registration.claim} + AND staged.registration_claim_expires_at = ${registrationExpiresAt} + AND staged.registration_claim_expires_at > ${now} + AND staged.repair_generation = ${repairGeneration} + ) + AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + RETURNING lookup_id + `.execute(database(env)); + return claimed.rows.map((row) => row.lookup_id).sort(); +} + +export async function retireObsoleteSandboxCredentialPolicyReference( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + repairGeneration: string, + lookupId: string, + ownershipFence: SandboxCredentialPolicyOwnershipFence, + registrationExpiresAt: number, +): Promise { + const now = Date.now(); + const retired = await sql<{ lookup_id: string }>` + DELETE FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND lookup_id = ${lookupId} + AND state = 'cleanup_pending' + AND registration_generation = ${repairGeneration} + AND cleanup_claim = ${registration.claim} + AND cleanup_claim_expires_at = ${registrationExpiresAt} + AND EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations AS staged + WHERE staged.session_id = ${sessionId} + AND staged.sandbox_id = ${sandboxId} + AND staged.state = 'registering' + AND staged.registration_generation = ${registration.generation} + AND staged.registration_claim = ${registration.claim} + AND staged.registration_claim_expires_at = ${registrationExpiresAt} + AND staged.registration_claim_expires_at > ${now} + AND staged.repair_generation = ${repairGeneration} + ) + AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + RETURNING lookup_id + `.execute(database(env)); + return retired.rows[0]?.lookup_id === lookupId; +} + +export async function deferSandboxCredentialPolicyRollback( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + reason: string, +): Promise { + const now = Date.now(); + await database(env) + .updateTable("interactive_session_credential_policy_registrations") + .set({ + registration_claim_expires_at: now, + last_error: reason, + updated_at: now, + }) + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .where("state", "=", "registering") + .where("registration_generation", "=", registration.generation) + .where("registration_claim", "=", registration.claim) + .where("rollback_policies_json", "is not", null) + .execute(); +} + +export async function finishSandboxCredentialPolicyRegistration( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + ownershipFence: SandboxCredentialPolicyOwnershipFence, +): Promise { + const registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + sessionId, + sandboxId, + registration, + ownershipFence, + ); + if (!registrationExpiresAt) return false; + const now = Date.now(); + let batchError: unknown; + try { + await executeBatch( + env, + sandboxCredentialPolicyPromotionQueries( + env, + sessionId, + sandboxId, + registration, + registrationExpiresAt, + ownershipFence, + now, + ), + ); + } catch (error) { + batchError = error; + } + try { + if (await sandboxCredentialPolicyPromotionCompleted(env, sessionId, sandboxId, registration)) { + return true; + } + } catch (readError) { + try { + if ( + await sandboxCredentialPolicyPromotionCompleted(env, sessionId, sandboxId, registration) + ) { + return true; + } + } catch { + // Preserve the first observable failure when verification remains unavailable. + } + if (batchError) throw batchError; + throw readError; + } + if (batchError) throw batchError; + return false; +} + +async function sandboxCredentialPolicyPromotionCompleted( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, +): Promise { + const db = database(env); const active = await db .selectFrom("interactive_session_credential_policies") .select(["lookup_id", "state", "registration_generation", "registration_claim"]) @@ -623,7 +1119,14 @@ export async function finishSandboxCredentialPolicyRegistration( .where("sandbox_id", "=", sandboxId) .where("lookup_id", "in", registration.lookupIds) .execute(); + const staged = await db + .selectFrom("interactive_session_credential_policy_registrations") + .select("registration_generation") + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .executeTakeFirst(); return ( + !staged && active.length === registration.lookupIds.length && active.every( (row) => @@ -643,13 +1146,9 @@ export async function abandonSandboxCredentialPolicyRegistration( ): Promise { const now = Date.now(); await database(env) - .updateTable("interactive_session_credential_policies") + .updateTable("interactive_session_credential_policy_registrations") .set({ - state: sql<"registering" | "cleanup_pending">`CASE - WHEN ${sandboxCredentialPolicyCleanupAuthorizedCondition(sessionId, sandboxId, now)} - THEN 'cleanup_pending' - ELSE 'registering' - END`, + state: "cleanup_pending", registration_claim: null, registration_claim_expires_at: null, last_error: reason, @@ -662,6 +1161,112 @@ export async function abandonSandboxCredentialPolicyRegistration( .execute(); } +export function sandboxCredentialPolicyPromotionQueries( + env: RuntimeEnv, + sessionId: string, + sandboxId: string, + registration: SandboxCredentialPolicyRegistration, + registrationExpiresAt: number, + ownershipFence: SandboxCredentialPolicyOwnershipFence, + now: number, +): CompilableQuery[] { + const promotionAuthorized = sql` + EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'registering' + AND registration_generation = ${registration.generation} + AND registration_claim = ${registration.claim} + AND registration_claim_expires_at = ${registrationExpiresAt} + AND registration_claim_expires_at > ${now} + ) + AND ${noLivePolicyTableRegistrationCondition(sessionId, sandboxId, now)} + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'cleanup_pending' + ) + AND ${sandboxCredentialPolicyOwnerCondition(sessionId, sandboxId, ownershipFence, now)} + `; + const promotions = registration.lookupIds.map( + (lookupId) => sql` + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + attempt_count, + last_attempt_at, + last_error, + cleanup_claim, + cleanup_claim_expires_at, + created_at, + updated_at + ) + SELECT + ${sessionId}, + ${sandboxId}, + ${lookupId}, + 'active', + ${registration.generation}, + NULL, + NULL, + 0, + NULL, + NULL, + NULL, + NULL, + ${now}, + ${now} + WHERE ${promotionAuthorized} + ON CONFLICT(session_id, sandbox_id, lookup_id) DO UPDATE SET + state = 'active', + registration_generation = excluded.registration_generation, + registration_claim = NULL, + registration_claim_expires_at = NULL, + last_error = NULL, + cleanup_claim = NULL, + cleanup_claim_expires_at = NULL, + updated_at = excluded.updated_at + WHERE interactive_session_credential_policies.state != 'cleanup_pending' + AND ${promotionAuthorized} + `, + ); + const promotionComplete = sql` + ( + SELECT count(DISTINCT lookup_id) + FROM interactive_session_credential_policies + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND lookup_id IN (${sql.join(registration.lookupIds)}) + AND state = 'active' + AND registration_generation = ${registration.generation} + AND registration_claim IS NULL + ) = ${registration.lookupIds.length} + `; + return [ + ...promotions, + sql` + DELETE FROM interactive_session_credential_policy_registrations + WHERE session_id = ${sessionId} + AND sandbox_id = ${sandboxId} + AND state = 'registering' + AND registration_generation = ${registration.generation} + AND registration_claim = ${registration.claim} + AND registration_claim_expires_at = ${registrationExpiresAt} + AND registration_claim_expires_at > ${now} + AND ${promotionComplete} + `, + ]; +} + export async function standaloneSandboxPolicyExpiresAt( env: RuntimeEnv, sessionId: string, diff --git a/src/worker/sandbox-credential-policy-rollback.ts b/src/worker/sandbox-credential-policy-rollback.ts new file mode 100644 index 00000000..edd0ba13 --- /dev/null +++ b/src/worker/sandbox-credential-policy-rollback.ts @@ -0,0 +1,133 @@ +import { + credentialPolicyRollbackExpiresAt, + credentialPolicyRollbackRecord, + isCurrentCredentialPolicyGeneration, + type CredentialPolicyRollbackRecord, +} from "../credential-policy-fence.ts"; +import type { + SandboxCredentialPolicy, + SandboxCredentialPolicyRegistration, + StoredSandboxCredentialPolicy, +} from "./session-control-policy.ts"; + +const credentialPolicyRollbackClaimMs = 60_000; + +export type SandboxCredentialPolicyRollbackRecord = + CredentialPolicyRollbackRecord; + +type SessionControlStub = Pick; + +export async function captureSandboxCredentialPolicyRollback( + stub: SessionControlStub, + lookupIds: readonly string[], + expectedGeneration: string | null, + sessionId: string, +): Promise { + const records = await Promise.all( + lookupIds.map(async (lookupId): Promise => { + const response = await stub.fetch( + `https://crabfleet.internal/api/session-control/egress/${encodeURIComponent(lookupId)}`, + ); + if (response.status === 404) return null; + if (!response.ok) throw new Error("sandbox credential policy rollback snapshot failed"); + const generation = response.headers.get("x-crabfleet-policy-generation"); + const policy = (await response.json()) as SandboxCredentialPolicy; + if ( + !isCurrentCredentialPolicyGeneration(generation) || + policy.sessionId !== sessionId || + policy.sandboxId !== lookupId + ) { + throw new Error("sandbox credential policy rollback snapshot is inconsistent"); + } + return { generation, policy }; + }), + ); + const present = records.filter((record): record is SandboxCredentialPolicyRollbackRecord => + Boolean(record), + ); + if (!expectedGeneration) { + if (present.length > 0) { + throw new Error("sandbox credential policy has no durable rollback owner"); + } + return []; + } + if ( + present.length !== lookupIds.length || + present.some((record) => record.generation !== expectedGeneration) + ) { + throw new Error("sandbox credential policy rollback generation is incomplete"); + } + return present; +} + +export function parseSandboxCredentialPolicyRollback( + value: string, + lookupIds: readonly string[], + sessionId: string, +): SandboxCredentialPolicyRollbackRecord[] { + let parsed: unknown; + try { + parsed = JSON.parse(value); + } catch { + throw new Error("sandbox credential policy rollback snapshot is invalid"); + } + if (!Array.isArray(parsed)) { + throw new Error("sandbox credential policy rollback snapshot is invalid"); + } + const records = parsed.map((item) => + credentialPolicyRollbackRecord(item), + ); + if (records.some((record) => !record)) { + throw new Error("sandbox credential policy rollback snapshot is invalid"); + } + const rollback = records as SandboxCredentialPolicyRollbackRecord[]; + const expectedLookups = new Set(lookupIds); + const actualLookups = new Set(rollback.map((record) => record.policy.sandboxId)); + const generations = new Set(rollback.map((record) => record.generation)); + if ( + rollback.some( + (record) => + record.policy.sessionId !== sessionId || !expectedLookups.has(record.policy.sandboxId), + ) || + actualLookups.size !== rollback.length || + (rollback.length > 0 && (rollback.length !== expectedLookups.size || generations.size !== 1)) + ) { + throw new Error("sandbox credential policy rollback snapshot is inconsistent"); + } + return rollback; +} + +export async function restoreSandboxCredentialPolicyRollback( + stub: SessionControlStub, + registration: SandboxCredentialPolicyRegistration, + registrationExpiresAt: number, + rollbackJson: string, + sessionId: string, +): Promise { + const rollback = parseSandboxCredentialPolicyRollback( + rollbackJson, + registration.lookupIds, + sessionId, + ); + const now = Date.now(); + const rollbackExpiresAt = credentialPolicyRollbackExpiresAt( + registrationExpiresAt, + now, + credentialPolicyRollbackClaimMs, + ); + for (const record of rollback) { + const response = await stub.fetch("https://crabfleet.internal/api/session-control/register", { + method: "POST", + body: JSON.stringify({ + generation: record.generation, + registrationClaim: `rollback:${registration.claim}`, + registrationExpiresAt: rollbackExpiresAt, + policy: record.policy, + } satisfies StoredSandboxCredentialPolicy), + headers: { "content-type": "application/json" }, + }); + if (!response.ok) { + throw new Error("sandbox credential policy rollback restore failed"); + } + } +} diff --git a/src/worker/sandbox-credential-policy-scanner.ts b/src/worker/sandbox-credential-policy-scanner.ts index 63956fda..579bb77b 100644 --- a/src/worker/sandbox-credential-policy-scanner.ts +++ b/src/worker/sandbox-credential-policy-scanner.ts @@ -9,11 +9,23 @@ import { database, executeBatch, type Database } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import type { InteractiveSessionStatus } from "./models.ts"; import { + abandonSandboxCredentialPolicyRegistration, + claimSandboxCredentialPolicyRegistrationRecovery, + finishSandboxCredentialPolicyRegistration, + markSandboxCredentialPolicyRegistrationWriteStarted, recordSandboxCredentialPolicyRefs, sandboxCredentialPolicyCleanupAuthorizedCondition, + sandboxCredentialPolicyPersistedLookupIds, + sandboxLookupIds, type SandboxCredentialPolicyOwnershipFence, } from "./sandbox-credential-policy-repository.ts"; +import { parseSandboxCredentialPolicyRollback } from "./sandbox-credential-policy-rollback.ts"; import { sandboxLeaseInfo, sandboxLeasePrefix } from "./sandbox-lease.ts"; +import { + sandboxCredentialPolicyRegistrationLookupIds, + sandboxCredentialPolicyRollbackLookupIds, + type SandboxCredentialPolicyRegistration, +} from "./session-control-policy.ts"; const credentialPolicyScanLimit = 32; export const credentialPolicyProvisioningStaleMs = 15 * 60_000; @@ -45,19 +57,60 @@ export type CredentialPolicyScanRow = { standalone_updated_at: number | null; }; +type CredentialPolicyOwnershipRow = Pick< + CredentialPolicyScanRow, + | "session_id" + | "sandbox_id" + | "matched_session_id" + | "session_adapter" + | "session_lease_id" + | "session_sandbox_refresh_sandbox_id" + | "session_sandbox_refresh_claim" + | "session_sandbox_refresh_claim_expires_at" + | "matched_standalone_id" + | "standalone_state" + | "standalone_claim" + | "standalone_claim_expires_at" +>; + +type StagedCredentialPolicyScanRow = CredentialPolicyOwnershipRow & { + registration_generation: string; + registration_claim: string; + registration_claim_expires_at: number; + lookup_ids_json: string | null; + rollback_policies_json: string | null; +}; + export type SandboxCredentialPolicyExists = ( env: RuntimeEnv, sandboxId: string, generation: string, + lookupIds?: readonly string[], ) => Promise; +export type RestoreSandboxCredentialPolicyRollback = (input: { + registration: SandboxCredentialPolicyRegistration; + registrationExpiresAt: number; + rollbackJson: string; + sessionId: string; +}) => Promise; + export async function scanCredentialPolicyCleanupPage( env: RuntimeEnv, now: number, policyExists: SandboxCredentialPolicyExists, sessionId?: string, + restoreRollback?: RestoreSandboxCredentialPolicyRollback, ): Promise { const db = database(env); + await scanStagedCredentialPolicyRegistrations( + env, + db, + now, + policyExists, + sessionId, + restoreRollback, + ); const state = sessionId ? null : await db @@ -250,6 +303,185 @@ export async function scanCredentialPolicyCleanupPage( } } +async function sandboxCredentialPolicyRollbackIsSuperseded( + db: Kysely, + sessionId: string, + sandboxId: string, + lookupIds: readonly string[], + rollbackGeneration: string | null, +): Promise { + const rows = await db + .selectFrom("interactive_session_credential_policies") + .select("registration_generation") + .where("session_id", "=", sessionId) + .where("sandbox_id", "=", sandboxId) + .where("lookup_id", "in", [...new Set(lookupIds)]) + .where("state", "=", "active") + .where("registration_claim", "is", null) + .execute(); + return rows.some( + (row) => + isCurrentCredentialPolicyGeneration(row.registration_generation) && + row.registration_generation !== rollbackGeneration, + ); +} + +async function scanStagedCredentialPolicyRegistrations( + env: RuntimeEnv, + db: Kysely, + now: number, + policyExists: SandboxCredentialPolicyExists, + sessionId?: string, + restoreRollback?: RestoreSandboxCredentialPolicyRollback, +): Promise { + const sessionFilter = sessionId ? sql`AND registration.session_id = ${sessionId}` : sql``; + const result = await sql` + SELECT + registration.session_id, + registration.sandbox_id, + registration.registration_generation, + registration.registration_claim, + registration.registration_claim_expires_at, + registration.lookup_ids_json, + registration.rollback_policies_json, + session.id AS matched_session_id, + session.adapter AS session_adapter, + session.lease_id AS session_lease_id, + session.sandbox_refresh_sandbox_id AS session_sandbox_refresh_sandbox_id, + session.sandbox_refresh_claim AS session_sandbox_refresh_claim, + session.sandbox_refresh_claim_expires_at AS session_sandbox_refresh_claim_expires_at, + standalone.id AS matched_standalone_id, + standalone.state AS standalone_state, + standalone.ownership_claim AS standalone_claim, + standalone.ownership_claim_expires_at AS standalone_claim_expires_at + FROM interactive_session_credential_policy_registrations AS registration + LEFT JOIN interactive_sessions AS session ON session.id = registration.session_id + LEFT JOIN standalone_sandbox_provisions AS standalone + ON standalone.id = registration.session_id + AND standalone.sandbox_id = registration.sandbox_id + WHERE registration.state = 'registering' + AND registration.registration_claim_expires_at <= ${now} + ${sessionFilter} + ORDER BY registration.updated_at ASC + LIMIT ${credentialPolicyScanLimit} + `.execute(db); + for (const row of result.rows) { + try { + const rollbackLookupIds = + row.rollback_policies_json === null + ? [] + : sandboxCredentialPolicyRollbackLookupIds(row.rollback_policies_json, row.session_id); + const persistedLookupIds = await sandboxCredentialPolicyPersistedLookupIds( + env, + row.session_id, + row.sandbox_id, + ); + const currentLookupIds = sandboxLookupIds(env, row.sandbox_id); + const registration: SandboxCredentialPolicyRegistration = { + generation: row.registration_generation, + claim: row.registration_claim, + lookupIds: sandboxCredentialPolicyRegistrationLookupIds( + row.lookup_ids_json, + row.sandbox_id, + currentLookupIds, + [...persistedLookupIds, ...rollbackLookupIds], + ), + }; + const ownershipFence = credentialPolicyScanOwnershipFence(row, now); + if (!ownershipFence) { + await abandonSandboxCredentialPolicyRegistration( + env, + row.session_id, + row.sandbox_id, + registration, + "sandbox credential policy owner is no longer current", + ); + continue; + } + const recovery = await claimSandboxCredentialPolicyRegistrationRecovery( + env, + row.session_id, + row.sandbox_id, + registration, + row.registration_claim_expires_at, + ownershipFence, + ); + if (!recovery) continue; + if ( + (await policyExists( + env, + row.sandbox_id, + row.registration_generation, + registration.lookupIds, + )) && + (await finishSandboxCredentialPolicyRegistration( + env, + row.session_id, + row.sandbox_id, + recovery.registration, + ownershipFence, + )) + ) { + continue; + } + if (row.rollback_policies_json !== null) { + if (!restoreRollback) throw new Error("sandbox credential policy rollback is unavailable"); + const rollbackExpiresAt = await markSandboxCredentialPolicyRegistrationWriteStarted( + env, + row.session_id, + row.sandbox_id, + recovery.registration, + ownershipFence, + ); + if (!rollbackExpiresAt) { + throw new Error("sandbox credential policy registration claim was revoked"); + } + const rollbackGeneration = + parseSandboxCredentialPolicyRollback( + row.rollback_policies_json, + rollbackLookupIds, + row.session_id, + )[0]?.generation ?? null; + const rollbackSuperseded = await sandboxCredentialPolicyRollbackIsSuperseded( + db, + row.session_id, + row.sandbox_id, + [...persistedLookupIds, ...currentLookupIds, ...rollbackLookupIds], + rollbackGeneration, + ); + if (rollbackSuperseded) { + await abandonSandboxCredentialPolicyRegistration( + env, + row.session_id, + row.sandbox_id, + recovery.registration, + "sandbox credential policy generation advanced before rollback", + ); + continue; + } + await restoreRollback({ + registration: { + ...recovery.registration, + lookupIds: rollbackLookupIds, + }, + registrationExpiresAt: rollbackExpiresAt, + rollbackJson: row.rollback_policies_json, + sessionId: row.session_id, + }); + } + await abandonSandboxCredentialPolicyRegistration( + env, + row.session_id, + row.sandbox_id, + recovery.registration, + "sandbox credential policy registration did not complete", + ); + } catch (error) { + console.error("staged sandbox credential policy recovery failed", error); + } + } +} + async function readCredentialPolicyScanPage( db: Kysely, cursor: number, @@ -341,7 +573,7 @@ async function repairActiveSandboxCredentialPolicyRegistration( } export function credentialPolicyScanOwnershipFence( - row: CredentialPolicyScanRow, + row: CredentialPolicyOwnershipRow, now: number, ): SandboxCredentialPolicyOwnershipFence | null { if ( diff --git a/src/worker/session-cleanup.ts b/src/worker/session-cleanup.ts index 48f8e9e8..542926e4 100644 --- a/src/worker/session-cleanup.ts +++ b/src/worker/session-cleanup.ts @@ -13,13 +13,18 @@ import { cleanupSessionLogArchiveObjects } from "./session-log-archive.ts"; const terminalCleanupDeletePending = 2; type SessionReference = string | RawBuilder; -function hasNoCredentialPolicy(sessionId: SessionReference): RawBuilder { +export function hasNoCredentialPolicyLifecycle(sessionId: SessionReference): RawBuilder { return sql` NOT EXISTS ( SELECT 1 FROM interactive_session_credential_policies WHERE session_id = ${sessionId} ) + AND NOT EXISTS ( + SELECT 1 + FROM interactive_session_credential_policy_registrations + WHERE session_id = ${sessionId} + ) `; } @@ -140,7 +145,7 @@ export async function readInteractiveSessionCleanupCandidates( .selectAll() .where("status", "in", deadInteractiveSessionStatuses) .where("terminal_finalize_pending", "=", 0) - .where(hasNoCredentialPolicy(sql.ref("interactive_sessions.id"))) + .where(hasNoCredentialPolicyLifecycle(sql.ref("interactive_sessions.id"))) .where(archiveCoversAllEvents(sql.ref("interactive_sessions.id"))) .where(sql` EXISTS ( @@ -185,7 +190,7 @@ export async function deleteFinalizedInteractiveSession( .where("status", "=", row.status) .where("updated_at", "=", row.updated_at) .where("terminal_finalize_pending", "=", 0) - .where(hasNoCredentialPolicy(row.id)) + .where(hasNoCredentialPolicyLifecycle(row.id)) .where(hasNoActiveDescendants(row.id)) .where(sql` ${archive ? 1 : 0} = 1 diff --git a/src/worker/session-control-do.ts b/src/worker/session-control-do.ts index 1c951f36..03cf5d35 100644 --- a/src/worker/session-control-do.ts +++ b/src/worker/session-control-do.ts @@ -7,10 +7,25 @@ import { } from "../credential-policy-fence.ts"; import type { FleetSandboxPolicySummary } from "../fleet-state.ts"; import { - forwardGitHubActionsRelayMessage, + attachGitHubActionsRunnerProtocol, + attachGitHubActionsViewerProtocol, + createGitHubActionsRelayGeneration, + gitHubActionsRelayGeneration, + gitHubActionsRelayUsesGenerations, + githubActionsLegacyRelayGeneration, githubActionsRelayRole, + githubActionsRunnerProtocolHeader, + githubActionsRunnerProtocolQuery, + githubActionsViewerGenerationHeader, + githubActionsViewerProtocolHeader, + githubActionsViewerProtocolQuery, notifyGitHubActionsViewers, + parseGitHubActionsRunnerProtocol, + parseGitHubActionsRunnerProtocolOffer, + parseGitHubActionsViewerProtocol, + relayGitHubActionsWebSocketMessage, replaceGitHubActionsRunner, + type GitHubActionsRelayProtocol, } from "../github-actions-runtime.ts"; import type { RuntimeEnv } from "./env.ts"; import { json } from "./http.ts"; @@ -51,14 +66,27 @@ export class SessionControlDO extends DurableObject { request.method === "GET" && url.pathname === "/api/session-control/github-actions/runner" ) { - return this.openGitHubActionsRelay("runner"); + const offeredProtocol = parseGitHubActionsRunnerProtocolOffer( + request.headers.get(githubActionsRunnerProtocolHeader), + ); + return this.openGitHubActionsRelay( + "runner", + offeredProtocol ?? + parseGitHubActionsRunnerProtocol( + url.searchParams.get(githubActionsRunnerProtocolQuery), + ), + offeredProtocol, + ); } if ( request.method === "GET" && url.pathname === "/api/session-control/github-actions/viewer" ) { - return this.openGitHubActionsRelay("viewer"); + return this.openGitHubActionsRelay( + "viewer", + parseGitHubActionsViewerProtocol(url.searchParams.get(githubActionsViewerProtocolQuery)), + ); } if ( @@ -200,8 +228,9 @@ export class SessionControlDO extends DurableObject { socket.close(1008, "unknown relay peer"); return; } - forwardGitHubActionsRelayMessage( + relayGitHubActionsWebSocketMessage( role, + socket, message, this.ctx.getWebSockets("github-actions-runner"), this.ctx.getWebSockets("github-actions-viewer"), @@ -214,6 +243,7 @@ export class SessionControlDO extends DurableObject { notifyGitHubActionsViewers( this.ctx.getWebSockets("github-actions-viewer"), "runner_disconnected", + gitHubActionsRelayGeneration(socket) ?? githubActionsLegacyRelayGeneration, ); } } @@ -222,24 +252,56 @@ export class SessionControlDO extends DurableObject { socket.close(1011, "relay peer error"); } - private openGitHubActionsRelay(role: "runner" | "viewer"): Response { + private openGitHubActionsRelay( + role: "runner" | "viewer", + protocol: GitHubActionsRelayProtocol | null = null, + confirmedRunnerProtocol: GitHubActionsRelayProtocol | null = null, + ): Response { const pair = new WebSocketPair(); const client = pair[0]; const server = pair[1]; + let initialRunnerGeneration: string | undefined; if (role === "runner") { + const generation = createGitHubActionsRelayGeneration(); replaceGitHubActionsRunner(this.ctx.getWebSockets("github-actions-runner")); + attachGitHubActionsRunnerProtocol(server, protocol, generation); this.ctx.acceptWebSocket(server, ["github-actions-runner"]); notifyGitHubActionsViewers( this.ctx.getWebSockets("github-actions-viewer"), "runner_connected", + generation, ); } else { + attachGitHubActionsViewerProtocol(server, protocol); this.ctx.acceptWebSocket(server, ["github-actions-viewer"]); - if (this.ctx.getWebSockets("github-actions-runner").length === 0) { - server.send(JSON.stringify({ type: "runner_waiting" })); + const runner = this.ctx + .getWebSockets("github-actions-runner") + .find((socket) => socket.readyState === WebSocket.OPEN); + if (!runner) { + if (gitHubActionsRelayUsesGenerations(server)) { + initialRunnerGeneration = "none"; + } + notifyGitHubActionsViewers([server], "runner_waiting"); + } else if (gitHubActionsRelayUsesGenerations(server)) { + initialRunnerGeneration = + gitHubActionsRelayGeneration(runner) ?? githubActionsLegacyRelayGeneration; + notifyGitHubActionsViewers([server], "runner_connected", initialRunnerGeneration); } } - return new Response(null, { status: 101, webSocket: client }); + const responseInit: ResponseInit = { status: 101, webSocket: client }; + if (role === "runner" && confirmedRunnerProtocol) { + responseInit.headers = { + [githubActionsRunnerProtocolHeader]: confirmedRunnerProtocol, + }; + } else if (role === "viewer" && protocol) { + responseInit.headers = { + [githubActionsViewerProtocolHeader]: protocol, + ...(initialRunnerGeneration + ? { [githubActionsViewerGenerationHeader]: initialRunnerGeneration } + : {}), + }; + } + return new Response(null, responseInit); } } diff --git a/src/worker/session-control-policy.ts b/src/worker/session-control-policy.ts index f927ed78..a8acf392 100644 --- a/src/worker/session-control-policy.ts +++ b/src/worker/session-control-policy.ts @@ -1,5 +1,6 @@ import { credentialPolicyCleanupMatches, + credentialPolicyRollbackRecord, isCurrentCredentialPolicyGeneration, type CredentialPolicyGenerationRecord, type CredentialPolicyGenerationTombstone, @@ -29,6 +30,75 @@ export type SandboxCredentialPolicyRegistration = { lookupIds: string[]; }; +export function sandboxCredentialPolicyRegistrationLookupIds( + value: string | null | undefined, + sandboxId: string, + expectedLookupIds: readonly string[], + historicalLookupIds: readonly string[] = [], +): string[] { + if (value !== null && value !== undefined) { + try { + const parsed = JSON.parse(value) as unknown; + if ( + Array.isArray(parsed) && + parsed.length > 0 && + parsed.every(validSandboxCredentialPolicyLookupId) + ) { + const lookupIds = [...new Set(parsed)]; + if (lookupIds.includes(sandboxId)) return lookupIds; + } + } catch { + // Malformed persisted state cannot authorize additional lookup identities. + } + return [sandboxId]; + } + const currentLookupIds = expectedLookupIds.includes(sandboxId) ? expectedLookupIds : []; + return [ + ...new Set( + [sandboxId, ...currentLookupIds, ...historicalLookupIds].filter( + validSandboxCredentialPolicyLookupId, + ), + ), + ]; +} + +export function sandboxCredentialPolicyRollbackLookupIds( + value: string, + sessionId: string, +): string[] { + let parsed: unknown; + try { + parsed = JSON.parse(value); + } catch { + throw new Error("sandbox credential policy rollback snapshot is invalid"); + } + if (!Array.isArray(parsed)) { + throw new Error("sandbox credential policy rollback snapshot is invalid"); + } + const generations = new Set(); + const lookupIds = parsed.map((item) => { + const record = credentialPolicyRollbackRecord(item); + const lookupId = record?.policy.sandboxId; + if ( + !record || + record.policy.sessionId !== sessionId || + !validSandboxCredentialPolicyLookupId(lookupId) + ) { + throw new Error("sandbox credential policy rollback snapshot is invalid"); + } + generations.add(record.generation); + return lookupId; + }); + if (new Set(lookupIds).size !== lookupIds.length || generations.size > 1) { + throw new Error("sandbox credential policy rollback snapshot is inconsistent"); + } + return lookupIds; +} + +function validSandboxCredentialPolicyLookupId(value: unknown): value is string { + return typeof value === "string" && value.length > 0 && value.length <= 200; +} + export function storedSandboxCredentialPolicy( value: unknown, ): StoredSandboxCredentialPolicy | undefined { diff --git a/src/worker/session-creation.ts b/src/worker/session-creation.ts index 6d7ea8fa..79fe9fcc 100644 --- a/src/worker/session-creation.ts +++ b/src/worker/session-creation.ts @@ -16,6 +16,7 @@ import type { InteractiveProvisionResult, SandboxProvisionOwnership, } from "./provisioning/types.ts"; +import type { RuntimeAdapterWorkspaceRegistration } from "./provisioning/runtime-adapter-release-service.ts"; export type InteractiveSessionCreateOptions = { createdBy?: string; @@ -45,6 +46,7 @@ export type InteractiveSessionCreationReservation = { export type InteractiveSessionProvisionRecoveryInput = { sessionId: string; adapterName: string; + adapterRegistration: RuntimeAdapterWorkspaceRegistration | null; sandboxLeasePrefix: string; now: number; }; @@ -123,6 +125,7 @@ export type InteractiveSessionCreationStore = { stopSupersededAdapter( sessionId: string, adapterWorkspaceId: string, + registration: RuntimeAdapterWorkspaceRegistration | null, createPending: boolean, now: number, ): Promise; @@ -166,7 +169,8 @@ export class InteractiveSessionCreationService { request.createdBy, lineage, ); - const preparationReservation = Boolean(options.afterReserve || supervisedRootSessionId); + // Keep the insert removable until request evidence is durable. + const preparationReservation = true; const now = this.store.now(); for (let attempt = 0; attempt < this.configuration.maximumAttempts; attempt += 1) { @@ -240,6 +244,12 @@ export class InteractiveSessionCreationService { { sessionId: id, adapterName: this.configuration.adapterName, + adapterRegistration: context.adapterControlPlane + ? { + profile: request.profile, + controlPlane: context.adapterControlPlane, + } + : null, sandboxLeasePrefix: this.configuration.sandboxLeasePrefix, now: this.store.now(), }, @@ -279,6 +289,7 @@ export class InteractiveSessionCreationService { } try { await prepare?.(); + await this.store.recordRequest(reservation.id, reservation.insertedAt); } catch (error) { await this.store.rollbackReservation(reservation.id, reservation.insertedAt); throw error; @@ -290,7 +301,6 @@ export class InteractiveSessionCreationService { reservation.adapterWorkspaceId, ); } - await this.store.recordRequest(reservation.id, reservation.insertedAt); return provision(); } @@ -372,6 +382,7 @@ export class InteractiveSessionCreationService { await this.store.stopSupersededAdapter( input.sessionId, result.adapterWorkspaceId, + input.adapterRegistration, result.createPending === true, input.now, ); diff --git a/src/worker/session-events.ts b/src/worker/session-events.ts index 4bc81a48..c46a95e1 100644 --- a/src/worker/session-events.ts +++ b/src/worker/session-events.ts @@ -106,6 +106,14 @@ export async function appendInteractiveSessionEventRecord( input: AppendInteractiveSessionEventInput, archive: InteractiveSessionEventArchive = (sessionId, now) => archiveInteractiveSessionLogs(env, sessionId, now), +): Promise { + await persistInteractiveSessionEventRecord(env, input); + await archive(input.sessionId, input.now).catch(() => undefined); +} + +export async function persistInteractiveSessionEventRecord( + env: RuntimeEnv, + input: AppendInteractiveSessionEventInput, ): Promise { const db = database(env); await executeBatch(env, [ @@ -117,7 +125,6 @@ export async function appendInteractiveSessionEventRecord( }), terminalFinalizationPendingQuery(db, input.sessionId), ]); - await archive(input.sessionId, input.now).catch(() => undefined); } export async function appendStructuredInteractiveSessionEventRecord( @@ -307,7 +314,7 @@ function structuredPayloadJson(value: unknown): string { if ( !Object.hasOwn(record, "version") || typeof version !== "number" || - !Number.isInteger(version) || + !Number.isSafeInteger(version) || version < 1 ) { throw badRequest("payload.version must be a positive integer"); @@ -432,6 +439,9 @@ function canonicalJsonValue( } if (typeof value === "number") { if (!Number.isFinite(value)) throw badRequest("payload must contain valid JSON values"); + if (Number.isInteger(value) && (!Number.isSafeInteger(value) || Object.is(value, -0))) { + throw badRequest("payload integers must be safe and round-trippable"); + } return value; } if (typeof value !== "object") { @@ -466,6 +476,9 @@ function consumePayloadMembers(budget: { members: number }, count: number): void } function assertPayloadStringSize(value: string): void { + if (hasLoneSurrogate(value)) { + throw badRequest("payload strings must contain valid Unicode"); + } if (encoder.encode(value).byteLength > structuredEventPayloadMaxStringBytes) { throw badRequest( `payload strings must be at most ${structuredEventPayloadMaxStringBytes} UTF-8 bytes`, @@ -477,12 +490,27 @@ function requiredString(value: unknown, name: string, maximum: number): string { if (typeof value !== "string") throw badRequest(`${name} must be a string`); const normalized = value.trim(); if (!normalized) throw badRequest(`${name} is required`); + if (hasLoneSurrogate(normalized)) throw badRequest(`${name} must contain valid Unicode`); if (normalized.length > maximum) { throw badRequest(`${name} must be at most ${maximum} characters`); } return normalized; } +function hasLoneSurrogate(value: string): boolean { + for (let index = 0; index < value.length; index += 1) { + const code = value.charCodeAt(index); + if (code >= 0xd800 && code <= 0xdbff) { + const next = value.charCodeAt(index + 1); + if (!(next >= 0xdc00 && next <= 0xdfff)) return true; + index += 1; + } else if (code >= 0xdc00 && code <= 0xdfff) { + return true; + } + } + return false; +} + function clean(value: unknown, maximum: number): string { return String(value ?? "") .trim() diff --git a/src/worker/session-grant-repository.ts b/src/worker/session-grant-repository.ts index 3fba3f2d..24649bf3 100644 --- a/src/worker/session-grant-repository.ts +++ b/src/worker/session-grant-repository.ts @@ -2,7 +2,7 @@ import { sql } from "kysely"; import { trustedProxyConfigured } from "../trusted-proxy-auth.ts"; import { authorize, trustedProxyAutomaticRole } from "./auth.ts"; -import { database, executeBatch, type InteractiveSessionGrantRow } from "./database.ts"; +import { database, type InteractiveSessionGrantRow } from "./database.ts"; import type { RuntimeEnv } from "./env.ts"; import type { User } from "./models.ts"; import type { @@ -152,12 +152,12 @@ export class InteractiveSessionGrantRepository { async revoke(sessionId: string, subject: string, now = Date.now()): Promise { const db = database(this.env); - const deleted = await db - .deleteFrom("interactive_session_grants") - .where("session_id", "=", sessionId) - .where("subject", "=", subject) - .executeTakeFirst(); - if (deleted.numDeletedRows < 1n) return false; + const grantExists = sql`EXISTS ( + SELECT 1 + FROM interactive_session_grants + WHERE session_id = ${sessionId} + AND subject = ${subject} + )`; const clearPendingControl = db .updateTable("interactive_sessions") .set({ @@ -166,6 +166,7 @@ export class InteractiveSessionGrantRepository { control_requested_at: null, }) .where("id", "=", sessionId) + .where(grantExists) .where("control_requested_by_subject", "=", subject); const clearDelegatedControl = db .updateTable("interactive_sessions") @@ -176,17 +177,26 @@ export class InteractiveSessionGrantRepository { control_expires_at: null, }) .where("id", "=", sessionId) + .where(grantExists) .where("controller_subject", "=", subject); const advanceSessionRevision = db .updateTable("interactive_sessions") .set({ updated_at: sql`MAX(updated_at + 1, ${now})` }) - .where("id", "=", sessionId); - await executeBatch(this.env, [ - clearPendingControl, - clearDelegatedControl, - advanceSessionRevision, - ]); - return true; + .where("id", "=", sessionId) + .where(grantExists); + const deleteGrant = db + .deleteFrom("interactive_session_grants") + .where("session_id", "=", sessionId) + .where("subject", "=", subject); + const results = await this.env.DB.batch( + [clearPendingControl, clearDelegatedControl, advanceSessionRevision, deleteGrant].map( + (query) => { + const compiled = query.compile(); + return this.env.DB.prepare(compiled.sql).bind(...compiled.parameters); + }, + ), + ); + return (results[3]?.meta.changes ?? 0) > 0; } } diff --git a/src/worker/session-reconciliation.ts b/src/worker/session-reconciliation.ts index 93961870..ee64ece4 100644 --- a/src/worker/session-reconciliation.ts +++ b/src/worker/session-reconciliation.ts @@ -5,6 +5,7 @@ import { database, type CompilableQuery, type InteractiveSessionRow } from "./da import type { RuntimeEnv } from "./env.ts"; import type { InteractiveSessionStatus } from "./models.ts"; import type { InteractiveProvisionResult } from "./provisioning/types.ts"; +import type { RuntimeAdapterWorkspaceRegistration } from "./provisioning/runtime-adapter-release-service.ts"; import type { InteractiveSession } from "./session-model.ts"; export type RuntimeAdapterReconciliationTransition = { @@ -37,6 +38,7 @@ export type InteractiveSessionReconciliationStore = { stopSuperseded( sessionId: string, adapterWorkspaceId: string, + registration: RuntimeAdapterWorkspaceRegistration | null, createPending: boolean, now: number, ): Promise; @@ -120,6 +122,12 @@ export class InteractiveSessionReconciliationService { await this.store.stopSuperseded( row.id, inspection.adapterWorkspaceId, + row.adapter_control_plane + ? { + profile: row.profile, + controlPlane: row.adapter_control_plane, + } + : null, inspection.createPending === true, this.store.now(), ); diff --git a/src/worker/session-terminal-finalization.ts b/src/worker/session-terminal-finalization.ts index 3887ce8f..82609de8 100644 --- a/src/worker/session-terminal-finalization.ts +++ b/src/worker/session-terminal-finalization.ts @@ -1,4 +1,4 @@ -import { sql, type Kysely } from "kysely"; +import { sql, type Kysely, type RawBuilder } from "kysely"; import { retainedRuntimeAdapterFailureMessage } from "../runtime-adapter.ts"; import { completeTerminalFinalization } from "../terminal-finalization.ts"; @@ -6,6 +6,7 @@ import { database, executeBatch, type CompilableQuery, type Database } from "./d import type { RuntimeEnv } from "./env.ts"; import { deadInteractiveSessionStatuses } from "./models.ts"; import { archiveInteractiveSessionLogs } from "./session-log-archive.ts"; +import { hasNoCredentialPolicyLifecycle } from "./session-cleanup.ts"; import { countInteractiveSessionEvents } from "./session-repository.ts"; export type TerminalInteractiveSessionStatus = "stopped" | "expired" | "failed"; @@ -41,6 +42,50 @@ export function terminalInteractiveSessionFinalizationMessage( return status === "expired" ? "interactive workspace expired" : "interactive workspace stopped"; } +export function terminalFinalizationClearPendingQuery( + id: string, + status: TerminalInteractiveSessionStatus, + sessionLogsEnabled: boolean, +): RawBuilder { + return sql` + UPDATE interactive_sessions + SET terminal_finalize_pending = 0 + WHERE id = ${id} + AND status = ${status} + AND terminal_finalize_pending > 0 + AND EXISTS ( + SELECT 1 + FROM interactive_session_log_archives AS archive + WHERE archive.session_id = interactive_sessions.id + AND archive.session_updated_at = interactive_sessions.updated_at + ) + AND ${hasNoCredentialPolicyLifecycle(id)} + AND COALESCE( + ( + SELECT event_count + FROM interactive_session_log_archives + WHERE session_id = ${id} + ), + -1 + ) >= ( + SELECT count(*) + FROM interactive_session_events + WHERE session_id = ${id} + ) + AND ( + ${sessionLogsEnabled ? 1 : 0} = 0 + OR EXISTS ( + SELECT 1 + FROM interactive_session_log_archives + WHERE session_id = ${id} + AND events_key IS NOT NULL + AND transcript_key IS NOT NULL + AND summary_key IS NOT NULL + ) + ) + `; +} + export async function finalizeTerminalInteractiveSession( env: RuntimeEnv, id: string, @@ -105,47 +150,11 @@ export async function finalizeTerminalInteractiveSession( }, archive: () => archiveInteractiveSessionLogs(env, id, now, { force: true }), clearPending: async () => { - const cleared = await sql` - UPDATE interactive_sessions - SET terminal_finalize_pending = 0 - WHERE id = ${id} - AND status = ${status} - AND terminal_finalize_pending > 0 - AND EXISTS ( - SELECT 1 - FROM interactive_session_log_archives AS archive - WHERE archive.session_id = interactive_sessions.id - AND archive.session_updated_at = interactive_sessions.updated_at - ) - AND NOT EXISTS ( - SELECT 1 - FROM interactive_session_credential_policies - WHERE session_id = ${id} - ) - AND COALESCE( - ( - SELECT event_count - FROM interactive_session_log_archives - WHERE session_id = ${id} - ), - -1 - ) >= ( - SELECT count(*) - FROM interactive_session_events - WHERE session_id = ${id} - ) - AND ( - ${env.SESSION_LOGS ? 1 : 0} = 0 - OR EXISTS ( - SELECT 1 - FROM interactive_session_log_archives - WHERE session_id = ${id} - AND events_key IS NOT NULL - AND transcript_key IS NOT NULL - AND summary_key IS NOT NULL - ) - ) - `.execute(db); + const cleared = await terminalFinalizationClearPendingQuery( + id, + status, + Boolean(env.SESSION_LOGS), + ).execute(db); if ((cleared.numAffectedRows ?? 0n) > 0n) return true; const current = await db .selectFrom("interactive_sessions") diff --git a/src/worker/terminal-hub.ts b/src/worker/terminal-hub.ts index 88888d27..f1fb5d3d 100644 --- a/src/worker/terminal-hub.ts +++ b/src/worker/terminal-hub.ts @@ -13,17 +13,46 @@ import { normalizeWebSocketMessageData, sendOutputAcknowledgement, } from "@openclaw/libterminal/worker"; +import { + createGitHubActionsRelayInputId, + encodeGitHubActionsRelayInput, + githubActionsRuntime, + parseGitHubActionsRelayInputAcknowledgement, + parseGitHubActionsRelayOutput, + parseGitHubActionsRelayEvent, + type GitHubActionsRelayInputAcknowledgement, +} from "../github-actions-runtime.ts"; import { redactedAdapterMessage } from "../runtime-adapter.ts"; import { badRequest, unauthorized } from "./http.ts"; import type { User } from "./models.ts"; import type { InteractiveSession } from "./session-model.ts"; const encoder = new TextEncoder(); -const terminalFrameLimits = { maxFrameBytes: 16 * 1024 * 1024 }; +const terminalMaxFrameBytes = 16 * 1024 * 1024; +const terminalFrameLimits = { maxFrameBytes: terminalMaxFrameBytes }; +const terminalInputQueueMaxBytes = terminalMaxFrameBytes; +const terminalInputQueueMaxFrames = 32; +const terminalInputAcknowledgementTimeoutMs = 5_000; +const terminalPreAuthorizationRelayEventMax = 32; + +type PendingTerminalInputAcknowledgement = { + inputId: string; + runnerGeneration: number | string; + promise: Promise; + resolve(result: TerminalInputAcknowledgementResult): void; + timeout: ReturnType; +}; + +type TerminalInputAcknowledgementResult = GitHubActionsRelayInputAcknowledgement & { + deliveryUnknown?: boolean; +}; export type TerminalUpstream = { socket: WebSocket; markConnected: () => Promise; + inputAcknowledgements?: boolean; + inputGenerations?: boolean; + initialRunnerGeneration?: string | null; outputAcknowledgements: boolean; }; @@ -37,6 +66,15 @@ export type TerminalHubSubscription = { viewCheck: ReturnType | null; cols: number; rows: number; + inputAcknowledgements: boolean; + inputQueue: Promise; + inputQueueBytes: number; + inputQueueFrames: number; + inputQueueRejections: number; + inputQueueRejectionScheduled: boolean; + inputGenerations: boolean; + pendingInputAcknowledgements: Map; + runnerGeneration: number | string; outputAcknowledgements: boolean; outputAcknowledgementBytes: number; }; @@ -87,6 +125,7 @@ export type TerminalHubDependencies = { error?: unknown, ): Promise; markDetached(user: User | null, sessionId: string, message: string): Promise; + inputAcknowledgementTimeoutMs?: number; }; type PendingTerminalSubscription = { @@ -119,6 +158,7 @@ export class TerminalHub { ok: true, version: TERMINAL_WS_VERSION, multiplex: true, + inputAcknowledgements: true, }); const closeSubscription = (id: string, code = 1000, reason = "unsubscribed") => { @@ -154,6 +194,7 @@ export class TerminalHub { ok: true, version: TERMINAL_WS_VERSION, multiplex: true, + inputAcknowledgements: true, }); return; } @@ -212,22 +253,108 @@ export class TerminalHub { return; } if (frame.type === TerminalMessageType.Input || frame.type === TerminalMessageType.Key) { - const canInput = await subscription.canInput(); - updateTerminalInputCapability(server, subscription, canInput); - if (!canInput) { + if ( + subscription.inputQueueFrames >= terminalInputQueueMaxFrames || + subscription.inputQueueBytes + frame.payload.byteLength > terminalInputQueueMaxBytes + ) { + scheduleTerminalInputBacklogRejection(server, subscription, frame.sessionId); return; } - if (subscription.upstream.readyState === WebSocket.OPEN) { - const inputs = await this.dependencies.inputPayloads( - subscription, - user, - frame.payload, - ); - for (const [index, input] of inputs.entries()) { - if (index > 0) await sleep(index === inputs.length - 1 ? 80 : 2); - subscription.upstream.send(input); - } - } + subscription.inputQueueBytes += frame.payload.byteLength; + subscription.inputQueueFrames += 1; + subscription.inputQueue = subscription.inputQueue + .catch(() => undefined) + .then(async () => { + try { + const canInput = await subscription.canInput(); + updateTerminalInputCapability(server, subscription, canInput); + if (!canInput) { + sendTerminalJson(server, TerminalMessageType.Event, frame.sessionId, { + type: "input-rejected", + error: "terminal control is not granted", + }); + return; + } + if (subscription.upstream.readyState !== WebSocket.OPEN) { + sendTerminalJson(server, TerminalMessageType.Error, frame.sessionId, { + error: "terminal upstream is not open", + }); + return; + } + const inputs = await this.dependencies.inputPayloads( + subscription, + user, + frame.payload, + ); + const acknowledgements: PendingTerminalInputAcknowledgement[] = []; + for (const [index, input] of inputs.entries()) { + if (index > 0) await sleep(index === inputs.length - 1 ? 80 : 2); + if ( + subscriptions.get(frame.sessionId) !== subscription || + subscription.upstream.readyState !== WebSocket.OPEN + ) { + sendTerminalJson(server, TerminalMessageType.Error, frame.sessionId, { + error: "terminal upstream is not open", + }); + return; + } + const inputId = subscription.inputAcknowledgements + ? createGitHubActionsRelayInputId() + : null; + const acknowledgement = inputId + ? beginTerminalInputAcknowledgement( + subscription, + inputId, + subscription.runnerGeneration, + this.dependencies.inputAcknowledgementTimeoutMs ?? + terminalInputAcknowledgementTimeoutMs, + ) + : null; + if (acknowledgement) acknowledgements.push(acknowledgement); + try { + subscription.upstream.send( + inputId + ? encodeGitHubActionsRelayInput( + inputId, + input, + subscription.inputGenerations + ? String(subscription.runnerGeneration) + : undefined, + ) + : input, + ); + } catch { + if (acknowledgement) { + completeTerminalInputAcknowledgement( + subscription, + acknowledgement.inputId, + { + inputId: acknowledgement.inputId, + accepted: false, + error: "terminal upstream send failed", + ...(subscription.inputGenerations + ? { generation: String(acknowledgement.runnerGeneration) } + : {}), + }, + ); + break; + } + sendTerminalJson(server, TerminalMessageType.Error, frame.sessionId, { + error: "terminal upstream send failed", + }); + return; + } + } + await reportTerminalInputCompletion( + server, + frame.sessionId, + acknowledgements.map((acknowledgement) => acknowledgement.promise), + ); + } finally { + subscription.inputQueueBytes -= frame.payload.byteLength; + subscription.inputQueueFrames -= 1; + } + }); return; } if (frame.type === TerminalMessageType.Resize) { @@ -302,8 +429,12 @@ export class TerminalHub { }); return; } - if (subscriptions.has(id)) { - sendTerminalJson(client, TerminalMessageType.Event, id, { type: "subscribed" }); + const existingSubscription = subscriptions.get(id); + if (existingSubscription) { + sendTerminalJson(client, TerminalMessageType.Event, id, { + type: "subscribed", + canInput: existingSubscription.canInputGranted, + }); return; } @@ -377,7 +508,46 @@ export class TerminalHub { return; } const upstream = upstreamConnection.socket; - if (!(await canView())) { + const inputAcknowledgements = + upstreamConnection.inputAcknowledgements ?? session.runtime === githubActionsRuntime; + const inputGenerations = upstreamConnection.inputGenerations ?? false; + let runnerGeneration: number | string = inputGenerations + ? (upstreamConnection.initialRunnerGeneration ?? "none") + : 0; + const bufferedRelayEvents: Array< + NonNullable> + > = []; + let captureRelayEvents = true; + upstream.addEventListener("message", (event) => { + if (!captureRelayEvents || !inputAcknowledgements) return; + const relayEvent = parseSynchronousGitHubActionsRelayEvent(event.data); + if (!relayEvent) return; + if (bufferedRelayEvents.length === terminalPreAuthorizationRelayEventMax) { + bufferedRelayEvents.shift(); + } + bufferedRelayEvents.push(relayEvent); + if (relayEvent.generation) { + if (relayEvent.type === "runner_disconnected") { + if (runnerGeneration === relayEvent.generation) runnerGeneration = "none"; + } else { + runnerGeneration = relayEvent.generation; + } + } else if (relayEvent.type === "runner_connected" && !inputGenerations) { + runnerGeneration = (runnerGeneration as number) + 1; + } + }); + let canViewNow: boolean; + try { + canViewNow = await canView(); + } catch (error) { + captureRelayEvents = false; + if (upstream.readyState < WebSocket.CLOSING) { + upstream.close(1011, "view authorization failed"); + } + throw error; + } + if (!canViewNow) { + captureRelayEvents = false; if (upstream.readyState < WebSocket.CLOSING) upstream.close(1008, "share revoked"); sendTerminalJson(client, TerminalMessageType.Error, id, { error: "interactive session not found", @@ -385,6 +555,7 @@ export class TerminalHub { return; } if (!isHubOpen() || client.readyState !== WebSocket.OPEN) { + captureRelayEvents = false; if (upstream.readyState < WebSocket.CLOSING) upstream.close(1000, "client closed"); return; } @@ -400,6 +571,15 @@ export class TerminalHub { viewCheck, cols, rows, + inputAcknowledgements, + inputQueue: Promise.resolve(), + inputQueueBytes: 0, + inputQueueFrames: 0, + inputQueueRejections: 0, + inputQueueRejectionScheduled: false, + inputGenerations, + pendingInputAcknowledgements: new Map(), + runnerGeneration, outputAcknowledgements: outputAcknowledgements && upstreamConnection.outputAcknowledgements, outputAcknowledgementBytes: 0, }; @@ -429,6 +609,7 @@ export class TerminalHub { }); }, 5000); activeSubscription.viewCheck = viewCheck; + captureRelayEvents = false; subscriptions.set(id, activeSubscription); let outputQueue = Promise.resolve(); sendTerminalJson(client, TerminalMessageType.Event, id, { @@ -437,17 +618,123 @@ export class TerminalHub { }); upstream.addEventListener("message", (event) => { const raw = event.data; + const receivedRelayEvent = activeSubscription.inputAcknowledgements + ? parseSynchronousGitHubActionsRelayEvent(raw) + : null; + let runnerGenerationAtReceipt = activeSubscription.runnerGeneration; + if (receivedRelayEvent?.generation) { + runnerGenerationAtReceipt = receivedRelayEvent.generation; + if (receivedRelayEvent.type === "runner_disconnected") { + if (activeSubscription.runnerGeneration === receivedRelayEvent.generation) { + activeSubscription.runnerGeneration = "none"; + } + } else { + activeSubscription.runnerGeneration = receivedRelayEvent.generation; + } + } else if ( + receivedRelayEvent?.type === "runner_connected" && + !activeSubscription.inputGenerations + ) { + runnerGenerationAtReceipt = ++(activeSubscription.runnerGeneration as number); + } outputQueue = outputQueue .catch(() => undefined) .then(async () => { const data = await normalizeWebSocketMessageData(raw); if (client.readyState !== WebSocket.OPEN || !viewGranted) return; - if (typeof data === "string") { - const parsed = parseTerminalControlMessage(data); - if (parsed) { - sendTerminalJson(client, TerminalMessageType.Event, id, parsed); + if (activeSubscription.inputAcknowledgements && typeof data !== "string") { + const inputAcknowledgement = parseGitHubActionsRelayInputAcknowledgement(data); + if (inputAcknowledgement) { + completeTerminalInputAcknowledgement( + activeSubscription, + inputAcknowledgement.inputId, + inputAcknowledgement, + ); return; } + const relayEvent = receivedRelayEvent ?? parseGitHubActionsRelayEvent(data); + if (relayEvent) { + if (relayEvent.type === "runner_disconnected") { + if (relayEvent.generation && activeSubscription.inputGenerations) { + completeTerminalInputAcknowledgements( + activeSubscription, + (pending) => pending.runnerGeneration === relayEvent.generation, + { + accepted: false, + deliveryUnknown: true, + error: + "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ); + } else { + completeAllTerminalInputAcknowledgements(activeSubscription, { + accepted: false, + deliveryUnknown: true, + error: + "terminal input delivery outcome is unknown; the runner may still complete it", + }); + } + } else if (relayEvent.type === "runner_connected") { + if (relayEvent.generation && activeSubscription.inputGenerations) { + runnerGenerationAtReceipt = relayEvent.generation; + activeSubscription.runnerGeneration = relayEvent.generation; + completeTerminalInputAcknowledgements( + activeSubscription, + (pending) => pending.runnerGeneration !== relayEvent.generation, + { + accepted: false, + deliveryUnknown: true, + error: + "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ); + } else { + if (!receivedRelayEvent) { + runnerGenerationAtReceipt = ++(activeSubscription.runnerGeneration as number); + } + completeTerminalInputAcknowledgementsBeforeGeneration( + activeSubscription, + runnerGenerationAtReceipt as number, + { + accepted: false, + deliveryUnknown: true, + error: + "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ); + } + } else if ( + relayEvent.type === "runner_waiting" && + activeSubscription.inputGenerations + ) { + activeSubscription.runnerGeneration = relayEvent.generation ?? "none"; + completeAllTerminalInputAcknowledgements(activeSubscription, { + accepted: false, + error: "GitHub Actions runner disconnected before accepting input", + }); + } + sendTerminalJson(client, TerminalMessageType.Event, id, relayEvent); + return; + } + const relayOutput = parseGitHubActionsRelayOutput(data); + if (!relayOutput) return; + const output = new Uint8Array(relayOutput); + sendTerminalFrame(client, TerminalMessageType.Output, id, output); + if (activeSubscription.outputAcknowledgements) { + activeSubscription.outputAcknowledgementBytes += output.byteLength; + } else if (upstreamConnection.outputAcknowledgements) { + sendOutputAcknowledgement(upstream, output.byteLength); + } + return; + } + if (typeof data === "string") { + if (!activeSubscription.inputAcknowledgements) { + const parsed = parseTerminalControlMessage(data); + if (parsed) { + sendTerminalJson(client, TerminalMessageType.Event, id, parsed); + return; + } + } const output = encoder.encode(data); sendTerminalFrame(client, TerminalMessageType.Output, id, output); if (activeSubscription.outputAcknowledgements) { @@ -466,7 +753,15 @@ export class TerminalHub { } }); }); + for (const relayEvent of bufferedRelayEvents) { + sendTerminalJson(client, TerminalMessageType.Event, id, relayEvent); + } upstream.addEventListener("close", (event) => { + completeAllTerminalInputAcknowledgements(activeSubscription, { + accepted: false, + deliveryUnknown: true, + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }); const closeReason = consumeCloseReason(); const safeUpstreamReason = event.reason ? redactedAdapterMessage( @@ -491,6 +786,11 @@ export class TerminalHub { } }); upstream.addEventListener("error", () => { + completeAllTerminalInputAcknowledgements(activeSubscription, { + accepted: false, + deliveryUnknown: true, + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }); const closeReason = closingReason; if (subscriptions.delete(id)) this.dependencies.releaseInputState(id); if (viewCheck !== null) clearInterval(viewCheck); @@ -519,6 +819,159 @@ export class TerminalHub { } } +function beginTerminalInputAcknowledgement( + subscription: TerminalHubSubscription, + inputId: string, + runnerGeneration: number | string, + timeoutMs: number, +): PendingTerminalInputAcknowledgement { + let resolve!: (result: TerminalInputAcknowledgementResult) => void; + const promise = new Promise((complete) => { + resolve = complete; + }); + const pending: PendingTerminalInputAcknowledgement = { + inputId, + runnerGeneration, + promise, + resolve, + timeout: setTimeout(() => { + if ( + completeAllTerminalInputAcknowledgements(subscription, { + accepted: false, + deliveryUnknown: true, + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }) === 0 + ) { + return; + } + if (subscription.upstream.readyState === WebSocket.OPEN) { + subscription.markClosing("input acknowledgement timed out"); + subscription.upstream.close(1011, "input acknowledgement timed out"); + } + }, timeoutMs), + }; + subscription.pendingInputAcknowledgements.set(inputId, pending); + return pending; +} + +function scheduleTerminalInputBacklogRejection( + socket: WebSocket, + subscription: TerminalHubSubscription, + sessionId: string, +): void { + subscription.inputQueueRejections += 1; + if (subscription.inputQueueRejectionScheduled) return; + subscription.inputQueueRejectionScheduled = true; + subscription.inputQueue = subscription.inputQueue + .catch(() => undefined) + .then(() => { + const rejections = subscription.inputQueueRejections; + subscription.inputQueueRejections = 0; + subscription.inputQueueRejectionScheduled = false; + for (let index = 0; index < rejections; index += 1) { + sendTerminalJson(socket, TerminalMessageType.Event, sessionId, { + type: "input-rejected", + error: "terminal input backlog exceeded", + }); + } + }); +} + +function completeTerminalInputAcknowledgement( + subscription: TerminalHubSubscription, + inputId: string, + result: GitHubActionsRelayInputAcknowledgement, +): boolean { + const pending = subscription.pendingInputAcknowledgements.get(inputId); + if (!pending) return false; + if (subscription.inputGenerations && result.generation !== String(pending.runnerGeneration)) { + return false; + } + subscription.pendingInputAcknowledgements.delete(inputId); + clearTimeout(pending.timeout); + pending.resolve(result); + return true; +} + +function completeAllTerminalInputAcknowledgements( + subscription: TerminalHubSubscription, + result: Omit, +): number { + return completeTerminalInputAcknowledgements(subscription, () => true, result); +} + +function completeTerminalInputAcknowledgementsBeforeGeneration( + subscription: TerminalHubSubscription, + runnerGeneration: number, + result: Omit, +): number { + return completeTerminalInputAcknowledgements( + subscription, + (pending) => (pending.runnerGeneration as number) < runnerGeneration, + result, + ); +} + +function completeTerminalInputAcknowledgements( + subscription: TerminalHubSubscription, + matches: (pending: PendingTerminalInputAcknowledgement) => boolean, + result: Omit, +): number { + const pending = [...subscription.pendingInputAcknowledgements.values()].filter(matches); + for (const acknowledgement of pending) { + subscription.pendingInputAcknowledgements.delete(acknowledgement.inputId); + clearTimeout(acknowledgement.timeout); + acknowledgement.resolve({ inputId: acknowledgement.inputId, ...result }); + } + return pending.length; +} + +function parseSynchronousGitHubActionsRelayEvent( + data: unknown, +): ReturnType { + if (typeof data === "string" || data instanceof ArrayBuffer) { + return parseGitHubActionsRelayEvent(data); + } + if (!ArrayBuffer.isView(data)) return null; + const copied = new Uint8Array(data.byteLength); + copied.set(new Uint8Array(data.buffer, data.byteOffset, data.byteLength)); + return parseGitHubActionsRelayEvent(copied.buffer); +} + +async function reportTerminalInputCompletion( + socket: WebSocket, + sessionId: string, + acknowledgements: Promise[], +): Promise { + if (acknowledgements.length === 0) { + sendTerminalJson(socket, TerminalMessageType.Event, sessionId, { + type: "input-accepted", + }); + return; + } + const results = await Promise.all(acknowledgements); + if (socket.readyState !== WebSocket.OPEN) return; + const unknown = results.find((result) => result.deliveryUnknown); + if (unknown) { + sendTerminalJson(socket, TerminalMessageType.Event, sessionId, { + type: "input-delivery-unknown", + error: unknown.error ?? "terminal input delivery outcome is unknown", + }); + return; + } + const rejection = results.find((result) => !result.accepted); + if (rejection) { + sendTerminalJson(socket, TerminalMessageType.Event, sessionId, { + type: "input-rejected", + error: rejection.error ?? "terminal input was not accepted", + }); + return; + } + sendTerminalJson(socket, TerminalMessageType.Event, sessionId, { + type: "input-accepted", + }); +} + function updateTerminalInputCapability( socket: WebSocket, subscription: TerminalHubSubscription, @@ -606,7 +1059,10 @@ function terminalCloseMessage(code: number, reason: string): string { function isPassiveTerminalClose(reason: string | undefined): boolean { return ( - reason === "unsubscribed" || reason === "client closed" || reason === "no terminals mounted" + reason === "unsubscribed" || + reason === "client closed" || + reason === "no terminals mounted" || + reason === "input acknowledgement timed out" ); } diff --git a/src/worker/worker-application.ts b/src/worker/worker-application.ts index 81ede5c1..87a60fe8 100644 --- a/src/worker/worker-application.ts +++ b/src/worker/worker-application.ts @@ -162,8 +162,12 @@ export class WorkerApplication { return { readState: (request, user) => this.readState(request, user, context), readFleet: (user) => this.readFleetState(user, undefined, context), - registerDesktopHost: (user, id, input) => this.desktopHosts().register(user, id, input), - removeDesktopHost: (user, id) => this.desktopHosts().remove(user, id), + registerDesktopHost: (user, id, input, ownershipMode, publicationID) => + this.desktopHosts().register(user, id, input, ownershipMode, publicationID), + recoverDesktopHost: (user, id, publicationID) => + this.desktopHosts().recover(user, id, publicationID), + removeDesktopHost: (user, id, ownershipToken) => + this.desktopHosts().remove(user, id, ownershipToken), searchGitHubRefs: (number) => this.githubReferenceService().search(number), createCard: async (request, user) => this.cardLifecycleService().create(await readJson(request), user), diff --git a/tests/app-navigation.test.ts b/tests/app-navigation.test.ts index a316ab38..b2ba86ab 100644 --- a/tests/app-navigation.test.ts +++ b/tests/app-navigation.test.ts @@ -1,7 +1,13 @@ import assert from "node:assert/strict"; import test from "node:test"; -import { normalizedAppView, sessionOpenTarget, topOpenDrawer } from "../src/app/app-navigation.js"; +import { + appNavigationLocationState, + normalizedAppView, + sessionOpenTarget, + shouldDisposeTerminalsForNavigation, + topOpenDrawer, +} from "../src/app/app-navigation.js"; test("navigation normalizes app views and closes the topmost drawer", () => { assert.equal(normalizedAppView("board"), "board"); @@ -48,3 +54,38 @@ test("session navigation derives focus and durable route targets", () => { grid: true, }); }); + +test("browser history locations reconcile view, drawers, and session focus", () => { + const focusedSession = appNavigationLocationState({ + pathname: "/sessions/IS-2", + search: "?token=shared", + }); + assert.deepEqual(focusedSession, { + appView: "fleet", + drawers: { sessions: true }, + focusedSessionId: "IS-2", + sharedSessionId: "IS-2", + sharedToken: "shared", + }); + assert.equal(shouldDisposeTerminalsForNavigation(focusedSession), false); + + const sessionGrid = appNavigationLocationState({ pathname: "/sessions", search: "" }); + assert.deepEqual(sessionGrid, { + appView: "fleet", + drawers: { sessions: true }, + focusedSessionId: null, + sharedSessionId: null, + sharedToken: null, + }); + assert.equal(shouldDisposeTerminalsForNavigation(sessionGrid), false); + + const board = appNavigationLocationState({ pathname: "/app/board", search: "" }); + assert.deepEqual(board, { + appView: "board", + drawers: {}, + focusedSessionId: null, + sharedSessionId: null, + sharedToken: null, + }); + assert.equal(shouldDisposeTerminalsForNavigation(board), true); +}); diff --git a/tests/app-routing.test.ts b/tests/app-routing.test.ts index 524e7dfd..a0b110db 100644 --- a/tests/app-routing.test.ts +++ b/tests/app-routing.test.ts @@ -26,6 +26,11 @@ test("app routing parses board and shared session locations", () => { id: null, token: null, }); + assert.deepEqual(parseSessionLink({ pathname: "/sessions/%", search: "?token=ignored" }), { + route: false, + id: null, + token: null, + }); }); test("app and session route builders preserve only owned URL state", () => { diff --git a/tests/application-architecture.test.ts b/tests/application-architecture.test.ts index dec08288..2181449c 100644 --- a/tests/application-architecture.test.ts +++ b/tests/application-architecture.test.ts @@ -59,6 +59,46 @@ test("worker entrypoint delegates OpenClaw and GitHub Actions composition", asyn assert.match(githubActions, /new GitHubActionsWorkStateService\(/); }); +test("GitHub Actions runner protocol is confirmed before the relay socket is accepted", async () => { + const [application, relay] = await Promise.all([ + readFile(new URL("../src/worker/github-actions-application.ts", import.meta.url), "utf8"), + readFile(new URL("../src/worker/session-control-do.ts", import.meta.url), "utf8"), + ]); + + assert.match(application, /headers: gitHubActionsRelayRunnerHeaders\(request\)/); + const attach = relay.indexOf("attachGitHubActionsRunnerProtocol(server, protocol, generation)"); + const accept = relay.indexOf( + 'this.ctx.acceptWebSocket(server, ["github-actions-runner"])', + attach, + ); + assert.notEqual(attach, -1); + assert.ok(accept > attach); + assert.match(relay, /\[githubActionsRunnerProtocolHeader\]: confirmedRunnerProtocol/); +}); + +test("GitHub Actions viewer protocol is requested and attached before relay acceptance", async () => { + const [terminal, relay] = await Promise.all([ + readFile(new URL("../src/worker/interactive-terminal-service.ts", import.meta.url), "utf8"), + readFile(new URL("../src/worker/session-control-do.ts", import.meta.url), "utf8"), + ]); + + assert.match(terminal, /stub\.fetch\(\s*buildGitHubActionsViewerRelayUrl\(\)/); + assert.match(terminal, /gitHubActionsViewerResponseUsesFramedProtocol\(upstreamResponse\)/); + const attach = relay.indexOf("attachGitHubActionsViewerProtocol(server, protocol)"); + const accept = relay.indexOf( + 'this.ctx.acceptWebSocket(server, ["github-actions-viewer"])', + attach, + ); + assert.notEqual(attach, -1); + assert.ok(accept > attach); + assert.match(relay, /\[githubActionsViewerProtocolHeader\]: protocol/); + assert.match(relay, /\[githubActionsViewerGenerationHeader\]: initialRunnerGeneration/); + assert.match( + terminal, + /initialRunnerGeneration: gitHubActionsViewerResponseGeneration\(upstreamResponse\)/, + ); +}); + test("worker entrypoint retains only routing and platform composition", async () => { const entrypoint = await readFile(new URL("../src/index.ts", import.meta.url), "utf8"); diff --git a/tests/card-repository.test.ts b/tests/card-repository.test.ts index eea0b7d8..924d4ea3 100644 --- a/tests/card-repository.test.ts +++ b/tests/card-repository.test.ts @@ -2,8 +2,10 @@ import assert from "node:assert/strict"; import test from "node:test"; import { CardRepository } from "../src/worker/card-repository.ts"; +import type { CardRunClaimInput } from "../src/worker/card-lifecycle-service.ts"; import type { RuntimeEnv } from "../src/worker/env.ts"; import type { User } from "../src/worker/models.ts"; +import { containerCapabilities } from "../src/worker/session-model.ts"; test("private card reads require the stable owner subject", async () => { const current: User = { @@ -110,3 +112,99 @@ test("card list batches related D1 reads below the bind-parameter limit", async assert.equal(executions.filter(({ sql }) => /from "run_attempts"/i.test(sql)).length, 3); assert.equal(executions.filter(({ sql }) => /from events/i.test(sql)).length, 3); }); + +test("card run claims batch the card transition with the run-attempt insert", async () => { + const batches: Array> = []; + const env = { + DB: { + prepare(sql: string) { + return { + bind(...parameters: unknown[]) { + return { + sql, + parameters, + async all() { + return { results: [], meta: { changes: 0 } }; + }, + async run() { + return { meta: { changes: 0 } }; + }, + }; + }, + }; + }, + async batch(statements: Array<{ sql: string; parameters: unknown[] }>) { + batches.push(statements); + return [{ meta: { changes: 1 } }, { meta: { changes: 1 } }]; + }, + } as unknown as D1Database, + } as RuntimeEnv; + const input = { + card: { id: "CY-101" }, + runId: "CY-101-R1", + attempt: 1, + cap: 2, + descriptor: { + runtime: "container", + reason: "repo default", + capabilities: containerCapabilities, + }, + now: 500, + } as CardRunClaimInput; + + assert.equal(await new CardRepository(env).claimRun(input), "claimed"); + assert.equal(batches.length, 1); + assert.equal(batches[0]?.length, 2); + assert.match(batches[0]?.[0]?.sql ?? "", /^\s*update cards/i); + assert.match(batches[0]?.[0]?.sql ?? "", /not exists/i); + assert.match(batches[0]?.[1]?.sql ?? "", /^\s*insert into run_attempts/i); + assert.ok(batches[0]?.[0]?.parameters.includes("CY-101-R1")); + assert.ok(batches[0]?.[1]?.parameters.includes("CY-101-R1")); +}); + +test("duplicate card run claims report active before global capacity", async () => { + const queries: string[] = []; + const env = { + DB: { + prepare(sql: string) { + queries.push(sql); + return { + bind() { + return { + async all() { + if (/from "run_attempts"/i.test(sql)) { + return { results: [{ id: "CY-101-R1" }], meta: { changes: 0 } }; + } + return { results: [{ count: 2 }], meta: { changes: 0 } }; + }, + async run() { + return { meta: { changes: 0 } }; + }, + }; + }, + }; + }, + async batch() { + return [{ meta: { changes: 0 } }, { meta: { changes: 0 } }]; + }, + } as unknown as D1Database, + } as RuntimeEnv; + const input = { + card: { id: "CY-101" }, + runId: "CY-101-R1", + attempt: 1, + cap: 2, + descriptor: { + runtime: "container", + reason: "repo default", + capabilities: containerCapabilities, + }, + now: 500, + } as CardRunClaimInput; + + assert.equal(await new CardRepository(env).claimRun(input), "active"); + const diagnosticQueries = queries.filter((query) => /^\s*select/i.test(query)); + assert.equal(diagnosticQueries.length, 1); + assert.match(diagnosticQueries[0] ?? "", /from "run_attempts"/i); + assert.doesNotMatch(diagnosticQueries[0] ?? "", /count\(\*\)/i); +}); diff --git a/tests/control-plane-routes.test.ts b/tests/control-plane-routes.test.ts index 4ce1a405..8086924b 100644 --- a/tests/control-plane-routes.test.ts +++ b/tests/control-plane-routes.test.ts @@ -6,6 +6,11 @@ import { handleControlPlaneRoute, type ControlPlaneRouteDependencies, } from "../src/worker/routes/control-plane.ts"; +import { + desktopHostOwnershipModeHeader, + desktopHostTokenOwnershipMode, +} from "../src/worker/desktop-host-service.ts"; +import { conflict } from "../src/worker/http.ts"; const viewer: User = { subject: "github:1", @@ -41,20 +46,33 @@ function dependencies(calls: string[]): ControlPlaneRouteDependencies { calls.push(`fleet:${user.login}`); return { handler: "fleet" }; }, - async registerDesktopHost(user, id, input) { - calls.push(`desktop-host:register:${user.login}:${id}:${input.name}`); + async registerDesktopHost(user, id, input, ownershipMode, publicationID) { + calls.push( + `desktop-host:register:${user.login}:${id}:${input.name}:${ownershipMode}:${publicationID ?? "none"}`, + ); + const registration = { + host: { + id, + owner: user.login ?? user.subject, + name: String(input.name), + address: String(input.address), + port: Number(input.port), + createdAt: 1, + updatedAt: 1, + }, + }; + return ownershipMode === desktopHostTokenOwnershipMode + ? { ...registration, ownershipToken: "ownership-token" } + : registration; + }, + async recoverDesktopHost(user, id, publicationID) { + calls.push(`desktop-host:recover:${user.login}:${id}:${String(publicationID)}`); return { - id, - owner: user.login ?? user.subject, - name: String(input.name), - address: String(input.address), - port: Number(input.port), - createdAt: 1, - updatedAt: 1, + ownershipToken: publicationID === "publication-id" ? "ownership-token" : null, }; }, - async removeDesktopHost(user, id) { - calls.push(`desktop-host:remove:${user.login}:${id}`); + async removeDesktopHost(user, id, ownershipToken) { + calls.push(`desktop-host:remove:${user.login}:${id}:${ownershipToken ?? "legacy"}`); }, async searchGitHubRefs(number) { calls.push(`github-refs:${number}`); @@ -151,10 +169,18 @@ test("control-plane read and card routes enforce their role boundaries", async ( test("desktop host routes register and remove only the authenticated user's host", async () => { const calls: string[] = []; const registered = await dispatch( - request("PUT", "/api/desktop-hosts/mac%2Dstudio", { - name: "Mac Studio", - address: "100.64.1.2", - port: 5901, + new Request("https://fleet.example/api/desktop-hosts/mac%2Dstudio", { + method: "PUT", + headers: { + "content-type": "application/json", + [desktopHostOwnershipModeHeader]: desktopHostTokenOwnershipMode, + "x-crabfleet-publication-id": "publication-id", + }, + body: JSON.stringify({ + name: "Mac Studio", + address: "100.64.1.2", + port: 5901, + }), }), viewer, calls, @@ -170,20 +196,121 @@ test("desktop host routes register and remove only the authenticated user's host createdAt: 1, updatedAt: 1, }, + ownershipToken: "ownership-token", }); const removed = await dispatch( - request("DELETE", "/api/desktop-hosts/mac%2Dstudio"), + new Request("https://fleet.example/api/desktop-hosts/mac%2Dstudio", { + method: "DELETE", + headers: { "x-crabfleet-ownership-token": "ownership-token" }, + }), viewer, calls, ); assert.equal(removed?.status, 200); + const recovered = await dispatch( + request("POST", "/api/desktop-hosts/mac%2Dstudio?recover=1", { + publicationID: "publication-id", + }), + viewer, + calls, + ); + assert.deepEqual(await recovered?.json(), { ownershipToken: "ownership-token" }); assert.deepEqual(calls, [ - "desktop-host:register:viewer:mac-studio:Mac Studio", - "desktop-host:remove:viewer:mac-studio", + "desktop-host:register:viewer:mac-studio:Mac Studio:token-v1:publication-id", + "desktop-host:remove:viewer:mac-studio:ownership-token", + "desktop-host:recover:viewer:mac-studio:publication-id", + ]); + + const legacyCalls: string[] = []; + const legacyRegistered = await dispatch( + request("PUT", "/api/desktop-hosts/legacy%2Dstudio", { + name: "Legacy Studio", + address: "100.64.1.3", + port: 5901, + }), + viewer, + legacyCalls, + ); + assert.deepEqual(await legacyRegistered?.json(), { + host: { + id: "legacy-studio", + owner: "viewer", + name: "Legacy Studio", + address: "100.64.1.3", + port: 5901, + createdAt: 1, + updatedAt: 1, + }, + }); + const legacyRemoved = await dispatch( + request("DELETE", "/api/desktop-hosts/legacy%2Dstudio"), + viewer, + legacyCalls, + ); + assert.equal(legacyRemoved?.status, 200); + assert.deepEqual(legacyCalls, [ + "desktop-host:register:viewer:legacy-studio:Legacy Studio:legacy:none", + "desktop-host:remove:viewer:legacy-studio:legacy", ]); }); +test("desktop host routes expose legacy ownership conflicts", async () => { + for (const method of ["PUT", "DELETE"]) { + const calls: string[] = []; + await assert.rejects( + dispatch( + request( + method, + "/api/desktop-hosts/token-owned", + method === "PUT" + ? { + name: "Legacy Studio", + address: "100.64.1.3", + port: 5901, + } + : undefined, + ), + viewer, + calls, + method === "PUT" + ? { + async registerDesktopHost() { + throw conflict("desktop host is owned by a token-aware registration"); + }, + } + : { + async removeDesktopHost() { + throw conflict("desktop host is owned by a token-aware registration"); + }, + }, + ), + (error) => { + assert.equal(status(error), 409); + return true; + }, + ); + } +}); + +test("desktop host recovery rejects malformed encoded ids with a client error", async () => { + const calls: string[] = []; + await assert.rejects( + dispatch( + request("POST", "/api/desktop-hosts/%?recover=1", { + publicationID: "publication-id", + }), + viewer, + calls, + ), + (error) => { + assert.equal(status(error), 400); + return true; + }, + ); + assert.deepEqual(calls, []); +}); + test("card actions derive viewer or maintainer authorization from the action", async () => { for (const action of ["attach", "watch"]) { const calls: string[] = []; diff --git a/tests/credential-policy-fence.test.ts b/tests/credential-policy-fence.test.ts index b00b40ce..f44f857e 100644 --- a/tests/credential-policy-fence.test.ts +++ b/tests/credential-policy-fence.test.ts @@ -3,6 +3,7 @@ import test from "node:test"; import { credentialPolicyCleanupMatches, + credentialPolicyRollbackExpiresAt, credentialPolicyRegistrationAccepted, credentialPolicySandboxIsExpected, type CredentialPolicyGenerationRecord, @@ -67,6 +68,32 @@ test("generation fences isolate new policies from stale cleanup", () => { ); }); +test("newer generations rotate policy for the same session only", () => { + const current = registration("generation-1", "claim-current", 300); + + assert.equal( + credentialPolicyRegistrationAccepted( + current, + undefined, + registration("generation-2", "claim-replacement", 301), + 100, + ), + true, + ); + assert.equal( + credentialPolicyRegistrationAccepted( + current, + undefined, + { + ...registration("generation-2", "claim-replacement", 301), + policy: { sessionId: "IS-102", value: "claim-replacement" }, + }, + 100, + ), + false, + ); +}); + test("same-generation registration claims advance monotonically", () => { const current = registration("generation-1", "claim-current", 300); @@ -124,6 +151,11 @@ test("delayed abandoned registration cannot replace a newer claim", () => { assert.equal(credentialPolicyRegistrationAccepted(newer, undefined, abandoned, 200), false); }); +test("rollback claims advance beyond the generation they replace", () => { + assert.equal(credentialPolicyRollbackExpiresAt(500, 100, 200), 501); + assert.equal(credentialPolicyRollbackExpiresAt(200, 100, 200), 300); +}); + test("live lease refresh fences both current and expected sandbox policies", () => { assert.equal( credentialPolicySandboxIsExpected("sandbox-old", "sandbox-old", null, null, null, 100), diff --git a/tests/deployment.test.ts b/tests/deployment.test.ts index 050a79b3..bf59024b 100644 --- a/tests/deployment.test.ts +++ b/tests/deployment.test.ts @@ -77,6 +77,37 @@ test("configured runtime profiles are allowlisted behaviorally", () => { ); }); +test("deployment reads tolerate adapter migration ambiguity and validate template routes", () => { + assert.equal( + deploymentConfig({ + CRABBOX_RUNTIME_ADAPTER_URL: "https://controller.example.test/adapter", + CRABBOX_RUNTIME_ADAPTER_URL_TEMPLATE: "https://controller.example.test/adapters/{profile}", + }).defaultProfile, + "default", + ); + assert.throws( + () => + deploymentConfig({ + CRABFLEET_DEFAULT_PROFILE: "Desktop.PROFILE_2026", + CRABFLEET_RUNTIME_PROFILES_JSON: JSON.stringify([ + { id: "Desktop.PROFILE_2026", label: "Desktop" }, + ]), + CRABBOX_RUNTIME_ADAPTER_URL_TEMPLATE: "https://controller.example.test/adapters/{profile}", + }), + /runtime profile Desktop\.PROFILE_2026 cannot be routed/, + ); + assert.equal( + deploymentConfig({ + CRABFLEET_DEFAULT_PROFILE: "Desktop.PROFILE_2026", + CRABFLEET_RUNTIME_PROFILES_JSON: JSON.stringify([ + { id: "Desktop.PROFILE_2026", label: "Desktop" }, + ]), + CRABBOX_RUNTIME_ADAPTER_URL: "https://controller.example.test/adapter", + }).defaultProfile, + "Desktop.PROFILE_2026", + ); +}); + test("public and client deployment views exclude server-only routing data", () => { const env: DeploymentEnv = { CRABFLEET_CANONICAL_URL: "https://backend.example", diff --git a/tests/desktop-host-migration.test.ts b/tests/desktop-host-migration.test.ts index b24723e8..85aaae97 100644 --- a/tests/desktop-host-migration.test.ts +++ b/tests/desktop-host-migration.test.ts @@ -9,8 +9,25 @@ test("desktop host migration creates an owner-scoped registry with bounded ports new URL("../migrations/0030_desktop_hosts.sql", import.meta.url), "utf8", ); + const ownershipMigration = readFileSync( + new URL("../migrations/0033_desktop_host_ownership.sql", import.meta.url), + "utf8", + ); database.exec(migration); database.exec(migration); + database.exec(ownershipMigration); + database.exec( + readFileSync( + new URL("../migrations/0038_desktop_host_publication_identity.sql", import.meta.url), + "utf8", + ), + ); + database.exec( + readFileSync( + new URL("../migrations/0041_desktop_host_ownership_errors.sql", import.meta.url), + "utf8", + ), + ); const insert = database.prepare(` INSERT INTO desktop_hosts @@ -25,6 +42,18 @@ test("desktop host migration creates an owner-scoped registry with bounded ports ?.count, 2, ); + assert.equal( + database + .prepare("SELECT ownership_token FROM desktop_hosts WHERE owner_subject = 'github:1'") + .get()?.ownership_token, + "", + ); + assert.equal( + database + .prepare("SELECT publication_id FROM desktop_hosts WHERE owner_subject = 'github:1'") + .get()?.publication_id, + "", + ); assert.throws( () => insert.run("github:3", "bad", "bad", "Bad", "100.64.1.4", 0, 1, 1), /constraint/i, @@ -36,3 +65,151 @@ test("desktop host migration creates an owner-scoped registry with bounded ports "idx_desktop_hosts_owner_updated", ); }); + +test("desktop host publication migration clears identities rotated by old workers", () => { + const database = new DatabaseSync(":memory:"); + for (const migration of [ + "0030_desktop_hosts.sql", + "0033_desktop_host_ownership.sql", + "0038_desktop_host_publication_identity.sql", + "0041_desktop_host_ownership_errors.sql", + ]) { + database.exec(readFileSync(new URL(`../migrations/${migration}`, import.meta.url), "utf8")); + } + database.exec(` + INSERT INTO desktop_hosts ( + owner_subject, id, owner, name, address, port, ownership_token, publication_id, + publication_write_token, created_at, updated_at + ) VALUES ( + 'github:1', 'studio', 'alice', 'Studio', '100.64.1.2', 5901, + 'token-a', 'publication-a', 'token-a', 1, 2 + ); + UPDATE desktop_hosts + SET ownership_token = 'token-b' + WHERE owner_subject = 'github:1' AND id = 'studio'; + `); + + assert.deepEqual( + { + ...database + .prepare(` + SELECT ownership_token, publication_id, publication_write_token + FROM desktop_hosts + WHERE id = 'studio' + `) + .get(), + }, + { ownership_token: "token-b", publication_id: "", publication_write_token: "" }, + ); +}); + +test("desktop host ownership migration blocks old-worker mutations of token-owned rows", () => { + const database = new DatabaseSync(":memory:"); + database.exec( + readFileSync(new URL("../migrations/0030_desktop_hosts.sql", import.meta.url), "utf8"), + ); + database.exec( + readFileSync(new URL("../migrations/0033_desktop_host_ownership.sql", import.meta.url), "utf8"), + ); + database.exec( + readFileSync( + new URL("../migrations/0041_desktop_host_ownership_errors.sql", import.meta.url), + "utf8", + ), + ); + database.exec(` + INSERT INTO desktop_hosts ( + owner_subject, id, owner, name, address, port, ownership_token, created_at, updated_at + ) VALUES + ('github:1', 'owned', 'alice', 'Owned Studio', '100.64.1.2', 5901, 'token-1', 1, 2), + ('github:1', 'legacy', 'alice', 'Legacy Studio', '100.64.1.3', 5901, '', 1, 2); + `); + + const oldWorkerUpsert = database.prepare(` + INSERT INTO desktop_hosts ( + owner_subject, id, owner, name, address, port, created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(owner_subject, id) DO UPDATE SET + owner = excluded.owner, + name = excluded.name, + address = excluded.address, + port = excluded.port, + updated_at = excluded.updated_at + `); + assert.throws( + () => + oldWorkerUpsert.run( + "github:1", + "owned", + "old-worker", + "Overwritten", + "100.64.1.99", + 5902, + 10, + 20, + ), + /token-owned desktop host update requires ownership token/, + ); + oldWorkerUpsert.run( + "github:1", + "legacy", + "old-worker", + "Updated Legacy", + "100.64.1.4", + 5902, + 10, + 20, + ); + + assert.deepEqual( + { + ...database + .prepare(` + SELECT owner, name, address, port, ownership_token, created_at, updated_at + FROM desktop_hosts + WHERE id = 'owned' + `) + .get(), + }, + { + owner: "alice", + name: "Owned Studio", + address: "100.64.1.2", + port: 5901, + ownership_token: "token-1", + created_at: 1, + updated_at: 2, + }, + ); + assert.equal( + database.prepare("SELECT name FROM desktop_hosts WHERE id = 'legacy'").get()?.name, + "Updated Legacy", + ); + + assert.throws( + () => + database.exec("DELETE FROM desktop_hosts WHERE owner_subject = 'github:1' AND id = 'owned'"), + /token-owned desktop host delete requires ownership token/, + ); + database.exec("DELETE FROM desktop_hosts WHERE owner_subject = 'github:1' AND id = 'legacy'"); + assert.equal( + database.prepare("SELECT count(*) AS count FROM desktop_hosts WHERE id = 'owned'").get()?.count, + 1, + ); + assert.equal( + database.prepare("SELECT count(*) AS count FROM desktop_hosts WHERE id = 'legacy'").get() + ?.count, + 0, + ); + + database.exec(` + UPDATE desktop_hosts + SET ownership_token = 'delete-authorized:test' + WHERE owner_subject = 'github:1' AND id = 'owned' AND ownership_token = 'token-1'; + DELETE FROM desktop_hosts + WHERE owner_subject = 'github:1' + AND id = 'owned' + AND ownership_token = 'delete-authorized:test'; + `); + assert.equal(database.prepare("SELECT count(*) AS count FROM desktop_hosts").get()?.count, 0); +}); diff --git a/tests/desktop-host-repository.test.ts b/tests/desktop-host-repository.test.ts index 6404b87d..197c277a 100644 --- a/tests/desktop-host-repository.test.ts +++ b/tests/desktop-host-repository.test.ts @@ -1,9 +1,72 @@ import assert from "node:assert/strict"; +import { readFileSync } from "node:fs"; +import { DatabaseSync } from "node:sqlite"; import test from "node:test"; import { DesktopHostRepository } from "../src/worker/desktop-host-repository.ts"; import type { RuntimeEnv } from "../src/worker/env.ts"; +type BoundStatement = { + execute(): { + results: Record[]; + success: true; + meta: { changes: number; last_row_id?: number }; + }; +}; + +function sqliteRuntimeEnv(sqlite: DatabaseSync): RuntimeEnv { + function execute(sql: string, parameters: unknown[]) { + const statement = sqlite.prepare(sql); + if (/^\s*(?:select|pragma|with)\b|\breturning\b/i.test(sql)) { + const results = statement.all(...parameters).map((row) => ({ ...row })); + const changes = Number(sqlite.prepare("SELECT changes() AS changes").get()?.changes ?? 0); + return { results, success: true as const, meta: { changes } }; + } + const result = statement.run(...parameters); + return { + results: [], + success: true as const, + meta: { + changes: Number(result.changes), + last_row_id: Number(result.lastInsertRowid), + }, + }; + } + return { + DB: { + prepare(sql: string) { + return { + bind(...parameters: unknown[]) { + const bound = { + execute: () => execute(sql, parameters), + async all() { + return bound.execute(); + }, + async run() { + return bound.execute(); + }, + }; + return bound; + }, + }; + }, + async batch(statements: D1PreparedStatement[]) { + sqlite.exec("BEGIN IMMEDIATE"); + try { + const results = statements.map((statement) => + (statement as unknown as BoundStatement).execute(), + ); + sqlite.exec("COMMIT"); + return results; + } catch (error) { + sqlite.exec("ROLLBACK"); + throw error; + } + }, + } as unknown as D1Database, + } as RuntimeEnv; +} + test("desktop host repository scopes reads, upserts, and deletes by owner subject", async () => { const executions: Array<{ sql: string; parameters: unknown[] }> = []; const stored = { @@ -13,6 +76,8 @@ test("desktop host repository scopes reads, upserts, and deletes by owner subjec name: "Studio", address: "100.64.1.2", port: 5901, + ownership_token: "ownership-token", + publication_id: "publication-id", created_at: 1, updated_at: 2, }; @@ -33,6 +98,9 @@ test("desktop host repository scopes reads, upserts, and deletes by owner subjec }, }; }, + async batch() { + return []; + }, } as unknown as D1Database, } as RuntimeEnv; const repository = new DesktopHostRepository(env); @@ -45,6 +113,8 @@ test("desktop host repository scopes reads, upserts, and deletes by owner subjec name: "Studio", address: "100.64.1.2", port: 5901, + ownershipToken: "ownership-token", + publicationID: "publication-id", createdAt: 1, updatedAt: 2, }, @@ -59,15 +129,353 @@ test("desktop host repository scopes reads, upserts, and deletes by owner subjec name: "Studio", address: "100.64.1.2", port: 5901, + ownershipToken: "ownership-token", + publicationID: "publication-id", createdAt: 1, updatedAt: 2, }); assert.equal(upserted.id, "studio"); assert.match(executions[1]?.sql ?? "", /^insert into "desktop_hosts"/i); - assert.match(executions[2]?.sql ?? "", /where "owner_subject" = \? and "id" = \?/i); - assert.deepEqual(executions[2]?.parameters, ["github:1", "studio"]); + assert.match(executions[1]?.sql ?? "", /\breturning \*/i); - await repository.remove("github:1", "studio"); + await repository.remove("github:1", "studio", "ownership-token"); + assert.match(executions[2]?.sql ?? "", /^update "desktop_hosts"/i); + assert.match(executions[2]?.sql ?? "", /"ownership_token" = \?/i); + assert.deepEqual(executions[2]?.parameters.slice(1), ["github:1", "studio", "ownership-token"]); + const deleteMarker = executions[2]?.parameters[0]; + assert.match(String(deleteMarker), /^delete-authorized:/); assert.match(executions[3]?.sql ?? "", /^delete from "desktop_hosts"/i); - assert.deepEqual(executions[3]?.parameters, ["github:1", "studio"]); + assert.deepEqual(executions[3]?.parameters, ["github:1", "studio", deleteMarker]); + + await repository.remove("github:1", "legacy-studio", null); + assert.match(executions[4]?.sql ?? "", /^delete from "desktop_hosts"/i); + assert.match(executions[4]?.sql ?? "", /"ownership_token" = \?/i); + assert.deepEqual(executions[4]?.parameters, ["github:1", "legacy-studio", ""]); +}); + +test("desktop host upsert returns the row written by the same atomic statement", async () => { + const executions: string[] = []; + const written = { + owner_subject: "github:1", + id: "studio", + owner: "alice", + name: "Host A", + address: "100.64.1.2", + port: 5901, + ownership_token: "token-a", + publication_id: "publication-a", + created_at: 1, + updated_at: 2, + }; + const competing = { + ...written, + owner: "bob", + name: "Host B", + address: "100.64.1.3", + ownership_token: "token-b", + updated_at: 3, + }; + const env = { + DB: { + prepare(sql: string) { + executions.push(sql); + return { + bind() { + return { + async all() { + return { + results: [/^insert into "desktop_hosts"/i.test(sql) ? written : competing], + meta: { changes: 1 }, + }; + }, + async run() { + return { meta: { changes: 1 } }; + }, + }; + }, + }; + }, + } as unknown as D1Database, + } as RuntimeEnv; + + const row = await new DesktopHostRepository(env).upsert({ + ownerSubject: written.owner_subject, + id: written.id, + owner: written.owner, + name: written.name, + address: written.address, + port: written.port, + ownershipToken: written.ownership_token, + publicationID: written.publication_id, + createdAt: written.created_at, + updatedAt: written.updated_at, + }); + + assert.equal(executions.length, 1); + assert.match(executions[0] ?? "", /^insert into "desktop_hosts".*\breturning \*/is); + assert.equal(row.name, "Host A"); + assert.equal(row.ownershipToken, "token-a"); }); + +test("legacy desktop host upserts reject token ownership", async () => { + let statement = ""; + const stored = { + owner_subject: "github:1", + id: "studio", + owner: "alice", + name: "Legacy Studio", + address: "100.64.1.2", + port: 5901, + ownership_token: "current-token", + publication_id: "current-publication", + created_at: 1, + updated_at: 2, + }; + const env = { + DB: { + prepare(sql: string) { + statement = sql; + return { + bind() { + return { + async all() { + return { results: [], meta: { changes: 0 } }; + }, + }; + }, + }; + }, + } as unknown as D1Database, + } as RuntimeEnv; + + await assert.rejects( + new DesktopHostRepository(env).upsert({ + ownerSubject: stored.owner_subject, + id: stored.id, + owner: stored.owner, + name: stored.name, + address: stored.address, + port: stored.port, + ownershipToken: "", + publicationID: "", + createdAt: stored.created_at, + updatedAt: stored.updated_at, + }), + (error) => { + assert.equal(status(error), 409); + return true; + }, + ); + + const updateClause = statement.split(/do update set/i)[1] ?? ""; + assert.match(updateClause, /where "desktop_hosts"\."ownership_token" = \?/i); +}); + +test("legacy desktop host writes and cleanup cannot mutate token-owned rows", async () => { + const sqlite = new DatabaseSync(":memory:"); + sqlite.exec( + readFileSync(new URL("../migrations/0030_desktop_hosts.sql", import.meta.url), "utf8"), + ); + sqlite.exec( + readFileSync(new URL("../migrations/0033_desktop_host_ownership.sql", import.meta.url), "utf8"), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0038_desktop_host_publication_identity.sql", import.meta.url), + "utf8", + ), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0041_desktop_host_ownership_errors.sql", import.meta.url), + "utf8", + ), + ); + sqlite.exec(` + INSERT INTO desktop_hosts ( + owner_subject, id, owner, name, address, port, ownership_token, publication_id, + created_at, updated_at + ) VALUES ( + 'github:1', 'studio', 'alice', 'Token Studio', '100.64.1.2', 5901, + 'current-token', 'current-publication', 1, 2 + ) + `); + const repository = new DesktopHostRepository(sqliteRuntimeEnv(sqlite)); + + await assert.rejects( + repository.upsert({ + ownerSubject: "github:1", + id: "studio", + owner: "legacy-worker", + name: "Overwritten Studio", + address: "100.64.1.99", + port: 5902, + ownershipToken: "", + publicationID: "", + createdAt: 10, + updatedAt: 20, + }), + (error) => { + assert.equal(status(error), 409); + return true; + }, + ); + await assert.rejects(repository.remove("github:1", "studio", null), (error) => { + assert.equal(status(error), 409); + return true; + }); + await repository.remove("github:1", "studio", "stale-token"); + assert.equal( + sqlite.prepare("SELECT ownership_token FROM desktop_hosts WHERE id = 'studio'").get() + ?.ownership_token, + "current-token", + ); + + await repository.remove("github:1", "studio", "current-token"); + assert.equal(sqlite.prepare("SELECT count(*) AS count FROM desktop_hosts").get()?.count, 0); +}); + +test("legacy desktop host writes and cleanup preserve legacy rows", async () => { + const sqlite = new DatabaseSync(":memory:"); + for (const migration of [ + "0030_desktop_hosts.sql", + "0033_desktop_host_ownership.sql", + "0038_desktop_host_publication_identity.sql", + "0041_desktop_host_ownership_errors.sql", + ]) { + sqlite.exec(readFileSync(new URL(`../migrations/${migration}`, import.meta.url), "utf8")); + } + const repository = new DesktopHostRepository(sqliteRuntimeEnv(sqlite)); + const host = { + ownerSubject: "github:1", + id: "studio", + owner: "alice", + name: "Legacy Studio", + address: "100.64.1.2", + port: 5901, + ownershipToken: "", + publicationID: "", + createdAt: 1, + updatedAt: 2, + }; + + await repository.upsert(host); + assert.equal( + ( + await repository.upsert({ + ...host, + name: "Updated Legacy Studio", + updatedAt: 3, + }) + ).name, + "Updated Legacy Studio", + ); + await repository.remove(host.ownerSubject, host.id, null); + assert.equal(sqlite.prepare("SELECT count(*) AS count FROM desktop_hosts").get()?.count, 0); +}); + +test("desktop host publication recovery matches only the current publication", async () => { + const sqlite = new DatabaseSync(":memory:"); + for (const migration of [ + "0030_desktop_hosts.sql", + "0033_desktop_host_ownership.sql", + "0038_desktop_host_publication_identity.sql", + "0041_desktop_host_ownership_errors.sql", + ]) { + sqlite.exec(readFileSync(new URL(`../migrations/${migration}`, import.meta.url), "utf8")); + } + const repository = new DesktopHostRepository(sqliteRuntimeEnv(sqlite)); + + await repository.upsert({ + ownerSubject: "github:1", + id: "studio", + owner: "alice", + name: "Studio", + address: "100.64.1.2", + port: 5901, + ownershipToken: "token-b", + publicationID: "publication-b", + createdAt: 1, + updatedAt: 2, + }); + + assert.equal( + await repository.ownershipTokenForPublication("github:1", "studio", "publication-a"), + null, + ); + assert.equal( + await repository.ownershipTokenForPublication("github:1", "studio", "publication-b"), + "token-b", + ); +}); + +test("same-publication retries remain recoverable after the publication migration", async () => { + const sqlite = new DatabaseSync(":memory:"); + for (const migration of [ + "0030_desktop_hosts.sql", + "0033_desktop_host_ownership.sql", + "0038_desktop_host_publication_identity.sql", + "0041_desktop_host_ownership_errors.sql", + ]) { + sqlite.exec(readFileSync(new URL(`../migrations/${migration}`, import.meta.url), "utf8")); + } + const repository = new DesktopHostRepository(sqliteRuntimeEnv(sqlite)); + const host = { + ownerSubject: "github:1", + id: "studio", + owner: "alice", + name: "Studio", + address: "100.64.1.2", + port: 5901, + publicationID: "publication-a", + createdAt: 1, + }; + + await repository.upsert({ + ...host, + ownershipToken: "token-a", + updatedAt: 2, + }); + await repository.upsert({ + ...host, + ownershipToken: "token-b", + updatedAt: 3, + }); + + assert.deepEqual( + { + ...sqlite + .prepare(` + SELECT ownership_token, publication_id, publication_write_token + FROM desktop_hosts + WHERE owner_subject = 'github:1' AND id = 'studio' + `) + .get(), + }, + { + ownership_token: "token-b", + publication_id: "publication-a", + publication_write_token: "token-b", + }, + ); + assert.equal( + await repository.ownershipTokenForPublication("github:1", "studio", "publication-a"), + "token-b", + ); + + sqlite.exec(` + UPDATE desktop_hosts + SET ownership_token = 'token-c' + WHERE owner_subject = 'github:1' AND id = 'studio' + `); + assert.equal( + await repository.ownershipTokenForPublication("github:1", "studio", "publication-a"), + null, + ); +}); + +function status(error: unknown): number | undefined { + return typeof error === "object" && error !== null && "status" in error + ? Number(error.status) + : undefined; +} diff --git a/tests/desktop-host-service.test.ts b/tests/desktop-host-service.test.ts index dddf93c0..97359a2a 100644 --- a/tests/desktop-host-service.test.ts +++ b/tests/desktop-host-service.test.ts @@ -2,7 +2,10 @@ import assert from "node:assert/strict"; import test from "node:test"; import type { DesktopHostRow, DesktopHostStore } from "../src/worker/desktop-host-repository.ts"; -import { DesktopHostService } from "../src/worker/desktop-host-service.ts"; +import { + DesktopHostService, + desktopHostTokenOwnershipMode, +} from "../src/worker/desktop-host-service.ts"; import type { User } from "../src/worker/models.ts"; const alice: User = { @@ -36,20 +39,44 @@ class MemoryDesktopHostStore implements DesktopHostStore { return stored; } - async remove(ownerSubject: string, id: string): Promise { - this.rows.delete(`${ownerSubject}:${id}`); + async ownershipTokenForPublication( + ownerSubject: string, + id: string, + publicationID: string, + ): Promise { + const row = this.rows.get(`${ownerSubject}:${id}`); + return row?.publicationID === publicationID ? row.ownershipToken : null; + } + + async remove(ownerSubject: string, id: string, ownershipToken: string | null): Promise { + const key = `${ownerSubject}:${id}`; + if (this.rows.get(key)?.ownershipToken === (ownershipToken ?? "")) { + this.rows.delete(key); + } } } test("desktop hosts are canonicalized and isolated to their stable owner", async () => { const store = new MemoryDesktopHostStore(); let now = 42; - const service = new DesktopHostService(store, () => now); - const host = await service.register(alice, " Studio.ONE ", { - name: " Peter's Mac Studio ", - address: "100.68.201.40", - port: 5901, - }); + const tokens = ["ownership-1", "ownership-2"]; + const service = new DesktopHostService( + store, + () => now, + () => tokens.shift() ?? "unexpected-token", + ); + const registration = await service.register( + alice, + " Studio.ONE ", + { + name: " Peter's Mac Studio ", + address: "100.68.201.40", + port: 5901, + }, + desktopHostTokenOwnershipMode, + "publication-1", + ); + const host = registration.host; assert.deepEqual(host, { id: "studio.one", @@ -60,22 +87,195 @@ test("desktop hosts are canonicalized and isolated to their stable owner", async createdAt: 42, updatedAt: 42, }); + assert.equal(registration.ownershipToken, "ownership-1"); assert.deepEqual(await service.list(alice), [host]); assert.deepEqual(await service.list(bob), []); now = 84; - const updated = await service.register(alice, host.id, { - name: "Renamed Studio", - address: host.address, - port: host.port, - }); + const updatedRegistration = await service.register( + alice, + host.id, + { + name: "Renamed Studio", + address: host.address, + port: host.port, + }, + desktopHostTokenOwnershipMode, + "publication-2", + ); + const updated = updatedRegistration.host; assert.equal(updated.createdAt, 42); assert.equal(updated.updatedAt, 84); assert.equal(updated.name, "Renamed Studio"); + assert.equal(updatedRegistration.ownershipToken, "ownership-2"); - await service.remove(bob, host.id); + await service.remove(bob, host.id, updatedRegistration.ownershipToken); assert.deepEqual(await service.list(alice), [updated]); - await service.remove(alice, host.id); + await service.remove(alice, host.id, updatedRegistration.ownershipToken); + assert.deepEqual(await service.list(alice), []); +}); + +test("stale desktop host cleanup cannot remove a newer registration", async () => { + const store = new MemoryDesktopHostStore(); + const tokens = ["old-process-token", "new-process-token"]; + const service = new DesktopHostService( + store, + () => 42, + () => tokens.shift() ?? "unexpected-token", + ); + const input = { name: "Studio", address: "100.64.1.2", port: 5901 }; + + const oldRegistration = await service.register( + alice, + "studio", + input, + desktopHostTokenOwnershipMode, + "old-publication", + ); + const newRegistration = await service.register( + alice, + "studio", + { + ...input, + name: "New Studio Process", + }, + desktopHostTokenOwnershipMode, + "new-publication", + ); + + await service.remove(alice, "studio", oldRegistration.ownershipToken); + assert.deepEqual(await service.list(alice), [newRegistration.host]); + + await service.remove(alice, "studio", newRegistration.ownershipToken); + assert.deepEqual(await service.list(alice), []); +}); + +test("tokenless cleanup removes only migrated legacy desktop hosts", async () => { + const store = new MemoryDesktopHostStore(); + const service = new DesktopHostService( + store, + () => 42, + () => "new-process-token", + ); + const legacy: DesktopHostRow = { + ownerSubject: alice.subject, + id: "legacy-studio", + owner: "alice", + name: "Legacy Studio", + address: "100.64.1.2", + port: 5901, + ownershipToken: "", + publicationID: "", + createdAt: 1, + updatedAt: 1, + }; + store.rows.set(`${alice.subject}:${legacy.id}`, legacy); + const registration = await service.register( + alice, + "new-studio", + { + name: "New Studio", + address: "100.64.1.3", + port: 5901, + }, + desktopHostTokenOwnershipMode, + "new-studio-publication", + ); + + await service.remove(alice, legacy.id, null); + await service.remove(alice, registration.host.id, null); + + assert.deepEqual(await service.list(alice), [registration.host]); +}); + +test("ambiguous desktop recovery cannot acquire a newer publication", async () => { + const store = new MemoryDesktopHostStore(); + const tokens = ["old-process-token", "new-process-token"]; + const service = new DesktopHostService( + store, + () => 42, + () => tokens.shift() ?? "unexpected-token", + ); + const input = { name: "Studio", address: "100.64.1.2", port: 5901 }; + + await service.register(alice, "studio", input, desktopHostTokenOwnershipMode, "publication-a"); + const newer = await service.register( + alice, + "studio", + { ...input, name: "Newer Studio" }, + desktopHostTokenOwnershipMode, + "publication-b", + ); + + assert.deepEqual(await service.recover(alice, "studio", "publication-a"), { + ownershipToken: null, + }); + assert.deepEqual(await service.list(alice), [newer.host]); + assert.deepEqual(await service.recover(alice, "studio", "publication-b"), { + ownershipToken: newer.ownershipToken, + }); + assert.deepEqual(await service.recover(bob, "studio", "publication-b"), { + ownershipToken: null, + }); + for (const publicationID of [null, undefined, "", "bad publication", "a".repeat(201)]) { + await assert.rejects( + service.recover(alice, "studio", publicationID), + /desktop host publication id/, + ); + } +}); + +test("token ownership requires a valid publication before minting or persistence", async () => { + const store = new MemoryDesktopHostStore(); + let tokenCreations = 0; + const service = new DesktopHostService( + store, + () => 42, + () => { + tokenCreations += 1; + return "ownership-token"; + }, + ); + const input = { name: "Studio", address: "100.64.1.2", port: 5901 }; + + for (const publicationID of [ + null, + undefined, + "", + "bad publication", + "bad\npublication", + "a".repeat(201), + ]) { + await assert.rejects( + service.register(alice, "studio", input, desktopHostTokenOwnershipMode, publicationID), + /desktop host publication id/, + ); + } + assert.equal(tokenCreations, 0); + assert.equal(store.rows.size, 0); + + const legacy = await service.register(alice, "studio", input, "legacy", "ignored publication"); + assert.equal(legacy.ownershipToken, undefined); + assert.equal(store.rows.get(`${alice.subject}:studio`)?.publicationID, ""); +}); + +test("legacy clients register tokenless rows they can remove after a server upgrade", async () => { + const store = new MemoryDesktopHostStore(); + const service = new DesktopHostService( + store, + () => 42, + () => "must-not-be-created", + ); + + const registration = await service.register(alice, "rolling-upgrade", { + name: "Rolling Upgrade", + address: "100.64.1.4", + port: 5901, + }); + + assert.equal(registration.ownershipToken, undefined); + assert.equal(store.rows.get(`${alice.subject}:rolling-upgrade`)?.ownershipToken, ""); + await service.remove(alice, registration.host.id, null); assert.deepEqual(await service.list(alice), []); }); @@ -98,4 +298,6 @@ test("desktop hosts accept only bounded metadata and Tailscale IPv4 endpoints", } await assert.rejects(service.register(alice, "studio", { ...valid, name: "bad\nname" }), /name/); await assert.rejects(service.register(alice, "studio", { ...valid, port: 0 }), /port/); + await assert.rejects(service.remove(alice, "studio", ""), /ownership token/); + await assert.rejects(service.remove(alice, "studio", "bad token"), /ownership token/); }); diff --git a/tests/dockerfile.test.ts b/tests/dockerfile.test.ts new file mode 100644 index 00000000..6be638f8 --- /dev/null +++ b/tests/dockerfile.test.ts @@ -0,0 +1,24 @@ +import assert from "node:assert/strict"; +import { readFile } from "node:fs/promises"; +import { test } from "node:test"; + +test("Crabbox image pins the default release and requires pinned version overrides", async () => { + const dockerfile = await readFile(new URL("../Dockerfile", import.meta.url), "utf8"); + + assert.match( + dockerfile, + /pinned_checksum="3c41839257e4622e28bcec8b0f0153f19d78d436fd548894a7c7d7726d922611"/, + ); + assert.match( + dockerfile, + /pinned_checksum="4bf87a0d2365441ee2f8cb34183cfd9ebeb065111697eb2d8dc867b3a627fdd2"/, + ); + assert.doesNotMatch(dockerfile, /checksums\.txt/); + assert.match( + dockerfile, + /if \[ "\$CRABBOX_VERSION" = "0\.17\.1" \]; then \\\n\s+checksum="\$pinned_checksum"; \\\n\s+else \\\n\s+checksum="\$override_checksum";/, + ); + assert.match(dockerfile, /an explicit \$checksum_arg is required/); + assert.match(dockerfile, /grep -Eq '\^\[0-9a-f\]\{64\}\$'/); + assert.match(dockerfile, /sha256sum -c -/); +}); diff --git a/tests/embedded-terminal-assets.test.ts b/tests/embedded-terminal-assets.test.ts index ff50477f..e6324ff7 100644 --- a/tests/embedded-terminal-assets.test.ts +++ b/tests/embedded-terminal-assets.test.ts @@ -1,41 +1,40 @@ import assert from "node:assert/strict"; -import { execFile } from "node:child_process"; import { createServer } from "node:http"; import { test } from "node:test"; -import { promisify } from "node:util"; import { loadGhosttyRuntime } from "@openclaw/libterminal/browser"; import { GHOSTTY_ASSET_PATHS, readGhosttyAsset } from "@openclaw/libterminal/node"; import { readGhosttyWorkerAsset } from "@openclaw/libterminal/worker-assets"; -const execFileAsync = promisify(execFile); +import { withGeneratedAssetsForTest } from "./helpers/generated-assets.ts"; test("libterminal Worker Ghostty assets are byte-exact and keep Crabfleet response policy", async () => { - await execFileAsync(process.execPath, ["scripts/generate-assets.mjs"]); - const generated = await import(`../src/generated.ts?terminal-assets=${Date.now()}`); - const { terminalAssetResponse } = await import( - `../src/worker/terminal-assets.ts?terminal-assets=${Date.now()}` - ); - assert.equal(generated.APP_HTML.includes("__GHOSTTY_WASM_PATH__"), false); - assert.equal(generated.APP_HTML.includes(GHOSTTY_ASSET_PATHS.wasm), true); - assert.equal("GHOSTTY_VT_WASM_BASE64" in generated, false); - - for (const pathname of Object.values(GHOSTTY_ASSET_PATHS)) { - const expected = await readGhosttyAsset(pathname); - const workerAsset = readGhosttyWorkerAsset(pathname); - const response = terminalAssetResponse(pathname); - assert.ok(expected); - assert.equal(workerAsset?.contentType, expected.contentType); - assert.equal( - Buffer.compare(Buffer.from(workerAsset?.body ?? []), Buffer.from(expected.body)), - 0, + await withGeneratedAssetsForTest(async () => { + const generated = await import(`../src/generated.ts?terminal-assets=${Date.now()}`); + const { terminalAssetResponse } = await import( + `../src/worker/terminal-assets.ts?terminal-assets=${Date.now()}` ); - assert.equal(response?.status, 200); - assert.equal(response?.headers.get("content-type"), expected.contentType); - assert.equal(response?.headers.get("cache-control"), "no-store"); - assert.equal(Buffer.compare(Buffer.from(await response!.arrayBuffer()), expected.body), 0); - } - assert.equal(terminalAssetResponse("/vendor/unknown.js"), null); + assert.equal(generated.APP_HTML.includes("__GHOSTTY_WASM_PATH__"), false); + assert.equal(generated.APP_HTML.includes(GHOSTTY_ASSET_PATHS.wasm), true); + assert.equal("GHOSTTY_VT_WASM_BASE64" in generated, false); + + for (const pathname of Object.values(GHOSTTY_ASSET_PATHS)) { + const expected = await readGhosttyAsset(pathname); + const workerAsset = readGhosttyWorkerAsset(pathname); + const response = terminalAssetResponse(pathname); + assert.ok(expected); + assert.equal(workerAsset?.contentType, expected.contentType); + assert.equal( + Buffer.compare(Buffer.from(workerAsset?.body ?? []), Buffer.from(expected.body)), + 0, + ); + assert.equal(response?.status, 200); + assert.equal(response?.headers.get("content-type"), expected.contentType); + assert.equal(response?.headers.get("cache-control"), "no-store"); + assert.equal(Buffer.compare(Buffer.from(await response!.arrayBuffer()), expected.body), 0); + } + assert.equal(terminalAssetResponse("/vendor/unknown.js"), null); + }); }); test("Ghostty loader injects the explicit WASM runtime into terminal modules", async () => { diff --git a/tests/fleet-state.test.ts b/tests/fleet-state.test.ts index b90fe3a1..a022518d 100644 --- a/tests/fleet-state.test.ts +++ b/tests/fleet-state.test.ts @@ -204,6 +204,7 @@ test("GitHub Actions sessions are attachable through the Worker relay", () => { runtime: "github_actions", leaseId: "github-actions:s1", attachUrl: null, + ptyAvailable: true, workKey: "openclaw/crabfleet:pr:42", workKind: "pr_repair", workState: "running", @@ -228,6 +229,30 @@ test("GitHub Actions sessions are attachable through the Worker relay", () => { assert.equal(fleet.sessions[0]?.workPhase, "fixing"); }); +test("GitHub Actions sessions require an available Worker relay", () => { + const fleet = buildFleetState( + [ + { + ...baseSession, + runtime: "github_actions", + leaseId: "github-actions:s1", + attachUrl: null, + ptyAvailable: false, + }, + ], + [], + { + canonicalUrl: "https://crabfleet.openclaw.ai", + defaultEgressHosts: [], + generatedAt: 100, + productUrl: "https://clawfleet.ai", + }, + ); + + assert.equal(fleet.totals.attachable, 0); + assert.equal(fleet.sessions[0]?.attachable, false); +}); + test("sandbox lease parser ignores non-sandbox leases", () => { assert.equal( sandboxIdFromLeaseId("sandbox:crabbox-s1-abcd1234:terminal-s1-abcd1234:autostart-v4"), diff --git a/tests/generated-assets.test.ts b/tests/generated-assets.test.ts new file mode 100644 index 00000000..15142249 --- /dev/null +++ b/tests/generated-assets.test.ts @@ -0,0 +1,16 @@ +import assert from "node:assert/strict"; +import { readFile } from "node:fs/promises"; +import test from "node:test"; + +import { withGeneratedAssetsForTest } from "./helpers/generated-assets.ts"; + +test("generated embedded specification matches the canonical markdown", async () => { + await withGeneratedAssetsForTest(async () => { + const { SPEC_MARKDOWN } = await import(`../src/generated.ts?spec-assets=${Date.now()}`); + const source = await readFile(new URL("../docs/spec.md", import.meta.url), "utf8"); + const markdown = source.replace(/^---\n[\s\S]*?\n---\n+/, ""); + + assert.equal(SPEC_MARKDOWN, markdown); + assert.match(SPEC_MARKDOWN, /relay-generation-fenced binary `CFR1` input/); + }); +}); diff --git a/tests/github-actions-docs.test.ts b/tests/github-actions-docs.test.ts new file mode 100644 index 00000000..10fea060 --- /dev/null +++ b/tests/github-actions-docs.test.ts @@ -0,0 +1,139 @@ +import assert from "node:assert/strict"; +import { readFile } from "node:fs/promises"; +import { test } from "node:test"; + +test("the documented Node runner acknowledges only delivered UTF-8 input", async () => { + const [readme, guide] = await Promise.all([ + readFile(new URL("../README.md", import.meta.url), "utf8"), + readFile(new URL("../docs/github-actions-sessions.md", import.meta.url), "utf8"), + ]); + + assert.match(readme, /complete byte-safe encoder, decoder, and Node PTY runner/); + assert.match(readme, /new WebSocket\(runnerPtyUrl, "cfr1-framed-io-v2"\)/); + assert.match(readme, /terminal\.protocol === "cfr1-framed-io-v2"/); + assert.match(readme, /ignore the offer leave `WebSocket\.protocol` empty/); + assert.doesNotMatch(readme, /New runners opt into[\s\S]*cfr1-framed-io-v1/); + assert.doesNotMatch(readme, /encodeCfr1Output|decodeCfr1Input|encodeCfr1Ack/); + assert.match(guide, /new WebSocket\(runnerPtyUrl, "cfr1-framed-io-v2"\)/); + assert.match(guide, /const framed = terminal\.protocol === "cfr1-framed-io-v2"/); + assert.match(guide, /empty protocol means an older relay kept this socket in legacy raw mode/); + assert.match(guide, /terminal\.send\(framed \? encodeUtf8Output\(outputText\) : outputText\)/); + assert.doesNotMatch(guide, /close\(1002, "framed protocol not negotiated"\)/); + assert.doesNotMatch(guide, /throw new Error\("relay did not negotiate cfr1-framed-io-v2"\)/); + assert.doesNotMatch(guide, /searchParams\.set\("runnerProtocol"/); + assert.match(guide, /`runnerProtocol` query remains compatibility-only/); + assert.match(guide, /New runners must not add it, close, or reconnect solely because/); + assert.equal(guide.match(/runnerProtocol/g)?.length, 1); + assert.match(guide, /let pendingInputs = \[\]/); + assert.match(guide, /let inputQueue = Promise\.resolve\(\)/); + assert.match(guide, /let terminalClosed = false;\s+let activeGeneration;/); + assert.match( + guide, + /const input = admitInput\(event\.data\);\s+if \(!input\) return;\s+inputQueue = inputQueue\.then\(\(\) => \{\s+if \(!inputIsActive\(input\)\) \{\s+releaseInputs\(\[input\]\);\s+return;\s+\}\s+return acceptInput\(input\);\s+\}\)/, + ); + assert.doesNotMatch(guide, /\.then\(\(\) => acceptInput\(event\.data\)\)/); + assert.match( + guide, + /if \(framed\) \{\s+if \(activeGeneration === undefined\) \{\s+activeGeneration = input\.generation;\s+\} else if \(input\.generation !== activeGeneration\) \{\s+sendAck\(input, false\);\s+return null;/, + ); + assert.match( + guide, + /function inputIsActive\(input\) \{\s+return \(\s+!terminalClosed &&\s+terminal\.readyState === WebSocket\.OPEN &&\s+\(!framed \|\| input\.generation === activeGeneration\)/, + ); + assert.match( + guide, + /function deactivateTerminal\(\) \{\s+if \(terminalClosed\) return;\s+terminalClosed = true;\s+activeGeneration = undefined;\s+rejectInputs\(takePendingInputs\(\), 1001, "terminal closed"\);\s+closeSteering\(\);/, + ); + assert.match(guide, /terminal\.addEventListener\("close", deactivateTerminal\)/); + assert.match(guide, /terminal\.addEventListener\("error", deactivateTerminal\)/); + assert.match(guide, /pendingInputs\.push\(input\)/); + assert.match(guide, /text = decodeCompleteUtf8\(payload\)/); + assert.match(guide, /if \(text === null\) \{\s+armPendingInputTimer\(\);\s+return;/); + assert.match( + guide, + /const inputs = takePendingInputs\(\);\s+try \{\s+await deliverSteeringInput\(text\);\s+settleInputs\(inputs, true\)/, + ); + assert.match(guide, /catch \{\s+rejectInputs\(inputs, 1011, "steering rejected input"\);/); + assert.match(guide, /const maxAdmittedInputBytes = 16 \* 1024/); + assert.match(guide, /const maxAdmittedInputFrames = 32/); + assert.match(guide, /const maxPendingInputAgeMs = 1_000/); + assert.match(guide, /const nextBytes = admittedInputBytes \+ input\.payload\.byteLength/); + assert.match(guide, /const nextFrames = admittedInputFrames \+ 1/); + assert.match(guide, /nextBytes > maxAdmittedInputBytes/); + assert.match(guide, /nextFrames > maxAdmittedInputFrames/); + assert.match(guide, /if \(framed\) \{\s+sendAck\(input, false\);/); + assert.match(guide, /terminal\.close\(1009, "input backlog exceeded"\)/); + assert.match(guide, /function decodeRawInput\(data\)/); + assert.match(guide, /typeof data === "string"/); + assert.match(guide, /data instanceof ArrayBuffer/); + assert.match(guide, /admittedInputBytes -= input\.payload\.byteLength/); + assert.match(guide, /admittedInputFrames -= inputs\.length/); + assert.match(guide, /clearTimeout\(pendingInputTimer\)/); + assert.match( + guide, + /const inputs = pendingInputs;\s+pendingInputs = \[\];\s+pendingInputBytes = 0;\s+return inputs;/, + ); + assert.match(guide, /new TextDecoder\("utf-8", \{ fatal: true, ignoreBOM: true \}\)/); + assert.doesNotMatch(guide, /inputDecoder\.decode/); + assert.match(guide, /subscribeSteeringOutput\(\(outputText\) => \{/); + assert.match(guide, /encodeUtf8Output\(outputText\)/); + assert.match(guide, /deliberately\s+a UTF-8 text\s+adapter/); + assert.match(guide, /byte-oriented restricted steering adapter/); + assert.match(guide, /Generation-fenced viewers add `viewerProtocol=cfr1-framed-io-v2`/); + assert.match(guide, /stale-generation input is rejected before it\s+can reach/); + assert.match(guide, /encodeAck\(input\.inputId, input\.generation, accepted\)/); + assert.match(guide, /Legacy viewers omit that query/); + assert.match( + guide, + /acknowledgement deadline expires\s+while the runner write may still be in flight[\s\S]*`input-delivery-unknown`, not\s+`input-rejected`, because that write may still complete/, + ); + assert.match( + guide, + /\{"type":"input-delivery-unknown","error":"terminal input delivery outcome is unknown; the runner may still complete it"\}/, + ); + assert.match(guide, /subscribeSteeringExit\(\(\) => \{/); + assert.match(guide, /terminal\.close\(1000, "pty exited"\)/); + assert.match(guide, /must never forward that input to a\s+shell or subprocess/); + assert.doesNotMatch(guide, /spawn\(process\.env\.SHELL|env:\s*process\.env|pty\.write/); + + const messageHandler = guide.indexOf('terminal.addEventListener("message"'); + const admission = guide.indexOf("const input = admitInput(event.data);", messageHandler); + const serialization = guide.indexOf("inputQueue = inputQueue.then(() => {", messageHandler); + const activeCheck = guide.indexOf("if (!inputIsActive(input))", serialization); + const deliveryAdmission = guide.indexOf("return acceptInput(input);", activeCheck); + assert.ok(messageHandler >= 0); + assert.ok(admission > messageHandler); + assert.ok(serialization > admission); + assert.ok(activeCheck > serialization); + assert.ok(deliveryAdmission > activeCheck); + + const timerArm = guide.indexOf("armPendingInputTimer();"); + const batchSnapshot = guide.indexOf("const inputs = takePendingInputs();"); + const delivery = guide.indexOf("await deliverSteeringInput(text);"); + assert.ok(timerArm > guide.indexOf("if (text === null)")); + assert.ok(batchSnapshot > timerArm); + assert.ok(delivery > batchSnapshot); + + const decodeCompleteUtf8 = (payload: Uint8Array) => { + const decoder = new TextDecoder("utf-8", { fatal: true, ignoreBOM: true }); + const text = decoder.decode(payload, { stream: true }); + return new TextEncoder().encode(text).byteLength === payload.byteLength ? text : null; + }; + assert.equal(decodeCompleteUtf8(Uint8Array.from([0xe2, 0x82])), null); + assert.equal(decodeCompleteUtf8(Uint8Array.from([0xe2, 0x82, 0xac])), "\u20ac"); + assert.throws(() => decodeCompleteUtf8(Uint8Array.from([0xe2, 0x28, 0xa1]))); +}); + +test("the architecture documents negotiated and legacy viewer output", async () => { + const architecture = await readFile(new URL("../docs/architecture.md", import.meta.url), "utf8"); + + assert.match(architecture, /Viewer framing is negotiated independently/); + assert.match( + architecture, + /opted-in viewers receive `CFR1` terminal, lifecycle, and acknowledgement/, + ); + assert.match( + architecture, + /legacy viewers retain raw terminal output and JSON control-message fallbacks/, + ); +}); diff --git a/tests/github-actions-event-auth.test.ts b/tests/github-actions-event-auth.test.ts index 70109485..1e63dcd6 100644 --- a/tests/github-actions-event-auth.test.ts +++ b/tests/github-actions-event-auth.test.ts @@ -5,8 +5,14 @@ import { sha256 } from "../src/worker/crypto.ts"; import type { RuntimeEnv } from "../src/worker/env.ts"; import { GitHubActionsApplication, + gitHubActionsRelayRunnerHeaders, + gitHubActionsRelayRunnerUrl, structuredEventRequestMaxBytes, } from "../src/worker/github-actions-application.ts"; +import { + githubActionsGenerationFencedCapability, + githubActionsRunnerProtocolHeader, +} from "../src/github-actions-runtime.ts"; import { terminalAgentEventGraceMs } from "../src/worker/session-agent-auth.ts"; import { handleServiceSessionRoute, @@ -143,6 +149,58 @@ test("GitHub Actions application rejects an event token issued to another sessio assert.equal(subject.mutationCount(), 0); }); +test("GitHub Actions application propagates only the exact runner protocol opt-in", () => { + const base = + "https://fleet.example/api/agent/interactive-sessions/IS-target/runner-pty?agentToken=secret"; + assert.equal( + gitHubActionsRelayRunnerUrl(new Request(base)), + "https://crabfleet.internal/api/session-control/github-actions/runner", + ); + assert.equal( + gitHubActionsRelayRunnerUrl(new Request(`${base}&runnerProtocol=cfr1-framed-io-v1`)), + "https://crabfleet.internal/api/session-control/github-actions/runner?runnerProtocol=cfr1-framed-io-v1", + ); + assert.equal( + gitHubActionsRelayRunnerUrl(new Request(`${base}&runnerProtocol=cfr1-framed-io-v2`)), + "https://crabfleet.internal/api/session-control/github-actions/runner?runnerProtocol=cfr1-framed-io-v2", + ); + assert.equal( + gitHubActionsRelayRunnerUrl(new Request(`${base}&runnerProtocol=cfr1-framed-io-v3`)), + "https://crabfleet.internal/api/session-control/github-actions/runner", + ); + + const offered = gitHubActionsRelayRunnerHeaders( + new Request(base, { + headers: { + [githubActionsRunnerProtocolHeader]: `ignored, ${githubActionsGenerationFencedCapability}`, + }, + }), + ); + assert.equal( + gitHubActionsRelayRunnerUrl( + new Request(base, { + headers: { + [githubActionsRunnerProtocolHeader]: githubActionsGenerationFencedCapability, + }, + }), + ), + "https://crabfleet.internal/api/session-control/github-actions/runner", + ); + assert.equal(offered.get("upgrade"), "websocket"); + assert.equal( + offered.get(githubActionsRunnerProtocolHeader), + githubActionsGenerationFencedCapability, + ); + assert.equal( + gitHubActionsRelayRunnerHeaders( + new Request(base, { + headers: { [githubActionsRunnerProtocolHeader]: "cfr1-framed-io-v1" }, + }), + ).get(githubActionsRunnerProtocolHeader), + null, + ); +}); + test("agent event endpoint rejects a wrong-session token before persistence", async () => { const subject = await authEnvironment(); const application = new GitHubActionsApplication(subject.env, { audit: async () => undefined }); diff --git a/tests/github-actions-repository.test.ts b/tests/github-actions-repository.test.ts index 2f491f86..5873aed0 100644 --- a/tests/github-actions-repository.test.ts +++ b/tests/github-actions-repository.test.ts @@ -17,7 +17,7 @@ type Execution = { kind: "all" | "run"; }; -function runtimeEnv(executions: Execution[]): RuntimeEnv { +function runtimeEnv(executions: Execution[], mutationChanges = 1): RuntimeEnv { const row = sessionRow({ id: "IS-101", runtime: "github_actions", @@ -35,7 +35,7 @@ function runtimeEnv(executions: Execution[]): RuntimeEnv { }, async run() { executions.push({ sql, parameters, kind: "run" }); - return { meta: { changes: 1 } }; + return { meta: { changes: mutationChanges } }; }, }; }, @@ -68,9 +68,19 @@ test("GitHub Actions repository owns registration and lifecycle SQL", async () = now: 100, }), ); - await repository.updateSession("IS-101", registrationUpdate); - await repository.updateSession("IS-101", workStateUpdate); - await repository.updateSession("IS-101", runnerConnectionUpdate); + await repository.updateSession("IS-101", registrationUpdate, { + kind: "registration", + registration: registrationExpectation, + }); + await repository.updateSession("IS-101", workStateUpdate, { + kind: "authenticated", + revision: 190, + terminalStatus: "attached", + }); + await repository.updateSession("IS-101", runnerConnectionUpdate, { + kind: "authenticated", + revision: 290, + }); assert.equal(executions.length, 6); assert.match(executions[0].sql, /select .* from "interactive_sessions"/i); @@ -86,8 +96,67 @@ test("GitHub Actions repository owns registration and lifecycle SQL", async () = for (const execution of executions.slice(3)) { assert.match(execution.sql, /update "interactive_sessions"/i); assert.match(execution.sql, /where "id" = \?/i); + assert.match(execution.sql, /"runtime" = \?/i); assert.ok(execution.parameters.includes("IS-101")); } + assert.match(executions[3].sql, /"updated_at" = \?/i); + assert.match(executions[3].sql, /"status" = \?/i); + assert.match(executions[3].sql, /"work_state" = \?/i); + assert.match(executions[3].sql, /"work_phase" = \?/i); + assert.match(executions[4].sql, /"updated_at" = \?/i); + assert.match(executions[4].sql, /"status" = \?/i); + assert.match(executions[4].sql, /"status" not in/i); + assert.ok(executions[4].parameters.includes("attached")); + assert.ok(executions[4].parameters.includes("expired")); + assert.match(executions[5].sql, /"updated_at" = \?/i); + assert.match(executions[3].sql, /"owner_subject" = \?/i); + assert.doesNotMatch(executions[3].sql, /"work_state" not in/i); + assert.match(executions[4].sql, /"work_state" not in/i); + assert.ok(executions[4].parameters.includes("blocked")); + assert.match(executions[5].sql, /"status" not in/i); + assert.ok(executions[5].parameters.includes("blocked")); +}); + +test("GitHub Actions repository rejects stale or invalid state transitions", async () => { + const executions: Execution[] = []; + const repository = new GitHubActionsRepository(runtimeEnv(executions, 0)); + + await assert.rejects( + repository.updateSession("IS-101", runnerConnectionUpdate, { + kind: "authenticated", + revision: 290, + }), + (error) => { + assert.equal( + typeof error === "object" && error && "status" in error ? error.status : undefined, + 409, + ); + return true; + }, + ); + assert.equal(executions.length, 1); +}); + +test("terminal work-state updates require an observed non-terminal status", async () => { + const executions: Execution[] = []; + const repository = new GitHubActionsRepository(runtimeEnv(executions)); + + await assert.rejects( + repository.updateSession( + "IS-101", + { + ...workStateUpdate, + status: "stopped", + work_state: "completed", + stopped_at: 200, + }, + { kind: "authenticated", revision: 190 }, + ), + { + message: "terminal GitHub Actions update requires expected session status", + }, + ); + assert.equal(executions.length, 0); }); const registrationUpdate: GitHubActionsSessionRegistrationUpdate = { @@ -118,6 +187,14 @@ const registrationUpdate: GitHubActionsSessionRegistrationUpdate = { completion_reason: null, }; +const registrationExpectation = { + agent_token_hash: "prior-agent-hash", + updated_at: 90, + status: "stopped", + work_state: "completed", + work_phase: "finished", +} as const; + const workStateUpdate: GitHubActionsWorkStateUpdate = { status: "attached", summary: "working", diff --git a/tests/github-actions-runner-connection.test.ts b/tests/github-actions-runner-connection.test.ts index b207e9ec..0a916335 100644 --- a/tests/github-actions-runner-connection.test.ts +++ b/tests/github-actions-runner-connection.test.ts @@ -27,21 +27,25 @@ function session(values: Parameters[0] = {}) { function connectionStore(): { store: GitHubActionsRunnerConnectionStore; updates: GitHubActionsRunnerConnectionUpdate[]; + expectedRevisions: number[]; events: string[]; operations: string[]; } { const updates: GitHubActionsRunnerConnectionUpdate[] = []; + const expectedRevisions: number[] = []; const events: string[] = []; const operations: string[] = []; return { updates, + expectedRevisions, events, operations, store: { now: () => 700, - persist: async (_id, values) => { + persist: async (_id, values, expectedRevision) => { operations.push("persist"); updates.push(values); + expectedRevisions.push(expectedRevision); }, appendEvent: async (_id, message) => { operations.push("event"); @@ -52,7 +56,7 @@ function connectionStore(): { } test("waiting runners become active with durable connection evidence", async () => { - const { store, updates, events, operations } = connectionStore(); + const { store, updates, expectedRevisions, events, operations } = connectionStore(); await new GitHubActionsRunnerConnectionService(store).connect( session({ status: "provisioning" }), ); @@ -69,6 +73,7 @@ test("waiting runners become active with durable connection evidence", async () }, ]); assert.deepEqual(events, [githubActionsRunnerConnectedEvent]); + assert.deepEqual(expectedRevisions, [session({ status: "provisioning" }).updatedAt]); assert.deepEqual(operations, ["persist", "event"]); }); @@ -87,6 +92,23 @@ test("reconnecting runners preserve active status, state, and phase", async () = assert.equal(updates[0]?.work_phase, "codex_turn"); }); +test("runner connections retain the exact revision authenticated before a token rotation", async () => { + const { store, expectedRevisions } = connectionStore(); + + await new GitHubActionsRunnerConnectionService(store).connect(session({ updated_at: 400 })); + + assert.deepEqual(expectedRevisions, [400]); +}); + +test("runner connections advance revisions when the authenticated clock is ahead", async () => { + const { store, updates, expectedRevisions } = connectionStore(); + + await new GitHubActionsRunnerConnectionService(store).connect(session({ updated_at: 800 })); + + assert.deepEqual(expectedRevisions, [800]); + assert.equal(updates[0]?.updated_at, 801); +}); + test("runner connections reject non-work sessions", async () => { const { store } = connectionStore(); const service = new GitHubActionsRunnerConnectionService(store); diff --git a/tests/github-actions-runner.test.ts b/tests/github-actions-runner.test.ts new file mode 100644 index 00000000..90f3c9e8 --- /dev/null +++ b/tests/github-actions-runner.test.ts @@ -0,0 +1,531 @@ +import assert from "node:assert/strict"; +import { test } from "node:test"; + +import { + acceptGitHubActionsRunnerInput, + sendGitHubActionsRunnerOutput, +} from "../src/github-actions-runner.ts"; +import { + encodeGitHubActionsRelayInput, + parseGitHubActionsRelayOutput, + parseGitHubActionsRelayInputAcknowledgement, + type GitHubActionsRelaySocket, +} from "../src/github-actions-runtime.ts"; + +function relaySocket(): GitHubActionsRelaySocket & { + closes: Array<{ code: number | undefined; reason: string | undefined }>; + sent: Array; +} { + return { + readyState: WebSocket.OPEN, + closes: [], + sent: [], + send(message) { + this.sent.push(message); + }, + close(code, reason) { + this.closes.push({ code, reason }); + this.readyState = WebSocket.CLOSED; + }, + }; +} + +test("runner acknowledges input only after the PTY write completes", async () => { + const socket = relaySocket(); + let completeWrite!: () => void; + const write = new Promise((resolve) => { + completeWrite = resolve; + }); + const handled = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-one", "steer"), + async (payload) => { + assert.equal(new TextDecoder().decode(payload), "steer"); + await write; + }, + ); + + await Promise.resolve(); + assert.deepEqual(socket.sent, []); + completeWrite(); + assert.equal(await handled, true); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(socket.sent[0]!), { + inputId: "input-one", + accepted: true, + }); +}); + +test("runner serializes concurrent input writes and acknowledgements per socket", async () => { + const socket = relaySocket(); + const writes: string[] = []; + let completeFirstWrite!: () => void; + const firstWrite = new Promise((resolve) => { + completeFirstWrite = resolve; + }); + + const first = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-first", "first"), + async (payload) => { + writes.push(`${new TextDecoder().decode(payload)}:start`); + await firstWrite; + writes.push("first:end"); + }, + ); + const second = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-second", "second"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.deepEqual(writes, ["first:start"]); + assert.deepEqual(socket.sent, []); + + completeFirstWrite(); + assert.deepEqual(await Promise.all([first, second]), [true, true]); + assert.deepEqual(writes, ["first:start", "first:end", "second"]); + assert.deepEqual( + socket.sent.map((message) => parseGitHubActionsRelayInputAcknowledgement(message)), + [ + { inputId: "input-first", accepted: true }, + { inputId: "input-second", accepted: true }, + ], + ); +}); + +test("runner bounds queued frames and rejects overflow without extending a stalled tail", async () => { + const socket = relaySocket(); + let completeFirstWrite!: () => void; + const firstWrite = new Promise((resolve) => { + completeFirstWrite = resolve; + }); + const writes: number[] = []; + const pending = Array.from({ length: 33 }, (_, index) => + acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput(`input-${index}`, new Uint8Array([index])), + async (payload) => { + writes.push(new Uint8Array(payload)[0]!); + if (index === 0) await firstWrite; + }, + ), + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.deepEqual(writes, [0]); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(socket.sent[0]!), { + inputId: "input-32", + accepted: false, + error: "GitHub Actions runner input backlog exceeded", + }); + completeFirstWrite(); + assert.deepEqual(await Promise.all(pending), Array(33).fill(true)); + assert.deepEqual( + writes, + Array.from({ length: 32 }, (_, index) => index), + ); + assert.deepEqual( + socket.sent.map((message) => parseGitHubActionsRelayInputAcknowledgement(message)), + [ + { + inputId: "input-32", + accepted: false, + error: "GitHub Actions runner input backlog exceeded", + }, + ...Array.from({ length: 32 }, (_, index) => ({ + inputId: `input-${index}`, + accepted: true, + })), + ], + ); +}); + +test("runner keeps overflow floods off the accepted input queue", async () => { + const socket = relaySocket(); + let completeFirstWrite!: () => void; + const firstWrite = new Promise((resolve) => { + completeFirstWrite = resolve; + }); + const writes: number[] = []; + const accepted = Array.from({ length: 32 }, (_, index) => + acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput(`accepted-${index}`, new Uint8Array([index])), + async (payload) => { + writes.push(new Uint8Array(payload)[0]!); + if (index === 0) await firstWrite; + }, + ), + ); + + let settledOverflows = 0; + const overflows = Array.from({ length: 512 }, (_, index) => + acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput(`overflow-${index}`, new Uint8Array([index])), + async () => { + assert.fail("overflow input must not reach the PTY"); + }, + ).then((handled) => { + settledOverflows += 1; + return handled; + }), + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.equal(settledOverflows, 512); + assert.deepEqual(await Promise.all(overflows), Array(512).fill(true)); + assert.deepEqual(writes, [0]); + + completeFirstWrite(); + assert.deepEqual(await Promise.all(accepted), Array(32).fill(true)); + assert.deepEqual( + writes, + Array.from({ length: 32 }, (_, index) => index), + ); + assert.equal(socket.sent.length, 544); +}); + +test("runner bounds queued bytes while a PTY write is stalled", async () => { + const socket = relaySocket(); + let completeFirstWrite!: () => void; + const firstWrite = new Promise((resolve) => { + completeFirstWrite = resolve; + }); + let writes = 0; + const first = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-first", new Uint8Array(9 * 1024 * 1024)), + async () => { + writes += 1; + await firstWrite; + }, + ); + const overflow = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-overflow", new Uint8Array(8 * 1024 * 1024)), + async () => { + writes += 1; + }, + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.equal(writes, 1); + assert.equal(await overflow, true); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(socket.sent[0]!), { + inputId: "input-overflow", + accepted: false, + error: "GitHub Actions runner input backlog exceeded", + }); + completeFirstWrite(); + assert.equal(await first, true); + assert.equal(writes, 1); + assert.deepEqual( + socket.sent.map((message) => parseGitHubActionsRelayInputAcknowledgement(message)), + [ + { + inputId: "input-overflow", + accepted: false, + error: "GitHub Actions runner input backlog exceeded", + }, + { inputId: "input-first", accepted: true }, + ], + ); +}); + +test("runner times out a stalled PTY write and retires its socket queue", async () => { + const socket = relaySocket(); + let completeWrite!: () => void; + const blockedWrite = new Promise((resolve) => { + completeWrite = resolve; + }); + const writes: string[] = []; + const first = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-first", "first"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + await blockedWrite; + }, + Date.now, + 10, + ); + const queued = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-queued", "queued"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + Date.now, + 10, + ); + + assert.deepEqual(await Promise.all([first, queued]), [true, true]); + assert.deepEqual(writes, ["first"]); + assert.deepEqual(socket.sent, []); + assert.deepEqual(socket.closes, [ + { code: 1012, reason: "GitHub Actions runner input write timed out" }, + ]); + + assert.equal( + await acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-late", "late"), + async () => { + assert.fail("retired input must not reach the PTY"); + }, + Date.now, + 10, + ), + true, + ); + completeWrite(); + await new Promise((resolve) => setImmediate(resolve)); + assert.deepEqual(socket.sent, []); +}); + +test("runner does not execute queued input after its socket is replaced", async () => { + const replaced = relaySocket(); + const replacement = relaySocket(); + let completeBlockedWrite!: () => void; + const blockedWrite = new Promise((resolve) => { + completeBlockedWrite = resolve; + }); + const writes: string[] = []; + const first = acceptGitHubActionsRunnerInput( + replaced, + encodeGitHubActionsRelayInput("old-first", "old-first", "old-generation"), + async (payload) => { + writes.push(`${new TextDecoder().decode(payload)}:start`); + await blockedWrite; + writes.push("old-first:end"); + }, + ); + const queued = acceptGitHubActionsRunnerInput( + replaced, + encodeGitHubActionsRelayInput("old-queued", "old-queued", "old-generation"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.deepEqual(writes, ["old-first:start"]); + replaced.close(1012, "runner replaced"); + const current = acceptGitHubActionsRunnerInput( + replacement, + encodeGitHubActionsRelayInput("new-input", "new-input", "new-generation"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + ); + assert.equal(await current, true); + + completeBlockedWrite(); + assert.deepEqual(await Promise.all([first, queued]), [true, true]); + assert.deepEqual(writes, ["old-first:start", "new-input", "old-first:end"]); + assert.deepEqual(replaced.sent, []); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(replacement.sent[0]!), { + inputId: "new-input", + accepted: true, + generation: "new-generation", + }); +}); + +test("runner retires queued input when its relay generation changes", async () => { + const socket = relaySocket(); + let completeBlockedWrite!: () => void; + const blockedWrite = new Promise((resolve) => { + completeBlockedWrite = resolve; + }); + const writes: string[] = []; + const first = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("first", "first", "generation-one"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + await blockedWrite; + }, + ); + const queued = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("queued", "queued", "generation-one"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.equal( + await acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("replacement", "replacement", "generation-two"), + async () => { + assert.fail("replacement input must not reach the retired socket"); + }, + ), + true, + ); + assert.deepEqual(socket.closes, [ + { code: 1012, reason: "GitHub Actions runner generation changed" }, + ]); + + completeBlockedWrite(); + assert.deepEqual(await Promise.all([first, queued]), [true, true]); + assert.deepEqual(writes, ["first"]); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(socket.sent[0]!), { + inputId: "replacement", + accepted: false, + error: "GitHub Actions runner generation changed", + generation: "generation-two", + }); +}); + +test("runner rejects queued input that outlives the viewer acknowledgement timeout", async () => { + const socket = relaySocket(); + let completeFirstWrite!: () => void; + const firstWrite = new Promise((resolve) => { + completeFirstWrite = resolve; + }); + let now = 0; + const writes: string[] = []; + const first = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-first", "first"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + await firstWrite; + }, + () => now, + ); + const expired = acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-expired", "expired"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + () => now, + ); + + await Promise.resolve(); + await Promise.resolve(); + assert.deepEqual(writes, ["first"]); + now = 5_000; + completeFirstWrite(); + assert.deepEqual(await Promise.all([first, expired]), [true, true]); + assert.deepEqual(writes, ["first"]); + assert.deepEqual( + socket.sent.map((message) => parseGitHubActionsRelayInputAcknowledgement(message)), + [ + { inputId: "input-first", accepted: true }, + { + inputId: "input-expired", + accepted: false, + error: "GitHub Actions runner input expired", + }, + ], + ); +}); + +test("runner copies the relay generation into its acknowledgement", async () => { + const socket = relaySocket(); + + assert.equal( + await acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-generated", "steer", "generation-one"), + async () => {}, + ), + true, + ); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(socket.sent[0]!), { + inputId: "input-generated", + accepted: true, + generation: "generation-one", + }); +}); + +test("runner retains its relay generation after the input queue drains", async () => { + const socket = relaySocket(); + const writes: string[] = []; + + assert.equal( + await acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("first", "first", "generation-one"), + async (payload) => { + writes.push(new TextDecoder().decode(payload)); + }, + ), + true, + ); + assert.equal( + await acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("replacement", "replacement", "generation-two"), + async () => { + assert.fail("replacement input must not reach the PTY after the queue drains"); + }, + ), + true, + ); + + assert.deepEqual(writes, ["first"]); + assert.deepEqual(socket.closes, [ + { code: 1012, reason: "GitHub Actions runner generation changed" }, + ]); + assert.deepEqual( + socket.sent.map((message) => parseGitHubActionsRelayInputAcknowledgement(message)), + [ + { inputId: "first", accepted: true, generation: "generation-one" }, + { + inputId: "replacement", + accepted: false, + error: "GitHub Actions runner generation changed", + generation: "generation-two", + }, + ], + ); +}); + +test("runner rejects failed writes and ignores unframed terminal data", async () => { + const socket = relaySocket(); + + assert.equal( + await acceptGitHubActionsRunnerInput( + socket, + encodeGitHubActionsRelayInput("input-two", "steer"), + async () => { + throw new Error("PTY closed"); + }, + ), + true, + ); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(socket.sent[0]!), { + inputId: "input-two", + accepted: false, + error: "GitHub Actions runner did not accept terminal input", + }); + assert.equal(await acceptGitHubActionsRunnerInput(socket, "raw input", async () => {}), false); + assert.equal(socket.sent.length, 1); +}); + +test("runner output helper envelopes PTY bytes for connection-negotiated runners", () => { + const socket = relaySocket(); + + const collision = encodeGitHubActionsRelayInput("looks-like-input", "terminal bytes"); + sendGitHubActionsRunnerOutput(socket, collision); + assert.deepEqual( + new Uint8Array(parseGitHubActionsRelayOutput(socket.sent[0]!)!), + new Uint8Array(collision), + ); +}); diff --git a/tests/github-actions-runtime.test.ts b/tests/github-actions-runtime.test.ts index 606ebe07..3fea9a1a 100644 --- a/tests/github-actions-runtime.test.ts +++ b/tests/github-actions-runtime.test.ts @@ -1,26 +1,57 @@ import assert from "node:assert/strict"; import { test } from "node:test"; import { + attachGitHubActionsRunnerProtocol, + attachGitHubActionsViewerProtocol, buildGitHubActionsRunnerPtyUrl, - forwardGitHubActionsRelayMessage, + buildGitHubActionsViewerRelayUrl, + createGitHubActionsRelayGeneration, + encodeGitHubActionsRelayInput, + encodeGitHubActionsRelayInputAcknowledgement, + encodeGitHubActionsRelayOutput, + gitHubActionsRelayGeneration, + gitHubActionsRelayUsesGenerations, gitHubActionsSessionStatus, githubActionsCapabilities, + githubActionsFramedRunnerCapability, + githubActionsGenerationFencedCapability, githubActionsRelayRole, + githubActionsRunnerProtocolHeader, + githubActionsRunnerProtocolQuery, githubActionsRuntimeLabel, + githubActionsViewerGenerationHeader, + githubActionsViewerProtocolHeader, + githubActionsViewerProtocolQuery, + gitHubActionsViewerResponseGeneration, + gitHubActionsViewerResponseUsesFramedProtocol, + gitHubActionsViewerResponseUsesGenerations, + gitHubActionsRunnerUsesFramedProtocol, + gitHubActionsViewerUsesFramedProtocol, isGitHubActionsViewerControlMessage, isTerminalGitHubActionsWorkState, notifyGitHubActionsViewers, + parseGitHubActionsRelayEvent, + parseGitHubActionsRelayInput, + parseGitHubActionsRelayInputAcknowledgement, + parseGitHubActionsRelayOutput, + parseGitHubActionsRunnerProtocol, + parseGitHubActionsRunnerProtocolOffer, + parseGitHubActionsViewerProtocol, parseGitHubActionsWorkState, + relayGitHubActionsWebSocketMessage, replaceGitHubActionsRunner, + sendGitHubActionsRelayInputAcknowledgement, type GitHubActionsRelaySocket, } from "../src/github-actions-runtime.ts"; function relaySocket(readyState = 1): GitHubActionsRelaySocket & { + attachment: unknown; closed: Array<[number | undefined, string | undefined]>; sent: Array; } { return { readyState, + attachment: null, sent: [], closed: [], send(message) { @@ -30,9 +61,27 @@ function relaySocket(readyState = 1): GitHubActionsRelaySocket & { this.closed.push([code, reason]); this.readyState = 3; }, + serializeAttachment(attachment) { + this.attachment = attachment; + }, + deserializeAttachment() { + return this.attachment; + }, }; } +function framedViewer() { + const viewer = relaySocket(); + attachGitHubActionsViewerProtocol(viewer, githubActionsFramedRunnerCapability); + return viewer; +} + +function generatedViewer() { + const viewer = relaySocket(); + attachGitHubActionsViewerProtocol(viewer, githubActionsGenerationFencedCapability); + return viewer; +} + test("github_actions exposes steerable terminal capabilities and label", () => { assert.equal(githubActionsRuntimeLabel("github_actions"), "GitHub Actions"); assert.equal(githubActionsRuntimeLabel("container"), ""); @@ -51,6 +100,69 @@ test("runner URL works without custom WebSocket headers", () => { buildGitHubActionsRunnerPtyUrl("https://crabfleet.openclaw.ai", "IS-123", "token with spaces"), "wss://crabfleet.openclaw.ai/api/agent/interactive-sessions/IS-123/runner-pty?agentToken=token+with+spaces", ); + assert.equal(parseGitHubActionsRunnerProtocol(null), null); + assert.equal( + parseGitHubActionsRunnerProtocol(githubActionsGenerationFencedCapability), + githubActionsGenerationFencedCapability, + ); + assert.equal( + parseGitHubActionsRunnerProtocol(githubActionsFramedRunnerCapability), + githubActionsFramedRunnerCapability, + ); + assert.equal(githubActionsRunnerProtocolHeader, "sec-websocket-protocol"); + assert.equal(parseGitHubActionsRunnerProtocolOffer(null), null); + assert.equal(parseGitHubActionsRunnerProtocolOffer(githubActionsFramedRunnerCapability), null); + assert.equal( + parseGitHubActionsRunnerProtocolOffer( + `other-protocol, ${githubActionsGenerationFencedCapability}`, + ), + githubActionsGenerationFencedCapability, + ); + assert.equal(githubActionsRunnerProtocolQuery, "runnerProtocol"); + assert.equal( + buildGitHubActionsViewerRelayUrl(), + "https://crabfleet.internal/api/session-control/github-actions/viewer?viewerProtocol=cfr1-framed-io-v2", + ); + assert.equal(parseGitHubActionsViewerProtocol(null), null); + assert.equal( + parseGitHubActionsViewerProtocol(githubActionsFramedRunnerCapability), + githubActionsFramedRunnerCapability, + ); + assert.equal( + parseGitHubActionsViewerProtocol(githubActionsGenerationFencedCapability), + githubActionsGenerationFencedCapability, + ); + assert.equal(githubActionsViewerProtocolQuery, "viewerProtocol"); + assert.equal(githubActionsViewerProtocolHeader, "x-crabfleet-viewer-protocol"); + assert.equal(githubActionsViewerGenerationHeader, "x-crabfleet-runner-generation"); + assert.equal( + gitHubActionsViewerResponseUsesFramedProtocol( + new Response(null, { + headers: { + [githubActionsViewerProtocolHeader]: githubActionsFramedRunnerCapability, + }, + }), + ), + true, + ); + assert.equal(gitHubActionsViewerResponseUsesFramedProtocol(new Response()), false); + const generatedResponse = new Response(null, { + headers: { + [githubActionsViewerGenerationHeader]: "generation-one", + [githubActionsViewerProtocolHeader]: githubActionsGenerationFencedCapability, + }, + }); + assert.equal(gitHubActionsViewerResponseUsesFramedProtocol(generatedResponse), true); + assert.equal(gitHubActionsViewerResponseUsesGenerations(generatedResponse), true); + assert.equal(gitHubActionsViewerResponseGeneration(generatedResponse), "generation-one"); + assert.equal( + gitHubActionsViewerResponseGeneration( + new Response(null, { + headers: { [githubActionsViewerGenerationHeader]: "invalid generation" }, + }), + ), + null, + ); }); test("work states preserve running phases and map terminal outcomes", () => { @@ -64,11 +176,11 @@ test("work states preserve running phases and map terminal outcomes", () => { assert.equal(gitHubActionsSessionStatus("failed"), "failed"); }); -test("relay replaces the current runner and routes messages by role", () => { +test("relay replaces the current runner and frames legacy raw runner output", () => { const oldRunner = relaySocket(); const runner = relaySocket(); - const viewerOne = relaySocket(); - const viewerTwo = relaySocket(); + const viewerOne = framedViewer(); + const viewerTwo = framedViewer(); assert.equal(replaceGitHubActionsRunner([oldRunner]), 1); assert.deepEqual(oldRunner.closed, [[1012, "runner replaced"]]); @@ -77,37 +189,403 @@ test("relay replaces the current runner and routes messages by role", () => { assert.deepEqual(stoppedRunner.closed, [[1000, "runner disconnected"]]); assert.equal( - forwardGitHubActionsRelayMessage("runner", "output", [runner], [viewerOne, viewerTwo]), + relayGitHubActionsWebSocketMessage( + "runner", + runner, + "output", + [runner], + [viewerOne, viewerTwo], + ), 2, ); - assert.deepEqual(viewerOne.sent, ["output"]); - assert.deepEqual(viewerTwo.sent, ["output"]); + assert.equal( + new TextDecoder().decode(parseGitHubActionsRelayOutput(viewerOne.sent[0]!)!), + "output", + ); + assert.equal( + new TextDecoder().decode(parseGitHubActionsRelayOutput(viewerTwo.sent[0]!)!), + "output", + ); + + const oldCapabilityMessage = + '{"type":"crabfleet_runner_capabilities","capabilities":["cfr1-framed-io-v1"]}'; + assert.equal( + relayGitHubActionsWebSocketMessage( + "runner", + runner, + oldCapabilityMessage, + [runner], + [viewerOne], + ), + 1, + ); + assert.equal( + new TextDecoder().decode(parseGitHubActionsRelayOutput(viewerOne.sent[1]!)!), + oldCapabilityMessage, + ); + assert.equal(gitHubActionsRunnerUsesFramedProtocol(runner), false); +}); + +test("legacy runners receive raw input and the relay acknowledges delivery", () => { + const runner = relaySocket(); + const viewer = framedViewer(); + const input = encodeGitHubActionsRelayInput("input-legacy", "steer"); + + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, input, [runner], []), 1); + assert.equal(new TextDecoder().decode(runner.sent[0] as ArrayBuffer), "steer"); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(viewer.sent[0]!), { + inputId: "input-legacy", + accepted: true, + }); +}); + +test("connection-time opt-in frames the first input without a pending handshake", () => { + const closedRunner = relaySocket(3); + const openRunner = relaySocket(); + const laterRunner = relaySocket(); + const viewer = framedViewer(); + const input = encodeGitHubActionsRelayInput("input-one", "steer"); + + attachGitHubActionsRunnerProtocol(openRunner, githubActionsFramedRunnerCapability); + assert.equal(gitHubActionsRunnerUsesFramedProtocol(openRunner), true); + assert.equal( + relayGitHubActionsWebSocketMessage( + "viewer", + viewer, + input, + [closedRunner, openRunner, laterRunner], + [], + ), + 1, + ); + assert.deepEqual(closedRunner.sent, []); + assert.deepEqual(openRunner.sent, [input]); + assert.deepEqual(laterRunner.sent, []); + assert.deepEqual(viewer.sent, []); + assert.deepEqual(parseGitHubActionsRelayInput(openRunner.sent[0]!), { + inputId: "input-one", + payload: new TextEncoder().encode("steer").buffer, + }); +}); + +test("relay rejects framed input only when no runner accepts the frame", () => { + const runner = relaySocket(); + runner.send = () => { + throw new Error("runner disconnected"); + }; + const viewer = framedViewer(); + const input = encodeGitHubActionsRelayInput("input-failed", "steer"); + + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, input, [runner], []), 0); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(viewer.sent[0]!), { + inputId: "input-failed", + accepted: false, + error: "GitHub Actions runner did not accept terminal input", + }); + + const waitingViewer = framedViewer(); + assert.equal(relayGitHubActionsWebSocketMessage("viewer", waitingViewer, input, [], []), 0); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(waitingViewer.sent[0]!), { + inputId: "input-failed", + accepted: false, + error: "GitHub Actions runner did not accept terminal input", + }); +}); + +test("runner acknowledgements retain correlation and fan out to viewers", () => { + const runner = relaySocket(); + const viewerOne = framedViewer(); + const viewerTwo = framedViewer(); + const acknowledgement = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "input-two", + accepted: true, + }); + attachGitHubActionsRunnerProtocol(runner, githubActionsFramedRunnerCapability); + + assert.equal( + relayGitHubActionsWebSocketMessage( + "runner", + runner, + acknowledgement, + [runner], + [viewerOne, viewerTwo], + ), + 2, + ); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(viewerOne.sent[0]!), { + inputId: "input-two", + accepted: true, + }); + assert.deepEqual(viewerTwo.sent, [acknowledgement]); +}); + +test("relay-owned generations fence stale input and bridge framed protocol versions", () => { + const runner = relaySocket(); + const viewer = generatedViewer(); + const generation = createGitHubActionsRelayGeneration(); + attachGitHubActionsRunnerProtocol(runner, githubActionsFramedRunnerCapability, generation); + + assert.equal(gitHubActionsRelayGeneration(runner), generation); + assert.equal(gitHubActionsRelayUsesGenerations(runner), false); + assert.equal(gitHubActionsRelayUsesGenerations(viewer), true); + + const staleInput = encodeGitHubActionsRelayInput("input-stale", "old", "old-generation"); + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, staleInput, [runner], []), 0); + assert.deepEqual(runner.sent, []); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(viewer.sent[0]!), { + inputId: "input-stale", + accepted: false, + error: "GitHub Actions runner did not accept terminal input", + generation: "old-generation", + }); + + const currentInput = encodeGitHubActionsRelayInput("input-current", "new", generation); + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, currentInput, [runner], []), 1); + assert.deepEqual(parseGitHubActionsRelayInput(runner.sent[0]!), { + inputId: "input-current", + payload: new TextEncoder().encode("new").buffer, + }); + + const acknowledgement = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "input-current", + accepted: true, + }); + assert.equal( + relayGitHubActionsWebSocketMessage("runner", runner, acknowledgement, [runner], [viewer]), + 1, + ); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(viewer.sent[1]!), { + inputId: "input-current", + accepted: true, + generation, + }); +}); + +test("generation acknowledgements bridge to exact v1 frames for v1 viewers", () => { + const runner = relaySocket(); + const viewer = framedViewer(); + attachGitHubActionsRunnerProtocol( + runner, + githubActionsGenerationFencedCapability, + "current-generation", + ); + const acknowledgement = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "input-current", + accepted: false, + error: "rejected", + generation: "current-generation", + }); assert.equal( - forwardGitHubActionsRelayMessage("viewer", "input", [runner], [viewerOne, viewerTwo]), + relayGitHubActionsWebSocketMessage("runner", runner, acknowledgement, [runner], [viewer]), 1, ); - assert.deepEqual(runner.sent, ["input"]); + const expected = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "input-current", + accepted: false, + error: "rejected", + }); + assert.deepEqual(new Uint8Array(viewer.sent[0] as ArrayBuffer), new Uint8Array(expected)); + assert.deepEqual(parseGitHubActionsRelayInputAcknowledgement(viewer.sent[0]!), { + inputId: "input-current", + accepted: false, + error: "rejected", + }); +}); + +test("replacement relay drops acknowledgements from the superseded runner", () => { + const oldRunner = relaySocket(); + const replacement = relaySocket(); + const viewer = generatedViewer(); + attachGitHubActionsRunnerProtocol( + oldRunner, + githubActionsGenerationFencedCapability, + "old-generation", + ); + attachGitHubActionsRunnerProtocol( + replacement, + githubActionsGenerationFencedCapability, + "new-generation", + ); + const acknowledgement = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "input-old", + accepted: true, + generation: "old-generation", + }); + + assert.equal( + relayGitHubActionsWebSocketMessage( + "runner", + oldRunner, + acknowledgement, + [replacement], + [viewer], + ), + 0, + ); + assert.deepEqual(viewer.sent, []); }); -test("relay consumes viewer resize controls without corrupting raw runner input", () => { +test("generation-fenced runners echo only their relay-owned generation", () => { const runner = relaySocket(); + const viewer = generatedViewer(); + attachGitHubActionsRunnerProtocol( + runner, + githubActionsGenerationFencedCapability, + "current-generation", + ); + const input = encodeGitHubActionsRelayInput("input-current", "steer", "current-generation"); + + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, input, [runner], []), 1); + assert.deepEqual(parseGitHubActionsRelayInput(runner.sent[0]!), { + inputId: "input-current", + generation: "current-generation", + payload: new TextEncoder().encode("steer").buffer, + }); + + const staleAcknowledgement = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "input-current", + accepted: true, + generation: "stale-generation", + }); + assert.equal( + relayGitHubActionsWebSocketMessage("runner", runner, staleAcknowledgement, [runner], [viewer]), + 0, + ); + assert.deepEqual(viewer.sent, []); +}); + +test("negotiated runners frame output so control-shaped terminal bytes stay output", () => { + const runner = relaySocket(); + const viewer = framedViewer(); + const controlShapedOutput = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "collision", + accepted: true, + }); + + attachGitHubActionsRunnerProtocol(runner, githubActionsFramedRunnerCapability); + const output = encodeGitHubActionsRelayOutput(controlShapedOutput); + assert.equal(relayGitHubActionsWebSocketMessage("runner", runner, output, [runner], [viewer]), 1); + assert.deepEqual( + new Uint8Array(parseGitHubActionsRelayOutput(viewer.sent[0]!)!), + new Uint8Array(controlShapedOutput), + ); +}); + +test("typed acknowledgements reject malformed ids and preserve collision-shaped terminal text", () => { + const viewer = framedViewer(); + const runner = relaySocket(); + const collision = '{"type":"github_actions_input_ack","inputId":"input-three","accepted":true}'; + + assert.equal(parseGitHubActionsRelayInputAcknowledgement(collision), null); + assert.equal( + relayGitHubActionsWebSocketMessage("runner", runner, collision, [runner], [viewer]), + 1, + ); + assert.equal( + new TextDecoder().decode(parseGitHubActionsRelayOutput(viewer.sent[0]!)!), + collision, + ); + assert.throws(() => encodeGitHubActionsRelayInput("bad id", "input"), { + message: "invalid GitHub Actions relay input id", + }); + assert.equal( + sendGitHubActionsRelayInputAcknowledgement(relaySocket(3), { + inputId: "input-three", + accepted: false, + }), + false, + ); +}); + +test("relay consumes viewer resize controls and rejects unframed input", () => { + const runner = relaySocket(); + const viewer = framedViewer(); const resize = JSON.stringify({ type: "resize", cols: 120, rows: 40 }); const typedJson = new TextEncoder().encode(resize).buffer; assert.equal(isGitHubActionsViewerControlMessage(resize), true); assert.equal(isGitHubActionsViewerControlMessage(typedJson), false); - assert.equal(forwardGitHubActionsRelayMessage("viewer", resize, [runner], []), 0); + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, resize, [runner], []), 0); + assert.deepEqual(runner.sent, []); + assert.deepEqual(viewer.sent, []); + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, typedJson, [runner], []), 0); assert.deepEqual(runner.sent, []); - assert.equal(forwardGitHubActionsRelayMessage("viewer", typedJson, [runner], []), 1); - assert.deepEqual(runner.sent, [typedJson]); }); test("relay tags and runner lifecycle notifications stay explicit", () => { - const viewer = relaySocket(); + const viewer = framedViewer(); assert.equal(githubActionsRelayRole(["github-actions-runner"]), "runner"); assert.equal(githubActionsRelayRole(["github-actions-viewer"]), "viewer"); assert.equal(githubActionsRelayRole([]), null); assert.equal(notifyGitHubActionsViewers([viewer], "runner_waiting"), 1); - assert.deepEqual(viewer.sent, ['{"type":"runner_waiting"}']); + assert.deepEqual(parseGitHubActionsRelayEvent(viewer.sent[0]!), { + type: "runner_waiting", + }); +}); + +test("generation-fenced lifecycle events identify the relay runner", () => { + const viewer = generatedViewer(); + assert.equal(notifyGitHubActionsViewers([viewer], "runner_connected", "generation-one"), 1); + assert.deepEqual(parseGitHubActionsRelayEvent(viewer.sent[0]!), { + type: "runner_connected", + generation: "generation-one", + }); + assert.equal(notifyGitHubActionsViewers([viewer], "runner_waiting"), 1); + assert.deepEqual(parseGitHubActionsRelayEvent(viewer.sent[1]!), { + type: "runner_waiting", + generation: "none", + }); +}); + +test("unnegotiated viewers retain raw relay compatibility", () => { + const runner = relaySocket(); + const viewer = relaySocket(); + + assert.equal(gitHubActionsViewerUsesFramedProtocol(viewer), false); + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, "steer", [runner], []), 1); + assert.deepEqual(runner.sent, ["steer"]); + assert.deepEqual(JSON.parse(viewer.sent[0] as string), { + type: "github_actions_input_ack", + accepted: true, + }); + + assert.equal( + relayGitHubActionsWebSocketMessage("runner", runner, "output", [runner], [viewer]), + 1, + ); + assert.equal(viewer.sent[1], "output"); + assert.equal(notifyGitHubActionsViewers([viewer], "runner_disconnected"), 1); + assert.deepEqual(JSON.parse(viewer.sent[2] as string), { + type: "runner_disconnected", + }); +}); + +test("raw viewers bridge through framed runners without receiving CFR1 controls", () => { + const runner = relaySocket(); + const viewer = relaySocket(); + attachGitHubActionsRunnerProtocol(runner, githubActionsFramedRunnerCapability); + + assert.equal(relayGitHubActionsWebSocketMessage("viewer", viewer, "steer", [runner], []), 1); + const input = parseGitHubActionsRelayInput(runner.sent[0]!); + assert.ok(input); + assert.equal(new TextDecoder().decode(input.payload), "steer"); + assert.deepEqual(JSON.parse(viewer.sent[0] as string), { + type: "github_actions_input_ack", + accepted: true, + }); + + const acknowledgement = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: input.inputId, + accepted: true, + }); + assert.equal( + relayGitHubActionsWebSocketMessage("runner", runner, acknowledgement, [runner], [viewer]), + 0, + ); + assert.equal(viewer.sent.length, 1); + + const output = encodeGitHubActionsRelayOutput("output"); + assert.equal(relayGitHubActionsWebSocketMessage("runner", runner, output, [runner], [viewer]), 1); + assert.equal(new TextDecoder().decode(viewer.sent[1] as ArrayBuffer), "output"); }); diff --git a/tests/github-actions-session-registration.test.ts b/tests/github-actions-session-registration.test.ts index ce181abc..481f4d19 100644 --- a/tests/github-actions-session-registration.test.ts +++ b/tests/github-actions-session-registration.test.ts @@ -8,6 +8,7 @@ import { actionWorkIdentifier, buildGitHubActionsSessionValues, optionalHttpUrl, + type GitHubActionsSessionRegistrationExpectation, type GitHubActionsSessionRegistrationStore, type GitHubActionsSessionRegistrationUpdate, } from "../src/worker/github-actions-session-registration.ts"; @@ -18,7 +19,11 @@ type StoreState = { rows: Map; workKeyReads: number; inserted: InteractiveSessionTable[]; - updates: Array<{ id: string; values: GitHubActionsSessionRegistrationUpdate }>; + updates: Array<{ + id: string; + values: GitHubActionsSessionRegistrationUpdate; + expected: GitHubActionsSessionRegistrationExpectation; + }>; events: string[]; audits: string[]; operations: string[]; @@ -73,9 +78,9 @@ function registrationStore(initialRows: InteractiveSessionRow[] = []): { state.rows.set(row.id, row); }, readById: async (id) => state.rows.get(id) ?? null, - updateSession: async (id, values) => { + updateSession: async (id, values, expected) => { state.operations.push("update"); - state.updates.push({ id, values }); + state.updates.push({ id, values, expected }); const row = state.rows.get(id); if (row) state.rows.set(id, { ...row, ...values }); }, @@ -260,6 +265,13 @@ test("GitHub Actions work keys can be resumed by the matching owner", async () = assert.equal(state.updates[0]?.values.owner, "operator"); assert.equal(state.updates[0]?.values.owner_subject, "github:42"); assert.equal(state.updates[0]?.values.agent_token_hash, "agent-token-hash"); + assert.deepEqual(state.updates[0]?.expected, { + agent_token_hash: existing.agent_token_hash, + updated_at: existing.updated_at, + status: existing.status, + work_state: existing.work_state, + work_phase: existing.work_phase, + }); }); test("GitHub Actions rejects work keys without a stable owner", async () => { @@ -352,6 +364,95 @@ test("registration adopts a concurrently inserted work key", async () => { assert.equal(result.session.id, "IS-concurrent"); assert.equal(state.workKeyReads, 2); assert.equal(state.updates[0]?.id, "IS-concurrent"); + assert.equal( + state.updates[0]?.values.updated_at, + Math.max(state.concurrentRow.updated_at + 1, 100), + ); + assert.equal(state.updates[0]?.expected.agent_token_hash, state.concurrentRow.agent_token_hash); +}); + +test("concurrent registration adoption rotates exactly one usable token", async () => { + const existing = sessionRow({ + id: "IS-concurrent", + runtime: "github_actions", + work_key: "issue:race-cas", + owner: "operator", + owner_subject: "github:42", + updated_at: 100, + }); + const { store, state } = registrationStore([existing]); + let tokenSequence = 0; + let arrivals = 0; + let releaseUpdates!: () => void; + const updatesReady = new Promise((resolve) => { + releaseUpdates = resolve; + }); + store.newAgentToken = () => `agent-token-${++tokenSequence}`; + store.hashToken = async (token) => `${token}-hash`; + store.updateSession = async (id, values, expected) => { + arrivals += 1; + if (arrivals === 2) releaseUpdates(); + await updatesReady; + const current = state.rows.get(id); + if ( + !current || + current.updated_at !== expected.updated_at || + current.agent_token_hash !== expected.agent_token_hash || + current.status !== expected.status || + current.work_state !== expected.work_state || + current.work_phase !== expected.work_phase + ) { + throw new Error("GitHub Actions session changed; retry"); + } + state.rows.set(id, { ...current, ...values }); + }; + + const input = { + workKey: "issue:race-cas", + workKind: "issue", + repo: "openclaw/crabfleet", + owner: "operator@example.test", + }; + const results = await Promise.allSettled([ + new GitHubActionsSessionRegistrationService(store).register(input), + new GitHubActionsSessionRegistrationService(store).register(input), + ]); + + const fulfilled = results.filter( + ( + result, + ): result is PromiseFulfilledResult< + Awaited> + > => result.status === "fulfilled", + ); + assert.equal(fulfilled.length, 1); + assert.equal(results.filter((result) => result.status === "rejected").length, 1); + assert.equal( + state.rows.get(existing.id)?.agent_token_hash, + `${fulfilled[0]?.value.agentToken}-hash`, + ); + assert.equal(state.rows.get(existing.id)?.updated_at, 101); +}); + +test("registration revisions advance monotonically when the stored clock is ahead", async () => { + const existing = sessionRow({ + id: "IS-future-revision", + runtime: "github_actions", + work_key: "issue:future-revision", + owner: "operator", + owner_subject: "github:42", + updated_at: 500, + }); + const { store, state } = registrationStore([existing]); + + await new GitHubActionsSessionRegistrationService(store).register({ + workKey: "issue:future-revision", + workKind: "issue", + repo: "openclaw/crabfleet", + owner: "operator@example.test", + }); + + assert.equal(state.rows.get(existing.id)?.updated_at, 501); }); test("registration rejects invalid input and work keys owned by another runtime", async () => { diff --git a/tests/github-actions-session-work-state.test.ts b/tests/github-actions-session-work-state.test.ts index 11a69d80..4310d639 100644 --- a/tests/github-actions-session-work-state.test.ts +++ b/tests/github-actions-session-work-state.test.ts @@ -13,6 +13,8 @@ import { sessionRow } from "./helpers/session-row.ts"; type WorkStateStoreState = { row: InteractiveSessionRow | null; update: GitHubActionsWorkStateUpdate | null; + expectedRevision: number | undefined; + expectedTerminalStatus: InteractiveSessionRow["status"] | undefined; events: string[]; operations: string[]; disconnectError: unknown; @@ -48,6 +50,8 @@ function workStateStore(values: Partial = {}): { ...values, }), update: null, + expectedRevision: undefined, + expectedTerminalStatus: undefined, events: [], operations: [], disconnectError: null, @@ -55,9 +59,11 @@ function workStateStore(values: Partial = {}): { const store: GitHubActionsWorkStateStore = { now: () => 500, readRow: async () => state.row, - persist: async (_id, update) => { + persist: async (_id, update, expectedRevision, expectedTerminalStatus) => { state.operations.push("persist"); state.update = update; + state.expectedRevision = expectedRevision; + state.expectedTerminalStatus = expectedTerminalStatus; if (state.row) state.row = { ...state.row, ...update }; }, appendEvent: async (_id, message) => { @@ -105,6 +111,7 @@ test("active work-state updates project fields and clear stale completion", asyn stopped_at: null, }); assert.deepEqual(state.events, ["running: codex_turn"]); + assert.equal(state.expectedRevision, workSession().updatedAt); assert.deepEqual(state.operations, ["persist", "event", "read"]); }); @@ -139,10 +146,57 @@ test("terminal work-state updates stop the session and disconnect the runner", a assert.equal(result.status, "failed"); assert.equal(state.update?.completion_reason, "existing reason"); assert.equal(state.update?.stopped_at, 500); + assert.equal(state.expectedTerminalStatus, "ready"); + assert.equal(state.expectedRevision, workSession().updatedAt); assert.deepEqual(state.events, ["failed: tests"]); assert.deepEqual(state.operations, ["persist", "event", "disconnect", "read"]); }); +test("terminal work-state updates carry the status observed before persistence", async () => { + const { store, state } = workStateStore({ + status: "attached", + work_state: "running", + }); + + await new GitHubActionsWorkStateService(store).update(workSession(), { + state: "completed", + }); + + assert.equal(state.expectedTerminalStatus, "attached"); +}); + +test("work-state updates retain the exact revision authenticated before a token rotation", async () => { + const authenticated = workSession({ updated_at: 400 }); + const { store, state } = workStateStore({ updated_at: 401, work_state: "registered" }); + store.persist = async (_id, _update, expectedRevision) => { + state.expectedRevision = expectedRevision; + if (state.row?.updated_at !== expectedRevision) { + throw new Error("GitHub Actions session changed; retry"); + } + }; + + await assert.rejects( + new GitHubActionsWorkStateService(store).update(authenticated, { + state: "running", + }), + { message: "GitHub Actions session changed; retry" }, + ); + + assert.equal(state.expectedRevision, 400); +}); + +test("work-state updates advance revisions when the authenticated clock is ahead", async () => { + const authenticated = workSession({ updated_at: 800 }); + const { store, state } = workStateStore({ updated_at: 800 }); + + await new GitHubActionsWorkStateService(store).update(authenticated, { + state: "running", + }); + + assert.equal(state.expectedRevision, 800); + assert.equal(state.update?.updated_at, 801); +}); + test("terminal runner disconnect races remain best effort", async () => { const { store, state } = workStateStore(); state.disconnectError = new Error("runner already disconnected"); diff --git a/tests/helpers/generated-assets.ts b/tests/helpers/generated-assets.ts new file mode 100644 index 00000000..cce15c67 --- /dev/null +++ b/tests/helpers/generated-assets.ts @@ -0,0 +1,65 @@ +import { execFile } from "node:child_process"; +import { createHash } from "node:crypto"; +import { createServer, type Server } from "node:net"; +import { setTimeout as delay } from "node:timers/promises"; +import { fileURLToPath } from "node:url"; +import { promisify } from "node:util"; + +const execFileAsync = promisify(execFile); +const lockWaitMs = 120_000; +const lockPortBase = 49_152; +const lockPortRange = 65_535 - lockPortBase + 1; +const lockPort = + lockPortBase + + (createHash("sha256") + .update(fileURLToPath(new URL("../../", import.meta.url))) + .digest() + .readUInt16BE(0) % + lockPortRange); + +export async function withGeneratedAssetsForTest(consume: () => T | Promise): Promise { + const lock = await acquireLock(); + try { + await execFileAsync(process.execPath, ["scripts/generate-assets.mjs"]); + return await consume(); + } finally { + await closeServer(lock); + } +} + +async function acquireLock(): Promise { + const deadline = Date.now() + lockWaitMs; + while (true) { + const server = createServer(); + try { + await listen(server); + return server; + } catch (error) { + if (!isAddressInUse(error)) throw error; + if (Date.now() >= deadline) + throw new Error("timed out waiting for generated asset test lock"); + await delay(50); + } + } +} + +function listen(server: Server): Promise { + return new Promise((resolve, reject) => { + const onError = (error: Error) => reject(error); + server.once("error", onError); + server.listen({ host: "127.0.0.1", port: lockPort, exclusive: true }, () => { + server.off("error", onError); + resolve(); + }); + }); +} + +function closeServer(server: Server): Promise { + return new Promise((resolve, reject) => { + server.close((error) => (error ? reject(error) : resolve())); + }); +} + +function isAddressInUse(error: unknown): boolean { + return error instanceof Error && "code" in error && error.code === "EADDRINUSE"; +} diff --git a/tests/http.test.ts b/tests/http.test.ts index 8540ecc0..5034d44d 100644 --- a/tests/http.test.ts +++ b/tests/http.test.ts @@ -67,6 +67,25 @@ test("JSON parsing and status errors retain stable messages and status codes", a "status" in error && error.status === 400, ); + const rejectedBody = new ReadableStream({ + start(controller) { + controller.error(new Error("request body aborted")); + }, + }); + await assert.rejects( + readJson( + new Request("https://fleet.example", { + method: "POST", + body: rejectedBody, + duplex: "half", + } as RequestInit & { duplex: "half" }), + ), + (error: unknown) => + error instanceof Error && + error.message === "invalid json" && + "status" in error && + error.status === 400, + ); for (const [error, status, message] of [ [unauthorized(), 401, "unauthorized"], @@ -113,6 +132,59 @@ test("bounded JSON parsing rejects declared and streamed bodies before unbounded } }); +test("JSON parsing rejects integers that cannot round-trip exactly", async () => { + for (const body of [ + '{"value":9007199254740993}', + '{"value":9007199254740991.1}', + '{"value":1.0000000000000001}', + '{"value":-0}', + '{"nested":[1e400]}', + ]) { + for (const parse of [ + () => readJson(new Request("https://fleet.example", { method: "POST", body })), + () => + readBoundedJson( + new Request("https://fleet.example", { method: "POST", body }), + body.length + 1, + ), + ]) { + await assert.rejects(parse(), (error: unknown) => { + assert.equal( + typeof error === "object" && error && "status" in error ? error.status : undefined, + 400, + ); + assert.match(error instanceof Error ? error.message : "", /round-trippable/); + return true; + }); + } + } +}); + +test("JSON parsing accepts exact integer-equivalent numeric forms", async () => { + for (const body of ['{"value":1.0}', '{"value":1e0}', '{"value":100e-2}']) { + assert.deepEqual( + await readJson<{ value: number }>( + new Request("https://fleet.example", { method: "POST", body }), + ), + { value: 1 }, + ); + } +}); + +test("JSON parsing handles deeply nested bounded payloads without exhausting the call stack", async () => { + const depth = 20_000; + const body = `${"[".repeat(depth)}0${"]".repeat(depth)}`; + let current = await readBoundedJson( + new Request("https://fleet.example", { method: "POST", body }), + body.length, + ); + for (let index = 0; index < depth; index += 1) { + assert.ok(Array.isArray(current)); + current = current[0]; + } + assert.equal(current, 0); +}); + test("bearer and cookie helpers normalize only their owned protocol surface", () => { assert.equal( bearerToken( diff --git a/tests/interactive-terminal-service.test.ts b/tests/interactive-terminal-service.test.ts index 753f31b8..43a6d4ad 100644 --- a/tests/interactive-terminal-service.test.ts +++ b/tests/interactive-terminal-service.test.ts @@ -5,6 +5,7 @@ import { readTerminalClipboardBytes, terminalClipboardFilename, } from "../src/worker/interactive-terminal.ts"; +import { TerminalInputStateRegistry } from "../src/worker/interactive-terminal-service.ts"; test("terminal clipboard filenames are bounded, sanitized, and typed", () => { assert.equal(terminalClipboardFilename("screen shot", "image/png"), "screen-shot.png"); @@ -40,3 +41,16 @@ test("terminal clipboard upload rejects empty and declared oversized bodies", as new Uint8Array([97, 98, 99]), ); }); + +test("terminal input state survives until the final subscriber releases it", () => { + const states = new TerminalInputStateRegistry(); + states.retain("IS-1"); + states.retain("IS-1"); + states.state("IS-1").line = "shared input"; + + states.release("IS-1"); + assert.equal(states.state("IS-1").line, "shared input"); + + states.release("IS-1"); + assert.equal(states.state("IS-1").line, ""); +}); diff --git a/tests/openclaw-repository.test.ts b/tests/openclaw-repository.test.ts index 7ed39176..700c25d9 100644 --- a/tests/openclaw-repository.test.ts +++ b/tests/openclaw-repository.test.ts @@ -30,6 +30,7 @@ type PreparedStatement = { function runtimeEnv( handler: D1Handler, batchHandler: (statements: PreparedStatement[]) => void = () => undefined, + sessionLogs?: Pick, ): RuntimeEnv { return { DB: { @@ -57,6 +58,7 @@ function runtimeEnv( return []; }, } as unknown as D1Database, + ...(sessionLogs ? { SESSION_LOGS: sessionLogs as R2Bucket } : {}), } as RuntimeEnv; } @@ -268,23 +270,68 @@ test("OpenClaw stale reservation reads are bounded and map persistence names", a test("OpenClaw reservation rollback deletes all owned records in one batch", async () => { let batch: PreparedStatement[] = []; + const deletedKeys: string[] = []; const removed = await removeInteractiveSessionReservation( runtimeEnv( (sql, parameters, kind) => { + if (/^update "interactive_sessions"/i.test(sql)) { + assert.equal(kind, "run"); + assert.match(sql, /set "reconcile_error" = .+, "updated_at" =/i); + assert.match(sql, /"updated_at" =/i); + assert.deepEqual(parameters, [ + "reservation-rollback:100", + 101, + "IS-2", + "provisioning", + 1, + 100, + 100, + ]); + return { changes: 1 }; + } assert.equal(kind, "all"); + if (/from "interactive_session_log_archives"/i.test(sql)) { + assert.deepEqual(parameters, ["IS-2"]); + return { + results: [ + { + events_key: "events", + transcript_key: "transcript", + summary_key: "summary", + }, + ], + }; + } assert.match(sql, /^select "id" from "interactive_sessions"/i); + if (parameters.length > 1) { + assert.deepEqual(parameters, [ + "IS-2", + "provisioning", + 1, + 100, + 101, + "reservation-rollback:100", + ]); + return { results: [{ id: "IS-2" }] }; + } assert.deepEqual(parameters, ["IS-2"]); return { results: [] }; }, (statements) => { batch = statements; }, + { + async delete(key) { + deletedKeys.push(key); + }, + }, ), "IS-2", 100, ); assert.equal(removed, true); + assert.deepEqual(deletedKeys.sort(), ["events", "summary", "transcript"]); assert.equal(batch.length, 4); assert.match(batch[0]?.sql ?? "", /^delete from "openclaw_request_replays"/i); assert.match(batch[1]?.sql ?? "", /^delete from "interactive_session_events"/i); @@ -292,6 +339,92 @@ test("OpenClaw reservation rollback deletes all owned records in one batch", asy assert.match(batch[3]?.sql ?? "", /^delete from "interactive_sessions"/i); assert.ok(batch.every((statement) => statement.parameters.includes("IS-2"))); assert.ok(batch.every((statement) => statement.parameters.includes(100))); + assert.ok(batch.every((statement) => statement.parameters.includes(101))); + assert.ok(batch.every((statement) => statement.parameters.includes("reservation-rollback:100"))); +}); + +test("OpenClaw reservation rollback retains its durable claim when archive cleanup fails", async () => { + let batched = false; + await assert.rejects( + removeInteractiveSessionReservation( + runtimeEnv( + (sql, _parameters, kind) => { + if (/^update "interactive_sessions"/i.test(sql)) { + assert.equal(kind, "run"); + return { changes: 1 }; + } + assert.equal(kind, "all"); + if (/from "interactive_session_log_archives"/i.test(sql)) { + return { + results: [ + { + events_key: "events", + transcript_key: "transcript", + summary_key: "summary", + }, + ], + }; + } + return { results: [{ id: "IS-2" }] }; + }, + () => { + batched = true; + }, + { + async delete() { + throw new Error("R2 unavailable"); + }, + }, + ), + "IS-2", + 100, + ), + /R2 unavailable/, + ); + assert.equal(batched, false); +}); + +test("OpenClaw reservation rollback resumes an existing archive cleanup claim", async () => { + let batched = false; + const deletedKeys: string[] = []; + const removed = await removeInteractiveSessionReservation( + runtimeEnv( + (sql, _parameters, kind) => { + if (/^update "interactive_sessions"/i.test(sql)) { + assert.equal(kind, "run"); + return { changes: 0 }; + } + assert.equal(kind, "all"); + if (/from "interactive_session_log_archives"/i.test(sql)) { + return { + results: [ + { + events_key: "events", + transcript_key: "transcript", + summary_key: "summary", + }, + ], + }; + } + if (/reconcile_error/i.test(sql)) return { results: [{ id: "IS-2" }] }; + return { results: [] }; + }, + () => { + batched = true; + }, + { + async delete(key) { + deletedKeys.push(key); + }, + }, + ), + "IS-2", + 100, + ); + + assert.equal(removed, true); + assert.equal(batched, true); + assert.deepEqual(deletedKeys.sort(), ["events", "summary", "transcript"]); }); test("OpenClaw reservation activation reports the fenced compare-and-set result", async () => { diff --git a/tests/runtime-adapter-release-service.test.ts b/tests/runtime-adapter-release-service.test.ts index 2b7fefa8..d00cf163 100644 --- a/tests/runtime-adapter-release-service.test.ts +++ b/tests/runtime-adapter-release-service.test.ts @@ -1,7 +1,17 @@ import assert from "node:assert/strict"; +import { readFileSync } from "node:fs"; +import { DatabaseSync } from "node:sqlite"; import test from "node:test"; import type { RuntimeEnv } from "../src/worker/env.ts"; +import { + claimRuntimeAdapterWorkspaceCleanup, + claimRuntimeAdapterWorkspaceCleanupBatch, + completeRuntimeAdapterWorkspaceCleanup, + markRuntimeAdapterWorkspaceCleanupDeletionObserved, + persistRuntimeAdapterWorkspaceCleanupEvidence, + stageRuntimeAdapterWorkspaceCleanup, +} from "../src/worker/provisioning/runtime-adapter-release-repository.ts"; import { clearRuntimeAdapterCreatePending, confirmRuntimeAdapterRelease, @@ -10,8 +20,15 @@ import { import { RuntimeAdapterReleaseService, type RuntimeAdapterReleaseServiceDependencies, + type RuntimeAdapterWorkspaceCleanup, + type RuntimeAdapterWorkspaceRegistration, } from "../src/worker/provisioning/runtime-adapter-release-service.ts"; +const registration: RuntimeAdapterWorkspaceRegistration = { + profile: "default", + controlPlane: "https://adapter.example.test/", +}; + type PreparedStatement = { sql: string; parameters: unknown[]; @@ -22,7 +39,27 @@ type PreparedStatement = { function releaseDependencies( overrides: Partial = {}, ): RuntimeAdapterReleaseServiceDependencies { + let stagedCleanup: RuntimeAdapterWorkspaceCleanup | null = null; return { + async stageCleanup(input) { + stagedCleanup = { + sessionId: input.sessionId, + adapterWorkspaceId: input.adapterWorkspaceId, + registration: input.registration, + createPending: input.createPending, + deletionObserved: false, + claim: "claim-1", + }; + }, + async claimCleanup() { + return stagedCleanup; + }, + async claimPendingCleanups() { + return []; + }, + async persistCleanupEvidence() {}, + async markCleanupDeletionObserved() {}, + async completeCleanup() {}, async clearCreatePending() {}, async stopWorkspace() { return { status: "stopped", message: "runtime workspace released" }; @@ -72,31 +109,52 @@ test("superseded release clears the create marker before stopping and confirming const calls: string[] = []; const service = new RuntimeAdapterReleaseService( releaseDependencies({ + async stageCleanup(input) { + calls.push(`stage:${input.sessionId}:${input.adapterWorkspaceId}`); + }, + async claimCleanup(sessionId, adapterWorkspaceId) { + return { + sessionId, + adapterWorkspaceId, + registration, + createPending: false, + deletionObserved: false, + claim: "claim-1", + }; + }, async clearCreatePending(sessionId, adapterWorkspaceId) { calls.push(`clear:${sessionId}:${adapterWorkspaceId}`); }, - async stopWorkspace(sessionId, adapterWorkspaceId) { - calls.push(`stop:${sessionId}:${adapterWorkspaceId}`); + async stopWorkspace(sessionId, adapterWorkspaceId, retained, createPending) { + calls.push( + `stop:${sessionId}:${adapterWorkspaceId}:${retained?.profile}:${retained?.controlPlane}:${createPending}`, + ); return { status: "stopped", message: "runtime workspace released" }; }, async confirmRelease(sessionId, adapterWorkspaceId, now, message) { calls.push(`confirm:${sessionId}:${adapterWorkspaceId}:${now}:${message}`); return "stopped"; }, + async completeCleanup(cleanup) { + calls.push(`complete:${cleanup.sessionId}:${cleanup.adapterWorkspaceId}`); + }, }), ); await service.stopSuperseded({ sessionId: "IS-101", adapterWorkspaceId: "fleet-a-is-101", + registration, createPending: false, now: 200, }); assert.deepEqual(calls, [ + "stage:IS-101:fleet-a-is-101", "clear:IS-101:fleet-a-is-101", - "stop:IS-101:fleet-a-is-101", + "stop:IS-101:fleet-a-is-101:default:https://adapter.example.test/:false", "confirm:IS-101:fleet-a-is-101:200:runtime workspace released", + "complete:IS-101:fleet-a-is-101", ]); }); @@ -116,6 +174,7 @@ test("superseded release preserves pending stop evidence", async () => { await service.stopSuperseded({ sessionId: "IS-101", adapterWorkspaceId: "fleet-a-is-101", + registration, createPending: true, now: 200, }); @@ -144,6 +203,7 @@ test("superseded release records redacted provider failures for retry", async () await service.stopSuperseded({ sessionId: "IS-101", adapterWorkspaceId: "fleet-a-is-101", + registration, createPending: true, now: 200, }); @@ -159,6 +219,325 @@ test("superseded release records redacted provider failures for retry", async () ]); }); +test("superseded cleanup survives ownership loss and retries only the old workspace", async () => { + const replacementWorkspaceId = "fleet-a-is-101-replacement"; + const cleanupRows = new Map(); + const stopped: string[] = []; + const sessionEvidence: string[] = []; + let stopAttempts = 0; + const service = new RuntimeAdapterReleaseService( + releaseDependencies({ + async stageCleanup(input) { + cleanupRows.set(input.adapterWorkspaceId, { + sessionId: input.sessionId, + adapterWorkspaceId: input.adapterWorkspaceId, + registration: input.registration, + createPending: input.createPending, + deletionObserved: false, + claim: "claim-1", + }); + }, + async claimCleanup(_sessionId, adapterWorkspaceId) { + return cleanupRows.get(adapterWorkspaceId) ?? null; + }, + async claimPendingCleanups() { + return [...cleanupRows.values()].map((cleanup) => ({ + ...cleanup, + claim: "claim-2", + })); + }, + async stopWorkspace(_sessionId, adapterWorkspaceId) { + stopped.push(adapterWorkspaceId); + stopAttempts += 1; + return stopAttempts === 1 + ? { status: "stopping", message: "provider stop pending" } + : { status: "stopped", message: "provider workspace released" }; + }, + async persistCleanupEvidence(cleanup, message) { + cleanupRows.set(cleanup.adapterWorkspaceId, { + ...cleanup, + claim: "", + }); + assert.equal(message, "provider stop pending"); + }, + async persistStopEvidence(_sessionId, adapterWorkspaceId) { + if (adapterWorkspaceId === replacementWorkspaceId) { + sessionEvidence.push(adapterWorkspaceId); + } + }, + async confirmRelease(_sessionId, adapterWorkspaceId) { + assert.notEqual(adapterWorkspaceId, replacementWorkspaceId); + return null; + }, + async completeCleanup(cleanup) { + cleanupRows.delete(cleanup.adapterWorkspaceId); + }, + }), + ); + + await service.stopSuperseded({ + sessionId: "IS-101", + adapterWorkspaceId: "fleet-a-is-101-old", + registration, + createPending: true, + now: 200, + }); + assert.equal(cleanupRows.size, 1); + + await service.retryPending(300); + + assert.deepEqual(stopped, ["fleet-a-is-101-old", "fleet-a-is-101-old"]); + assert.deepEqual(sessionEvidence, []); + assert.equal(cleanupRows.size, 0); +}); + +test("superseded provider failures remain independently retryable", async () => { + const cleanupRows: RuntimeAdapterWorkspaceCleanup[] = []; + let fail = true; + const service = new RuntimeAdapterReleaseService( + releaseDependencies({ + async stageCleanup(input) { + cleanupRows.push({ + sessionId: input.sessionId, + adapterWorkspaceId: input.adapterWorkspaceId, + registration: input.registration, + createPending: input.createPending, + deletionObserved: false, + claim: "claim-1", + }); + }, + async claimCleanup() { + return cleanupRows[0] ?? null; + }, + async claimPendingCleanups() { + return cleanupRows; + }, + async stopWorkspace() { + if (fail) { + fail = false; + throw new Error("provider unavailable"); + } + return { status: "stopped", message: "provider workspace released" }; + }, + async persistCleanupEvidence(cleanup, message, _now, reconcileError) { + assert.equal(cleanup.adapterWorkspaceId, "fleet-a-is-101-old"); + assert.equal(message, "superseded runtime adapter stop pending: provider unavailable"); + assert.equal(reconcileError, "provider unavailable"); + }, + async completeCleanup() { + cleanupRows.length = 0; + }, + }), + ); + + await service.stopSuperseded({ + sessionId: "IS-101", + adapterWorkspaceId: "fleet-a-is-101-old", + registration, + createPending: true, + now: 200, + }); + assert.equal(cleanupRows.length, 1); + + await service.retryPending(300); + assert.equal(cleanupRows.length, 0); +}); + +test("create-pending cleanup recovers from a crash before DELETE success is persisted", async () => { + let cleanup: RuntimeAdapterWorkspaceCleanup | null = null; + const retryMissing: boolean[] = []; + let stopAttempt = 0; + let deletionPersistenceAttempt = 0; + const service = new RuntimeAdapterReleaseService( + releaseDependencies({ + async stageCleanup(input) { + cleanup = { + sessionId: input.sessionId, + adapterWorkspaceId: input.adapterWorkspaceId, + registration: input.registration, + createPending: input.createPending, + deletionObserved: false, + claim: "claim-1", + }; + }, + async claimCleanup() { + return cleanup; + }, + async claimPendingCleanups() { + return cleanup ? [{ ...cleanup, claim: `claim-${stopAttempt + 1}` }] : []; + }, + async stopWorkspace(_sessionId, _adapterWorkspaceId, _registration, retry) { + retryMissing.push(retry); + stopAttempt += 1; + if (stopAttempt === 1) { + return { + status: "stopping", + message: "runtime adapter workspace not yet visible; cleanup retry pending", + }; + } + return { + status: "stopped", + message: + stopAttempt === 2 + ? "runtime adapter workspace released" + : "workspace stopped tombstone", + }; + }, + async persistCleanupEvidence(current) { + cleanup = { ...current, claim: "" }; + }, + async markCleanupDeletionObserved(current) { + deletionPersistenceAttempt += 1; + if (deletionPersistenceAttempt === 1) { + throw new Error("crash before DELETE success persistence"); + } + cleanup = { ...current, deletionObserved: true }; + }, + async completeCleanup() { + cleanup = null; + }, + providerError(error) { + assert.ok(error instanceof Error); + return error.message; + }, + }), + ); + + await service.stopSuperseded({ + sessionId: "IS-101", + adapterWorkspaceId: "fleet-a-is-101-old", + registration, + createPending: true, + now: 200, + }); + assert.ok(cleanup); + + await service.retryPending(15_200); + assert.ok(cleanup); + + await service.retryPending(30_200); + + assert.deepEqual(retryMissing, [true, true, true]); + assert.equal(deletionPersistenceAttempt, 2); + assert.equal(cleanup, null); +}); + +test("observed create-pending deletion survives completion persistence failure", async () => { + let cleanup: RuntimeAdapterWorkspaceCleanup | null = null; + const retryMissing: boolean[] = []; + let completionAttempts = 0; + const service = new RuntimeAdapterReleaseService( + releaseDependencies({ + async stageCleanup(input) { + cleanup = { + sessionId: input.sessionId, + adapterWorkspaceId: input.adapterWorkspaceId, + registration: input.registration, + createPending: input.createPending, + deletionObserved: false, + claim: "claim-1", + }; + }, + async claimCleanup() { + return cleanup; + }, + async claimPendingCleanups() { + return cleanup ? [{ ...cleanup, claim: "claim-2" }] : []; + }, + async stopWorkspace(_sessionId, _adapterWorkspaceId, _registration, retry) { + retryMissing.push(retry); + return { + status: "stopped", + message: retry + ? "runtime adapter workspace released" + : "runtime adapter workspace already gone", + }; + }, + async markCleanupDeletionObserved(current) { + cleanup = { ...current, deletionObserved: true }; + }, + async completeCleanup() { + completionAttempts += 1; + if (completionAttempts === 1) throw new Error("completion persistence unavailable"); + cleanup = null; + }, + providerError(error) { + assert.ok(error instanceof Error); + return error.message; + }, + }), + ); + + await service.stopSuperseded({ + sessionId: "IS-101", + adapterWorkspaceId: "fleet-a-is-101-old", + registration, + createPending: true, + now: 200, + }); + assert.equal(cleanup?.deletionObserved, true); + + await service.retryPending(15_200); + + assert.deepEqual(retryMissing, [true, false]); + assert.equal(completionAttempts, 2); + assert.equal(cleanup, null); +}); + +test("runtime adapter cleanup storage is independent and claim fenced", async () => { + const sqlite = new DatabaseSync(":memory:"); + sqlite.exec( + readFileSync( + new URL("../migrations/0037_runtime_adapter_workspace_cleanup.sql", import.meta.url), + "utf8", + ), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0039_runtime_adapter_cleanup_deletion_observed.sql", import.meta.url), + "utf8", + ), + ); + const env = sqliteRuntimeEnv(sqlite); + await stageRuntimeAdapterWorkspaceCleanup(env, { + sessionId: "IS-101", + adapterWorkspaceId: "fleet-a-is-101-old", + registration, + createPending: true, + now: 200, + }); + + const claimed = await claimRuntimeAdapterWorkspaceCleanup( + env, + "IS-101", + "fleet-a-is-101-old", + 200, + ); + assert.ok(claimed); + assert.equal(claimed.createPending, true); + assert.equal(claimed.deletionObserved, false); + assert.deepEqual(claimed.registration, registration); + assert.equal((await claimRuntimeAdapterWorkspaceCleanupBatch(env, 200, 3)).length, 0); + + await markRuntimeAdapterWorkspaceCleanupDeletionObserved(env, claimed, 201); + await persistRuntimeAdapterWorkspaceCleanupEvidence( + env, + claimed, + "provider stop pending", + 200, + null, + ); + assert.equal((await claimRuntimeAdapterWorkspaceCleanupBatch(env, 15_199, 3)).length, 0); + const retry = await claimRuntimeAdapterWorkspaceCleanupBatch(env, 15_200, 3); + assert.equal(retry.length, 1); + assert.equal(retry[0].deletionObserved, true); + await completeRuntimeAdapterWorkspaceCleanup(env, retry[0]); + assert.equal( + sqlite.prepare("SELECT COUNT(*) AS count FROM runtime_adapter_workspace_cleanups").get()?.count, + 0, + ); +}); + test("confirmed release waits for create resolution behind an exact lifecycle fence", async () => { let statements: PreparedStatement[] = []; const effects: string[] = []; @@ -280,7 +659,7 @@ test("confirmed stopped release persists provider evidence before finalization", assert.ok(statements[0].parameters.includes("runtime workspace released")); }); -test("create-pending clearing is fenced to the registered stopping workspace", async () => { +test("create-pending clearing fences the prior marker and advances its revision", async () => { const executions: Array<{ sql: string; parameters: unknown[] }> = []; const env = runtimeEnv((sql, parameters, kind) => { assert.equal(kind, "run"); @@ -296,9 +675,12 @@ test("create-pending clearing is fenced to the registered stopping workspace", a assert.match(executions[0].sql, /"adapter" = \?/i); assert.match(executions[0].sql, /"adapter_workspace_id" = \?/i); assert.match(executions[0].sql, /"status" = \?/i); + assert.match(executions[0].sql, /max\(updated_at \+ 1, \?\)/i); + assert.match(executions[0].sql, /where[\s\S]*"adapter_create_pending" = \?/i); assert.ok(executions[0].parameters.includes("IS-101")); assert.ok(executions[0].parameters.includes("fleet-a-is-101")); assert.ok(executions[0].parameters.includes("stopping")); + assert.ok(executions[0].parameters.includes(1)); }); function releaseEffects(calls: string[]): RuntimeAdapterReleaseEffects { @@ -311,3 +693,64 @@ function releaseEffects(calls: string[]): RuntimeAdapterReleaseEffects { }, }; } + +type BoundStatement = { + execute(): { + results: Record[]; + success: true; + meta: { changes: number; last_row_id?: number }; + }; +}; + +function sqliteRuntimeEnv(sqlite: DatabaseSync): RuntimeEnv { + function execute(sql: string, parameters: unknown[]) { + const statement = sqlite.prepare(sql); + if (/^\s*(?:select|pragma|with)\b|\breturning\b/i.test(sql)) { + const results = statement.all(...parameters).map((row) => ({ ...row })); + const changes = Number(sqlite.prepare("SELECT changes() AS changes").get()?.changes ?? 0); + return { results, success: true as const, meta: { changes } }; + } + const result = statement.run(...parameters); + return { + results: [], + success: true as const, + meta: { + changes: Number(result.changes), + last_row_id: Number(result.lastInsertRowid), + }, + }; + } + return { + DB: { + prepare(sql: string) { + return { + bind(...parameters: unknown[]) { + const bound = { + execute: () => execute(sql, parameters), + async all() { + return bound.execute(); + }, + async run() { + return bound.execute(); + }, + }; + return bound; + }, + }; + }, + async batch(statements: D1PreparedStatement[]) { + sqlite.exec("BEGIN IMMEDIATE"); + try { + const results = statements.map((statement) => + (statement as unknown as BoundStatement).execute(), + ); + sqlite.exec("COMMIT"); + return results; + } catch (error) { + sqlite.exec("ROLLBACK"); + throw error; + } + }, + } as unknown as D1Database, + } as RuntimeEnv; +} diff --git a/tests/runtime-adapter-workspaces.test.ts b/tests/runtime-adapter-workspaces.test.ts index 010d822c..4f76a1f4 100644 --- a/tests/runtime-adapter-workspaces.test.ts +++ b/tests/runtime-adapter-workspaces.test.ts @@ -5,6 +5,8 @@ import { containerCapabilities } from "../src/worker/session-model.ts"; import type { RuntimeEnv } from "../src/worker/env.ts"; import { RuntimeAdapterWorkspaceLifecycle, + runtimeAdapterCapabilitiesHeader, + runtimeAdapterDeleteTombstoneCapability, type RuntimeAdapterWorkspaceLifecycleDependencies, } from "../src/worker/runtime-adapter-workspaces.ts"; import type { InteractiveProvisionResult } from "../src/worker/provisioning/types.ts"; @@ -362,6 +364,214 @@ test("session-bound stop parses DELETE evidence and preserves the registered pat }); }); +test("superseded stop uses retained registration after the session row moves on", async () => { + let databaseReads = 0; + const requests: Array<{ url: string; method: string | undefined }> = []; + const service = new RuntimeAdapterWorkspaceLifecycle( + runtimeEnv(() => { + databaseReads += 1; + return []; + }), + dependencies({ + async fetch(input, init) { + requests.push({ url: input, method: init.method }); + return new Response(null, { status: 204 }); + }, + }), + ); + + const result = await service.stopForSession( + "IS-42", + "workspace-superseded", + { + profile: "default", + controlPlane: "https://adapter.example.test/", + }, + false, + ); + + assert.equal(databaseReads, 0); + assert.deepEqual(requests, [ + { + url: "https://adapter.example.test/v1/workspaces/workspace-superseded", + method: "DELETE", + }, + ]); + assert.deepEqual(result, { + status: "stopped", + message: "runtime adapter workspace released", + }); +}); + +test("superseded pending creates retry DELETE until the old workspace becomes visible", async () => { + const requests: Array<{ url: string; method: string | undefined }> = []; + let responseStatus = 404; + const service = new RuntimeAdapterWorkspaceLifecycle( + runtimeEnv(), + dependencies({ + async fetch(input, init) { + requests.push({ url: input, method: init.method }); + return responseStatus === 204 + ? new Response(null, { status: 204 }) + : Response.json( + { message: "workspace not found" }, + { + status: responseStatus, + headers: { + [runtimeAdapterCapabilitiesHeader]: runtimeAdapterDeleteTombstoneCapability, + }, + }, + ); + }, + }), + ); + const registration = { + profile: "default", + controlPlane: "https://adapter.example.test/", + }; + + assert.deepEqual( + await service.stopForSession("IS-42", "workspace-superseded", registration, true), + { + status: "stopping", + message: "runtime adapter workspace not yet visible; cleanup retry pending", + }, + ); + + responseStatus = 204; + assert.deepEqual( + await service.stopForSession("IS-42", "workspace-superseded", registration, true), + { + status: "stopped", + message: "runtime adapter workspace released", + }, + ); + assert.deepEqual(requests, [ + { + url: "https://adapter.example.test/v1/workspaces/workspace-superseded", + method: "DELETE", + }, + { + url: "https://adapter.example.test/v1/workspaces/workspace-superseded", + method: "DELETE", + }, + ]); +}); + +test("create-pending cleanup preserves legacy 404 release semantics", async () => { + let requestHeaders = new Headers(); + const service = new RuntimeAdapterWorkspaceLifecycle( + runtimeEnv(), + dependencies({ + async fetch(_input, init) { + requestHeaders = new Headers(init.headers); + return Response.json({ message: "workspace not found" }, { status: 404 }); + }, + }), + ); + + assert.deepEqual( + await service.stopForSession( + "IS-42", + "workspace-superseded", + { + profile: "default", + controlPlane: "https://adapter.example.test/", + }, + true, + ), + { + status: "stopped", + message: "workspace not found", + }, + ); + assert.equal( + requestHeaders.get(runtimeAdapterCapabilitiesHeader), + runtimeAdapterDeleteTombstoneCapability, + ); +}); + +test("create-pending cleanup recovers a lost DELETE response from the terminal tombstone", async () => { + let attempt = 0; + const service = new RuntimeAdapterWorkspaceLifecycle( + runtimeEnv(), + dependencies({ + async fetch() { + attempt += 1; + if (attempt === 1) { + return Response.json( + { message: "workspace not found" }, + { + status: 404, + headers: { + [runtimeAdapterCapabilitiesHeader]: runtimeAdapterDeleteTombstoneCapability, + }, + }, + ); + } + if (attempt === 2) { + throw new Error("response lost after delete commit"); + } + return Response.json({ + id: "workspace-superseded", + status: "stopped", + message: "workspace stopped", + }); + }, + }), + ); + const registration = { + profile: "default", + controlPlane: "https://adapter.example.test/", + }; + + assert.deepEqual( + await service.stopForSession("IS-42", "workspace-superseded", registration, true), + { + status: "stopping", + message: "runtime adapter workspace not yet visible; cleanup retry pending", + }, + ); + await assert.rejects( + service.stopForSession("IS-42", "workspace-superseded", registration, true), + /response lost after delete commit/u, + ); + assert.deepEqual( + await service.stopForSession("IS-42", "workspace-superseded", registration, true), + { + status: "stopped", + message: "workspace stopped", + }, + ); +}); + +test("superseded cleanup accepts missing workspaces after deletion was observed", async () => { + const service = new RuntimeAdapterWorkspaceLifecycle( + runtimeEnv(), + dependencies({ + async fetch() { + return Response.json({ message: "workspace not found" }, { status: 404 }); + }, + }), + ); + + assert.deepEqual( + await service.stopForSession( + "IS-42", + "workspace-superseded", + { + profile: "default", + controlPlane: "https://adapter.example.test/", + }, + false, + ), + { + status: "stopped", + message: "workspace not found", + }, + ); +}); + test("session-bound stop redacts provider credentials from failures", async () => { const env = runtimeEnv(() => [ { diff --git a/tests/runtime-adapter.test.ts b/tests/runtime-adapter.test.ts index 5d854924..49f6423e 100644 --- a/tests/runtime-adapter.test.ts +++ b/tests/runtime-adapter.test.ts @@ -1037,6 +1037,14 @@ test("sandbox credential cleanup is durably staged and retried", async () => { new URL("../migrations/0022_credential_policy_cleanup.sql", import.meta.url), "utf8", ); + const registrationStagingMigration = await readFile( + new URL("../migrations/0034_credential_policy_registration_staging.sql", import.meta.url), + "utf8", + ); + const registrationRollbackMigration = await readFile( + new URL("../migrations/0035_credential_policy_registration_rollback.sql", import.meta.url), + "utf8", + ); const scanStart = scannerSource.indexOf("type CredentialPolicyScanRow"); const scanSource = scannerSource.slice(scanStart); const batchStart = cleanupServiceSource.indexOf( @@ -1197,6 +1205,18 @@ test("sandbox credential cleanup is durably staged and retried", async () => { 'stub.fetch("https://crabfleet.internal/api/session-control/register"', ), ); + assert.ok( + registerSource.indexOf("captureSandboxCredentialPolicyRollback") < + registerSource.indexOf( + 'stub.fetch("https://crabfleet.internal/api/session-control/register"', + ), + ); + assert.ok( + registerSource.indexOf("recordSandboxCredentialPolicyRollback") < + registerSource.indexOf( + 'stub.fetch("https://crabfleet.internal/api/session-control/register"', + ), + ); assert.ok( registerSource.indexOf("renewSandboxCredentialPolicyRegistration") < registerSource.indexOf( @@ -1207,14 +1227,25 @@ test("sandbox credential cleanup is durably staged and retried", async () => { registerSource.indexOf('stub.fetch("https://crabfleet.internal/api/session-control/register"') < registerSource.indexOf("finishSandboxCredentialPolicyRegistration"), ); - assert.doesNotMatch(finishSource, /INSERT INTO|insertInto/); - assert.match(finishSource, /state: "active"/); - assert.match(finishSource, /where\(sandboxCredentialPolicyOwnerCondition/); - assert.doesNotMatch(finishSource, /cleanup_pending/); + assert.match(finishSource, /executeBatch/); + assert.match(finishSource, /sandboxCredentialPolicyPromotionQueries/); + assert.match(finishSource, /row\.state === "active"/); + assert.match(registrationLifecycleSource, /INSERT INTO interactive_session_credential_policies/); + assert.match( + registrationLifecycleSource, + /DELETE FROM interactive_session_credential_policy_registrations/, + ); assert.match(registerSource, /abandonSandboxCredentialPolicyRegistration/); - assert.match(abandonSource, /sandboxCredentialPolicyCleanupAuthorizedCondition/); - assert.match(abandonSource, /THEN 'cleanup_pending'/); - assert.match(abandonSource, /ELSE 'registering'/); + assert.match(registerSource, /restoreSandboxCredentialPolicyRollback/); + assert.ok( + scanSource.indexOf("restoreRollback") < + scanSource.indexOf("abandonSandboxCredentialPolicyRegistration"), + ); + assert.match( + abandonSource, + /updateTable\("interactive_session_credential_policy_registrations"\)/, + ); + assert.match(abandonSource, /state: "cleanup_pending"/); assert.ok( scanDecisionSource.indexOf("sandboxExpected") < scanDecisionSource.indexOf("if (registrationAbandoned) return true"), @@ -1229,6 +1260,12 @@ test("sandbox credential cleanup is durably staged and retried", async () => { assert.match(migration, /state IN \('registering', 'active', 'cleanup_pending'\)/); assert.match(migration, /registration_generation TEXT NOT NULL/); assert.match(migration, /registration_claim_expires_at INTEGER/); + assert.match( + registrationStagingMigration, + /CREATE TABLE IF NOT EXISTS interactive_session_credential_policy_registrations/, + ); + assert.match(registrationStagingMigration, /state IN \('registering', 'cleanup_pending'\)/); + assert.match(registrationRollbackMigration, /ADD COLUMN rollback_policies_json TEXT/); assert.match(migration, /CREATE TABLE IF NOT EXISTS credential_policy_reconcile_state/); assert.match(migration, /scan_max_rowid INTEGER NOT NULL/); assert.match(migration, /group_max_session_id TEXT NOT NULL/); diff --git a/tests/runtime-profiles.test.ts b/tests/runtime-profiles.test.ts index b4a8e15c..110a8ce6 100644 --- a/tests/runtime-profiles.test.ts +++ b/tests/runtime-profiles.test.ts @@ -12,6 +12,7 @@ import { runtimeProfileByID, runtimeProfileCapabilities, } from "../src/runtime-profiles.ts"; +import { runtimeAdapterControlPlaneForProfile } from "../src/runtime-adapter.ts"; test("runtime profile catalog preserves generic labels, targets, and capabilities", () => { const profiles = parseRuntimeProfiles( @@ -69,6 +70,10 @@ test("runtime profile catalog fails closed on malformed or ambiguous input", () '[{"id":"a","label":"A","capabilities":null}]', '[{"id":"a","label":"A","capabilities":{"unknown":true}}]', '[{"id":"a","label":"A","privateProvider":"hidden"}]', + '[{"id":"profile/escape","label":"Desktop"}]', + '[{"id":"-profile","label":"Desktop"}]', + '[{"id":"profile.","label":"Desktop"}]', + `[{"id":"${"a".repeat(121)}","label":"Desktop"}]`, '[{"id":"a","label":"A","codexSsh":null}]', '[{"id":"a","label":"A","codexSsh":{"aliasTemplate":"box {sessionId}"}}]', '[{"id":"a","label":"A","codexSsh":{"aliasTemplate":"box-{unknown}"}}]', @@ -81,10 +86,52 @@ test("runtime profile catalog fails closed on malformed or ambiguous input", () for (const value of invalid) { assert.throws(() => parseRuntimeProfiles(value)); } + assert.equal( + parseRuntimeProfiles(JSON.stringify([{ id: `A${"_".repeat(118)}Z`, label: "Maximum" }]))[0]?.id + .length, + 120, + ); assert.deepEqual(parseRuntimeProfiles(undefined), []); assert.deepEqual(parseRuntimeProfiles(""), []); }); +test("fixed adapters accept opaque profile ids without weakening profile-routed URLs", () => { + const profiles = parseRuntimeProfiles( + JSON.stringify([ + { id: "Desktop.PROFILE_2026", label: "Desktop" }, + { id: "desktop_profile", label: "Terminal" }, + ]), + ); + + for (const profile of profiles) { + assert.equal( + runtimeAdapterControlPlaneForProfile( + "https://adapter.example.test/base", + undefined, + profile.id, + ), + "https://adapter.example.test/base", + ); + assert.equal( + runtimeAdapterControlPlaneForProfile( + undefined, + "https://controller.example.test/adapters/{profile}", + profile.id, + ), + null, + ); + } + + assert.equal( + runtimeAdapterControlPlaneForProfile( + undefined, + "https://controller.example.test/adapters/{profile}", + "desktop-profile", + ), + "https://controller.example.test/adapters/desktop-profile", + ); +}); + test("runtime profiles resolve bounded Codex SSH handoff data", () => { const [profile] = parseRuntimeProfiles( JSON.stringify([ diff --git a/tests/sandbox-credential-policy-cleanup.test.ts b/tests/sandbox-credential-policy-cleanup.test.ts index 2f6cd82c..7b1d4d16 100644 --- a/tests/sandbox-credential-policy-cleanup.test.ts +++ b/tests/sandbox-credential-policy-cleanup.test.ts @@ -1,4 +1,5 @@ import assert from "node:assert/strict"; +import { readFile } from "node:fs/promises"; import test from "node:test"; import { @@ -195,7 +196,13 @@ test("terminal cleanup atomically stages the session and credential-policy refs" true, ); - assert.equal(batch.length, 3); + assert.equal(batch.length, 4); + assert.match( + batch[1]?.sql ?? "", + /update "interactive_session_credential_policy_registrations"/i, + ); + assert.match(batch[2]?.sql ?? "", /interactive_session_credential_policies/i); + assert.match(batch[3]?.sql ?? "", /update "interactive_session_credential_policies"/i); const sql = batch.map((statement) => statement.sql).join("\n"); const parameters = batch.flatMap((statement) => statement.parameters); assert.match(sql, /update "interactive_sessions"/i); @@ -205,9 +212,26 @@ test("terminal cleanup atomically stages the session and credential-policy refs" assert.match(sql, /"lease_id" = \?/i); assert.match(sql, /on conflict\s*\(session_id, sandbox_id, lookup_id\)/i); assert.match(sql, /update "interactive_session_credential_policies"/i); + assert.match(sql, /update "interactive_session_credential_policy_registrations"/i); assert.match(sql, /not exists/i); assert.ok(parameters.includes("failed")); assert.ok(parameters.includes("generation:test-1")); assert.ok(parameters.includes("sandbox-1")); assert.ok(parameters.includes(leaseId)); }); + +test("staged cleanup recovers durable historical lookup identities", async () => { + const source = await readFile( + new URL("../src/worker/sandbox-credential-policy-cleanup-service.ts", import.meta.url), + "utf8", + ); + const start = source.indexOf("async function reconcileStagedCredentialPolicyRegistration"); + const end = source.indexOf("async function normalizeCredentialPolicyCleanupGroups", start); + const stagedCleanup = source.slice(start, end); + + assert.match(stagedCleanup, /sandboxCredentialPolicyRegistrationLookupIds/); + assert.match(stagedCleanup, /sandboxCredentialPolicyPersistedLookupIds/); + assert.match(stagedCleanup, /sandboxCredentialPolicyRollbackLookupIds/); + assert.match(stagedCleanup, /registration\.lookup_ids_json/); + assert.match(stagedCleanup, /sandboxLookupIds\(env, registration\.sandbox_id\)/); +}); diff --git a/tests/sandbox-credential-policy-repository.test.ts b/tests/sandbox-credential-policy-repository.test.ts index de2948dc..1484c51a 100644 --- a/tests/sandbox-credential-policy-repository.test.ts +++ b/tests/sandbox-credential-policy-repository.test.ts @@ -1,17 +1,40 @@ import assert from "node:assert/strict"; +import { readFileSync } from "node:fs"; +import { DatabaseSync } from "node:sqlite"; import test from "node:test"; import { activeSandboxCredentialPolicyGeneration, + abandonSandboxCredentialPolicyRegistration, + beginSandboxCredentialPolicyRegistration, + claimObsoleteSandboxCredentialPolicyReferences, + claimSandboxCredentialPolicyRegistrationRecovery, currentSandboxCredentialPolicyGeneration, + finishSandboxCredentialPolicyRegistration, + incompleteSandboxCredentialPolicyGeneration, + markSandboxCredentialPolicyRegistrationWriteStarted, recordSandboxCredentialPolicyRefs, + recordSandboxCredentialPolicyRollback, + repairSandboxCredentialPolicyReferences, + retireObsoleteSandboxCredentialPolicyReference, + renewSandboxCredentialPolicyRegistration, + sandboxCredentialPolicyLookupIdsForGeneration, + sandboxCredentialPolicyPersistedLookupIds, sandboxCredentialPolicyRegistrationQueries, sandboxLookupIds, + stageSandboxCredentialPolicyReferenceRepair, type SandboxCredentialPolicyOwnershipFence, } from "../src/worker/sandbox-credential-policy-repository.ts"; +import { captureSandboxCredentialPolicyRollback } from "../src/worker/sandbox-credential-policy-rollback.ts"; +import { credentialPolicyRegistrationAccepted } from "../src/credential-policy-fence.ts"; import { database } from "../src/worker/database.ts"; import type { RuntimeEnv } from "../src/worker/env.ts"; import type { SandboxCredentialPolicyRegistration } from "../src/worker/session-control-policy.ts"; +import { + sandboxCredentialPolicyRegistrationLookupIds, + sandboxCredentialPolicyRollbackLookupIds, +} from "../src/worker/session-control-policy.ts"; +import type { StoredSandboxCredentialPolicy } from "../src/worker/session-control-policy.ts"; type PreparedStatement = { sql: string; @@ -20,6 +43,19 @@ type PreparedStatement = { run(): Promise; }; +type SqliteStatement = { + all(...parameters: unknown[]): Record[]; + run(...parameters: unknown[]): { changes: number | bigint; lastInsertRowid: number | bigint }; +}; + +type BoundStatement = PreparedStatement & { + execute(): { + results: Record[]; + success: true; + meta: { changes: number; last_row_id?: number }; + }; +}; + function runtimeEnv( handler: ( sql: string, @@ -68,6 +104,204 @@ function runtimeEnv( } as RuntimeEnv; } +function credentialPolicyDatabase(options: { applyMigrations?: boolean } = {}): DatabaseSync { + const db = new DatabaseSync(":memory:"); + db.exec(` + CREATE TABLE interactive_sessions ( + id TEXT PRIMARY KEY, + adapter TEXT, + status TEXT NOT NULL, + credential_cleanup_terminal_status TEXT, + agent_token_hash TEXT, + lease_id TEXT, + sandbox_refresh_sandbox_id TEXT, + sandbox_refresh_claim TEXT, + sandbox_refresh_claim_expires_at INTEGER + ); + CREATE TABLE interactive_session_credential_policies ( + session_id TEXT NOT NULL, + sandbox_id TEXT NOT NULL, + lookup_id TEXT NOT NULL, + state TEXT NOT NULL, + registration_generation TEXT NOT NULL, + registration_claim TEXT, + registration_claim_expires_at INTEGER, + attempt_count INTEGER NOT NULL DEFAULT 0, + last_attempt_at INTEGER, + last_error TEXT, + cleanup_claim TEXT, + cleanup_claim_expires_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (session_id, sandbox_id, lookup_id) + ); + INSERT INTO interactive_sessions ( + id, + adapter, + status, + credential_cleanup_terminal_status, + agent_token_hash, + lease_id + ) VALUES ( + 'IS-42', + NULL, + 'ready', + NULL, + 'agent-token', + 'sandbox:sandbox-1:terminal-1:autostart-v4' + ); + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + created_at, + updated_at + ) VALUES + ('IS-42', 'sandbox-1', 'sandbox-1', 'active', 'generation:existing', NULL, NULL, 1, 1), + ('IS-42', 'sandbox-1', 'do-1', 'active', 'generation:existing', NULL, NULL, 1, 1); + `); + if (options.applyMigrations !== false) { + db.exec( + readFileSync( + new URL("../migrations/0034_credential_policy_registration_staging.sql", import.meta.url), + "utf8", + ), + ); + db.exec( + readFileSync( + new URL("../migrations/0035_credential_policy_registration_rollback.sql", import.meta.url), + "utf8", + ), + ); + db.exec( + readFileSync( + new URL("../migrations/0036_credential_policy_lookup_repair.sql", import.meta.url), + "utf8", + ), + ); + db.exec( + readFileSync( + new URL( + "../migrations/0037_credential_policy_registration_lookup_ids.sql", + import.meta.url, + ), + "utf8", + ), + ); + db.exec( + readFileSync( + new URL( + "../migrations/0040_credential_policy_registration_write_fence.sql", + import.meta.url, + ), + "utf8", + ), + ); + } + return db; +} + +function sqliteRuntimeEnv( + sqlite: DatabaseSync, + options: { + durableObjectId?: string; + interruptAfterStatement?: number; + throwAfterCommit?: boolean; + failNextReadAfterBatch?: boolean; + } = {}, +): RuntimeEnv { + let failNextRead = false; + function execute(sql: string, parameters: unknown[]) { + if (failNextRead && /^\s*select\b/i.test(sql)) { + failNextRead = false; + throw new Error("simulated committed read failure"); + } + const statement = sqlite.prepare(sql) as unknown as SqliteStatement; + if (/^\s*(?:select|pragma|with)\b|\breturning\b/i.test(sql)) { + const results = statement.all(...parameters).map((row) => ({ ...row })); + const changes = Number(sqlite.prepare("SELECT changes() AS changes").get()?.changes ?? 0); + return { results, success: true as const, meta: { changes } }; + } + const result = statement.run(...parameters); + return { + results: [], + success: true as const, + meta: { + changes: Number(result.changes), + last_row_id: Number(result.lastInsertRowid), + }, + }; + } + return { + DB: { + prepare(sql: string) { + return { + bind(...parameters: unknown[]) { + const bound = { + sql, + parameters, + execute: () => execute(sql, parameters), + async all() { + return bound.execute(); + }, + async run() { + return bound.execute(); + }, + }; + return bound; + }, + }; + }, + async batch(statements: D1PreparedStatement[]) { + sqlite.exec("BEGIN IMMEDIATE"); + try { + const results = []; + for (const [index, statement] of statements.entries()) { + results.push((statement as unknown as BoundStatement).execute()); + if (options.interruptAfterStatement === index + 1) { + throw new Error("simulated batch interruption"); + } + } + sqlite.exec("COMMIT"); + failNextRead = options.failNextReadAfterBatch ?? false; + if (options.throwAfterCommit) { + throw new Error("simulated ambiguous committed batch"); + } + return results; + } catch (error) { + if (sqlite.isTransaction) sqlite.exec("ROLLBACK"); + throw error; + } + }, + } as unknown as D1Database, + SANDBOX: { + idFromName() { + return { toString: () => options.durableObjectId ?? "do-1" }; + }, + } as unknown as DurableObjectNamespace, + } as RuntimeEnv; +} + +const ownershipFence: SandboxCredentialPolicyOwnershipFence = { + leaseId: "sandbox:sandbox-1:terminal-1:autostart-v4", + sandboxId: "sandbox-1", +}; + +function activeCredentialPolicyRows(db: DatabaseSync): Record[] { + return db + .prepare(` + SELECT lookup_id, state, registration_generation, registration_claim + FROM interactive_session_credential_policies + ORDER BY lookup_id + `) + .all() + .map((row) => ({ ...row })); +} + const registration: SandboxCredentialPolicyRegistration = { generation: "generation:test-1", claim: "registration-1", @@ -100,6 +334,84 @@ test("credential-policy lookup identity includes the Sandbox durable object id e ]); }); +test("staged lookup identity decoder requires the stable sandbox lookup", () => { + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds('["sandbox-1","do-old"]', "sandbox-1", [ + "sandbox-1", + "do-current", + ]), + ["sandbox-1", "do-old"], + ); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds('["do-old"]', "sandbox-1", [ + "sandbox-1", + "do-current", + ]), + ["sandbox-1"], + ); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds('["sandbox-1","sandbox-1"]', "sandbox-1", [ + "sandbox-1", + "do-current", + ]), + ["sandbox-1"], + ); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds(null, "sandbox-1", ["sandbox-1", "do-current"]), + ["sandbox-1", "do-current"], + ); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds( + null, + "sandbox-1", + ["sandbox-1", "do-current"], + ["do-persisted", "do-rollback", "do-persisted"], + ), + ["sandbox-1", "do-current", "do-persisted", "do-rollback"], + ); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds(null, "sandbox-1", ["do-current"]), + ["sandbox-1"], + ); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds("{", "sandbox-1", ["sandbox-1", "do-current"]), + ["sandbox-1"], + ); +}); + +test("rollback lookup decoder preserves exact valid historical identities", () => { + const rollbackJson = JSON.stringify([ + { + generation: "generation:existing", + policy: { + allowedHosts: [], + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: "sandbox-1", + sessionId: "IS-42", + }, + }, + { + generation: "generation:existing", + policy: { + allowedHosts: [], + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: "do-old", + sessionId: "IS-42", + }, + }, + ]); + assert.deepEqual(sandboxCredentialPolicyRollbackLookupIds(rollbackJson, "IS-42"), [ + "sandbox-1", + "do-old", + ]); + assert.throws( + () => sandboxCredentialPolicyRollbackLookupIds(rollbackJson, "IS-other"), + /rollback snapshot is invalid/, + ); +}); + test("credential-policy generations reuse exactly one current identity", () => { assert.equal(currentSandboxCredentialPolicyGeneration([]), null); assert.equal( @@ -123,10 +435,14 @@ test("credential-policy registration SQL proves every supported ownership fence" sandboxId: "sandbox-1", }); assert.match(current.sql, /from interactive_sessions/i); + assert.match(current.sql, /interactive_session_credential_policy_registrations/i); assert.match(current.sql, /adapter is null|adapter !=/i); assert.match(current.sql, /agent_token_hash is not null/i); assert.match(current.sql, /lease_id =/i); assert.match(current.sql, /sandbox_refresh_claim is null/i); + assert.match(current.sql, /state = 'registering'/i); + assert.match(current.sql, /registration_claim_expires_at >/i); + assert.match(current.sql, /registration_write_started = 0/i); assert.ok(current.parameters.includes("sandbox:sandbox-1:terminal-1:autostart-v4")); assert.doesNotMatch(current.sql, /1 = 1/); @@ -155,35 +471,1921 @@ test("credential-policy registration SQL proves every supported ownership fence" assert.ok(standalone.parameters.includes("standalone-1")); }); -test("active credential-policy generation requires every exact lookup row", async () => { - const rows = [ +test("credential-policy rotation always claims a fresh generation", async () => { + let generation = ""; + let claim = ""; + let registrationExpiresAt = 0; + let statements: PreparedStatement[] = []; + const env = runtimeEnv( + (sql, _parameters, kind) => { + if ( + kind === "all" && + /select .*state/i.test(sql) && + /interactive_session_credential_policy_registrations/i.test(sql) + ) { + return { + results: [ + { + state: "registering", + registration_generation: generation, + registration_claim: claim, + registration_claim_expires_at: registrationExpiresAt, + lookup_ids_json: '["sandbox-1"]', + }, + ], + }; + } + return {}; + }, + (prepared) => { + statements = prepared; + const parameters = prepared.flatMap((statement) => statement.parameters); + generation = String( + parameters.find( + (parameter) => + typeof parameter === "string" && + parameter.startsWith("generation:") && + parameter !== "generation:existing", + ), + ); + claim = String( + parameters.find( + (parameter) => typeof parameter === "string" && parameter.startsWith("registration:"), + ), + ); + registrationExpiresAt = Math.max( + ...parameters.filter((parameter): parameter is number => typeof parameter === "number"), + ); + return prepared.map(() => ({ results: [], meta: { changes: 1 } })); + }, + ); + + const rotated = await beginSandboxCredentialPolicyRegistration(env, "IS-42", "sandbox-1", { + leaseId: "sandbox:sandbox-1:terminal-1:autostart-v4", + sandboxId: "sandbox-1", + }); + + assert.equal(statements.length, 1); + assert.match(rotated.generation, /^generation:/); + assert.notEqual(rotated.generation, "generation:existing"); + assert.equal(rotated.generation, generation); +}); + +test("migration leaves live legacy registrations unstaged while old workers renew them", () => { + const sqlite = credentialPolicyDatabase({ applyMigrations: false }); + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-worker', + registration_claim = 'legacy-registration', + registration_claim_expires_at = 1000, + updated_at = 500 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + sqlite.exec( + readFileSync( + new URL("../migrations/0034_credential_policy_registration_staging.sql", import.meta.url), + "utf8", + ), + ); + + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET registration_claim_expires_at = 3000, updated_at = 2000 + WHERE session_id = 'IS-42' + AND sandbox_id = 'sandbox-1' + AND registration_claim = 'legacy-registration' + `) + .run(); + + assert.equal( + sqlite + .prepare(` + SELECT count(*) AS count + FROM interactive_session_credential_policy_registrations + WHERE state = 'registering' AND registration_claim_expires_at <= 2000 + `) + .get()?.count, + 0, + ); + assert.deepEqual( + sqlite + .prepare(` + SELECT DISTINCT state, registration_generation, registration_claim, + registration_claim_expires_at + FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .all() + .map((row) => ({ ...row })), + [ + { + state: "registering", + registration_generation: "generation:legacy-worker", + registration_claim: "legacy-registration", + registration_claim_expires_at: 3000, + }, + ], + ); +}); + +test("pre-0037 recovery preserves current, persisted, and rollback lookup identities", async () => { + const sqlite = credentialPolicyDatabase({ applyMigrations: false }); + sqlite.exec( + readFileSync( + new URL("../migrations/0034_credential_policy_registration_staging.sql", import.meta.url), + "utf8", + ), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0035_credential_policy_registration_rollback.sql", import.meta.url), + "utf8", + ), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0036_credential_policy_lookup_repair.sql", import.meta.url), + "utf8", + ), + ); + const rollback = [ { - lookup_id: "sandbox-1", - state: "active", - registration_generation: "generation:test-1", - registration_claim: null, + generation: "generation:existing", + policy: { + allowedHosts: [], + githubCredentialSource: "none", + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: "sandbox-1", + sessionId: "IS-42", + }, }, { - lookup_id: "do-1", - state: "active", - registration_generation: "generation:test-1", - registration_claim: null, + generation: "generation:existing", + policy: { + allowedHosts: [], + githubCredentialSource: "none", + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: "do-rollback-old", + sessionId: "IS-42", + }, }, ]; - const env = runtimeEnv(() => ({ results: rows }), undefined, "do-1"); + sqlite + .prepare(` + INSERT INTO interactive_session_credential_policy_registrations ( + session_id, + sandbox_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + rollback_policies_json, + created_at, + updated_at + ) VALUES (?, ?, 'registering', ?, ?, ?, ?, 1, 1) + `) + .run( + "IS-42", + "sandbox-1", + "generation:staged", + "registration:staged", + Number.MAX_SAFE_INTEGER, + JSON.stringify(rollback), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0037_credential_policy_registration_lookup_ids.sql", import.meta.url), + "utf8", + ), + ); + + const row = sqlite + .prepare(` + SELECT lookup_ids_json, repair_generation, rollback_policies_json + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(); + const env = sqliteRuntimeEnv(sqlite, { durableObjectId: "do-current" }); + const persistedLookupIds = await sandboxCredentialPolicyPersistedLookupIds( + env, + "IS-42", + "sandbox-1", + ); + const rollbackLookupIds = sandboxCredentialPolicyRollbackLookupIds( + String(row?.rollback_policies_json), + "IS-42", + ); + assert.equal(row?.lookup_ids_json, null); + assert.deepEqual(persistedLookupIds, ["do-1", "sandbox-1"]); + assert.deepEqual(rollbackLookupIds, ["sandbox-1", "do-rollback-old"]); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds( + row?.lookup_ids_json as string | null, + "sandbox-1", + sandboxLookupIds(env, "sandbox-1"), + [...persistedLookupIds, ...rollbackLookupIds], + ), + ["sandbox-1", "do-current", "do-1", "do-rollback-old"], + ); + assert.equal(row?.repair_generation, null); +}); + +test("post-migration legacy staging recovers the current sandbox lookup set", () => { + const sqlite = credentialPolicyDatabase(); + sqlite + .prepare(` + INSERT INTO interactive_session_credential_policy_registrations ( + session_id, + sandbox_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + created_at, + updated_at + ) VALUES (?, ?, 'registering', ?, ?, ?, 1, 1) + `) + .run( + "IS-42", + "sandbox-1", + "generation:legacy-staged", + "registration:legacy-staged", + Number.MAX_SAFE_INTEGER, + ); + + const row = sqlite + .prepare(` + SELECT lookup_ids_json + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(); + const env = sqliteRuntimeEnv(sqlite); + assert.equal(row?.lookup_ids_json, null); + assert.deepEqual( + sandboxCredentialPolicyRegistrationLookupIds( + row?.lookup_ids_json as string | null, + "sandbox-1", + sandboxLookupIds(env, "sandbox-1"), + ), + ["sandbox-1", "do-1"], + ); +}); + +test("post-migration legacy registration claims block new staged generations", async () => { + const sqlite = credentialPolicyDatabase(); + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-worker', + registration_claim = 'legacy-registration', + registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + + await assert.rejects( + beginSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite), + "IS-42", + "sandbox-1", + ownershipFence, + ), + { message: "sandbox credential policy registration is unavailable" }, + ); + assert.equal( - await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), - "generation:test-1", + sqlite + .prepare("SELECT count(*) AS count FROM interactive_session_credential_policy_registrations") + .get()?.count, + 0, + ); + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => ({ + generation: row.registration_generation, + state: row.state, + claim: row.registration_claim, + })), + [ + { + generation: "generation:legacy-worker", + state: "registering", + claim: "legacy-registration", + }, + { + generation: "generation:legacy-worker", + state: "registering", + claim: "legacy-registration", + }, + ], + ); +}); + +test("staged rotations fence old-worker writes until the staged row is removed", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, ); + const legacyClaim = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-race', + registration_claim = 'legacy-race-claim', + registration_claim_expires_at = ?, + updated_at = 2000 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(legacyClaim.changes, 0); - rows[0] = { ...rows[0]!, registration_generation: "legacy:test-1" }; - rows[1] = { ...rows[1]!, registration_generation: "legacy:test-1" }; - assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => ({ + generation: row.registration_generation, + state: row.state, + })), + [ + { generation: "generation:existing", state: "active" }, + { generation: "generation:existing", state: "active" }, + ], + ); - rows[0] = { ...rows[0]!, registration_generation: "generation:test-1" }; - rows[1] = { ...rows[1]!, registration_generation: "generation:test-1" }; - rows[1] = { ...rows[1]!, registration_claim: "stale" }; - assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + await abandonSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + "simulated registration failure after rollback", + ); + + const legacyCompletion = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'active', + registration_claim = NULL, + registration_claim_expires_at = NULL, + updated_at = 3000 + WHERE session_id = 'IS-42' + AND sandbox_id = 'sandbox-1' + AND registration_generation = 'generation:legacy-race' + AND registration_claim = 'legacy-race-claim' + `) + .run(); + assert.equal(legacyCompletion.changes, 0); + + const legacyDelete = sqlite + .prepare(` + DELETE FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + assert.equal(legacyDelete.changes, 0); + + const legacyInsert = sqlite + .prepare(` + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + created_at, + updated_at + ) VALUES ( + 'IS-42', + 'sandbox-1', + 'legacy-extra', + 'registering', + 'generation:legacy-race', + 'legacy-race-claim', + ?, + 2000, + 2000 + ) + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(legacyInsert.changes, 0); + + assert.equal( + sqlite + .prepare(` + SELECT state + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get()?.state, + "cleanup_pending", + ); + + sqlite + .prepare(` + DELETE FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + const postFenceClaim = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-after-fence', + registration_claim = 'legacy-after-fence-claim', + registration_claim_expires_at = ?, + updated_at = 4000 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(postFenceClaim.changes, 2); +}); + +test("expired write-started registrations remain reserved for recovery", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const rollback = ["sandbox-1", "do-1"].map((lookupId) => ({ + generation: "generation:existing", + policy: { + allowedHosts: [], + githubCredentialSource: "none" as const, + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: lookupId, + sessionId: "IS-42", + }, + })); + assert.equal( + await recordSandboxCredentialPolicyRollback( + env, + "IS-42", + "sandbox-1", + staged, + rollback, + ownershipFence, + ), + true, + ); + assert.ok( + await markSandboxCredentialPolicyRegistrationWriteStarted( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + ); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = 0 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + + await assert.rejects( + beginSandboxCredentialPolicyRegistration(env, "IS-42", "sandbox-1", ownershipFence), + { message: "sandbox credential policy registration is unavailable" }, + ); + + assert.deepEqual( + { + ...sqlite + .prepare(` + SELECT + registration_generation, + registration_claim, + registration_write_started, + rollback_policies_json + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(), + }, + { + registration_generation: staged.generation, + registration_claim: staged.claim, + registration_write_started: 1, + rollback_policies_json: JSON.stringify(rollback), + }, + ); +}); + +test("expired pre-write staged rotations release the legacy worker compatibility fence", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + await beginSandboxCredentialPolicyRegistration(env, "IS-42", "sandbox-1", ownershipFence); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = 0, updated_at = 0 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + + const legacyClaim = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-rollback', + registration_claim = 'legacy-rollback-claim', + registration_claim_expires_at = ?, + updated_at = 1 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + + assert.equal(legacyClaim.changes, 2); + assert.equal( + sqlite + .prepare("SELECT count(*) AS count FROM interactive_session_credential_policy_registrations") + .get()?.count, + 1, + ); +}); + +test("started staged rotations retain the legacy fence after claim expiry", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + assert.ok( + await markSandboxCredentialPolicyRegistrationWriteStarted( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + ); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = 0, updated_at = 0 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + + const legacyClaim = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-rollback', + registration_claim = 'legacy-rollback-claim', + registration_claim_expires_at = ?, + updated_at = 1 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + const legacyCleanupInsert = sqlite + .prepare(` + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + created_at, + updated_at + ) VALUES ( + 'IS-42', + 'sandbox-1', + 'legacy-cleanup', + 'cleanup_pending', + 'generation:existing', + NULL, + NULL, + 1, + 1 + ) + `) + .run(); + const legacyDelete = sqlite + .prepare(` + DELETE FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + + assert.equal(legacyClaim.changes, 0); + assert.equal(legacyCleanupInsert.changes, 0); + assert.equal(legacyDelete.changes, 0); + assert.equal( + sqlite + .prepare(` + SELECT registration_write_started + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get()?.registration_write_started, + 1, + ); + assert.ok( + await claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + staged, + 0, + ownershipFence, + ), + ); +}); + +test("write-fence migration conservatively protects existing staged rotations", () => { + const sqlite = credentialPolicyDatabase({ applyMigrations: false }); + for (const migration of [ + "0034_credential_policy_registration_staging.sql", + "0035_credential_policy_registration_rollback.sql", + "0036_credential_policy_lookup_repair.sql", + "0037_credential_policy_registration_lookup_ids.sql", + ]) { + sqlite.exec(readFileSync(new URL(`../migrations/${migration}`, import.meta.url), "utf8")); + } + sqlite + .prepare(` + INSERT INTO interactive_session_credential_policy_registrations ( + session_id, + sandbox_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + lookup_ids_json, + created_at, + updated_at + ) VALUES (?, ?, 'registering', ?, ?, 0, ?, 1, 0) + `) + .run( + "IS-42", + "sandbox-1", + "generation:pre-write-fence", + "registration:pre-write-fence", + JSON.stringify(["sandbox-1", "do-1"]), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0040_credential_policy_registration_write_fence.sql", import.meta.url), + "utf8", + ), + ); + + assert.equal( + sqlite + .prepare(` + SELECT registration_write_started + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get()?.registration_write_started, + 1, + ); + assert.equal( + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET registration_generation = 'generation:legacy-after-rollback' + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run().changes, + 0, + ); +}); + +test("write-fence migration recovers a crashed namespace repair", async () => { + const sqlite = credentialPolicyDatabase({ applyMigrations: false }); + for (const migration of [ + "0034_credential_policy_registration_staging.sql", + "0035_credential_policy_registration_rollback.sql", + "0036_credential_policy_lookup_repair.sql", + "0037_credential_policy_registration_lookup_ids.sql", + ]) { + sqlite.exec(readFileSync(new URL(`../migrations/${migration}`, import.meta.url), "utf8")); + } + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET lookup_id = 'do-old' + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' AND lookup_id = 'do-1' + `) + .run(); + const expiredRegistration = { + generation: "generation:interrupted", + claim: "registration:interrupted", + lookupIds: ["sandbox-1", "do-1"], + }; + sqlite + .prepare(` + INSERT INTO interactive_session_credential_policy_registrations ( + session_id, + sandbox_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + lookup_ids_json, + repair_generation, + created_at, + updated_at + ) VALUES (?, ?, 'registering', ?, ?, 0, ?, 'generation:existing', 1, 0) + `) + .run( + "IS-42", + "sandbox-1", + expiredRegistration.generation, + expiredRegistration.claim, + JSON.stringify(expiredRegistration.lookupIds), + ); + sqlite.exec( + readFileSync( + new URL("../migrations/0040_credential_policy_registration_write_fence.sql", import.meta.url), + "utf8", + ), + ); + + const env = sqliteRuntimeEnv(sqlite); + const recovered = await claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + expiredRegistration, + 0, + ownershipFence, + ); + assert.ok(recovered); + assert.deepEqual( + await claimObsoleteSandboxCredentialPolicyReferences( + env, + "IS-42", + "sandbox-1", + recovered.registration, + "generation:existing", + ["do-old"], + ownershipFence, + recovered.registrationExpiresAt, + ), + ["do-old"], + ); +}); + +test("stale staged cleanup releases legacy deletion but an active cleanup claim stays fenced", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + await beginSandboxCredentialPolicyRegistration(env, "IS-42", "sandbox-1", ownershipFence); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET + state = 'cleanup_pending', + registration_claim = NULL, + registration_claim_expires_at = NULL, + cleanup_claim = 'cleanup:new-worker', + cleanup_claim_expires_at = ?, + updated_at = 0 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + + const fencedDelete = sqlite + .prepare(` + DELETE FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + assert.equal(fencedDelete.changes, 0); + + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET cleanup_claim = NULL, cleanup_claim_expires_at = NULL, updated_at = 0 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + const legacyDelete = sqlite + .prepare(` + DELETE FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + assert.equal(legacyDelete.changes, 2); +}); + +test("partial credential-policy rotation failure preserves the prior active generation", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const rollback = ["sandbox-1", "do-1"].map((lookupId) => ({ + generation: "generation:existing", + policy: { + allowedHosts: [], + githubCredentialSource: "none" as const, + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: lookupId, + sessionId: "IS-42", + }, + })); + assert.equal( + await recordSandboxCredentialPolicyRollback( + env, + "IS-42", + "sandbox-1", + staged, + rollback, + ownershipFence, + ), + true, + ); + + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => row.registration_generation), + ["generation:existing", "generation:existing"], + ); + + await abandonSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + "simulated Durable Object registration failure", + ); + + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => row.registration_generation), + ["generation:existing", "generation:existing"], + ); + assert.deepEqual( + { + ...sqlite + .prepare(` + SELECT + state, + registration_generation, + registration_claim, + rollback_policies_json, + last_error + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(), + }, + { + state: "cleanup_pending", + registration_generation: staged.generation, + registration_claim: null, + rollback_policies_json: JSON.stringify(rollback), + last_error: "simulated Durable Object registration failure", + }, + ); +}); + +test("stale foreground rollback cannot renew after recovery takes its claim", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const expiredAt = 1; + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(expiredAt); + const recovery = await claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + staged, + expiredAt, + ownershipFence, + ); + assert.ok(recovery); + + const renewed = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ); + + assert.equal(renewed, null); +}); + +test("expired staged registration claims cannot be revived by renewal", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const expiredAt = 1; + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(expiredAt); + + assert.equal( + await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + null, + ); + assert.equal( + sqlite + .prepare(` + SELECT registration_claim_expires_at + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get()?.registration_claim_expires_at, + expiredAt, + ); +}); + +test("staged renewal cannot extend across a live legacy registration claim", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = 1 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + const legacyClaim = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-renewal', + registration_claim = 'registration:legacy-renewal', + registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(legacyClaim.changes, 2); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + + assert.equal( + await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + null, + ); + assert.deepEqual( + sqlite + .prepare(` + SELECT DISTINCT registration_generation, registration_claim, registration_claim_expires_at + FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .all() + .map((row) => ({ ...row })), + [ + { + registration_generation: "generation:legacy-renewal", + registration_claim: "registration:legacy-renewal", + registration_claim_expires_at: Number.MAX_SAFE_INTEGER, + }, + ], + ); +}); + +test("expired staged recovery cannot race a live legacy registration owner", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const expiredAt = 1; + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(expiredAt); + const legacyClaim = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_generation = 'generation:legacy-recovery', + registration_claim = 'registration:legacy-recovery', + registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(legacyClaim.changes, 2); + + const recoveries = await Promise.all([ + claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + staged, + expiredAt, + ownershipFence, + ), + claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + staged, + expiredAt, + ownershipFence, + ), + ]); + + assert.deepEqual(recoveries, [null, null]); + assert.deepEqual( + { + ...sqlite + .prepare(` + SELECT registration_claim, registration_claim_expires_at + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(), + }, + { + registration_claim: staged.claim, + registration_claim_expires_at: expiredAt, + }, + ); + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => ({ + generation: row.registration_generation, + claim: row.registration_claim, + })), + [ + { + generation: "generation:legacy-recovery", + claim: "registration:legacy-recovery", + }, + { + generation: "generation:legacy-recovery", + claim: "registration:legacy-recovery", + }, + ], + ); +}); + +test("expired registration recovery grants one fresh exclusive claim", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const expiredAt = 1; + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(expiredAt); + + const claims = await Promise.all([ + claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + staged, + expiredAt, + ownershipFence, + ), + claimSandboxCredentialPolicyRegistrationRecovery( + env, + "IS-42", + "sandbox-1", + staged, + expiredAt, + ownershipFence, + ), + ]); + const winner = claims.find((claim) => claim !== null); + + assert.equal(claims.filter((claim) => claim !== null).length, 1); + assert.ok(winner); + assert.notEqual(winner.registration.claim, staged.claim); + assert.ok(winner.registrationExpiresAt > expiredAt); + assert.deepEqual( + { + ...sqlite + .prepare(` + SELECT registration_generation, registration_claim, registration_claim_expires_at + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(), + }, + { + registration_generation: staged.generation, + registration_claim: winner.registration.claim, + registration_claim_expires_at: winner.registrationExpiresAt, + }, + ); +}); + +test("completed credential-policy rotation atomically promotes every active lookup", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + + assert.equal( + await finishSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + true, + ); + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => ({ + generation: row.registration_generation, + state: row.state, + })), + [ + { generation: staged.generation, state: "active" }, + { generation: staged.generation, state: "active" }, + ], + ); + assert.equal( + sqlite + .prepare("SELECT count(*) AS count FROM interactive_session_credential_policy_registrations") + .get()?.count, + 0, + ); +}); + +test("expired credential-policy claims cannot promote active authority", async () => { + const sqlite = credentialPolicyDatabase(); + const env = sqliteRuntimeEnv(sqlite); + const staged = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + sqlite + .prepare(` + UPDATE interactive_session_credential_policy_registrations + SET registration_claim_expires_at = 0 + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(); + + assert.equal( + await finishSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + false, + ); + assert.equal( + activeCredentialPolicyRows(sqlite).some( + (row) => row.registration_generation === staged.generation && row.state === "active", + ), + false, + ); +}); + +test("completed credential-policy rotation tolerates an ambiguous committed batch", async () => { + const sqlite = credentialPolicyDatabase(); + const staged = await beginSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite), + "IS-42", + "sandbox-1", + ownershipFence, + ); + + assert.equal( + await finishSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite, { throwAfterCommit: true }), + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + true, + ); +}); + +test("completed credential-policy rotation retries an ambiguous verification read", async () => { + const sqlite = credentialPolicyDatabase(); + const staged = await beginSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite), + "IS-42", + "sandbox-1", + ownershipFence, + ); + + assert.equal( + await finishSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite, { failNextReadAfterBatch: true }), + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + true, + ); +}); + +test("interrupted credential-policy promotion rolls back every active lookup", async () => { + const sqlite = credentialPolicyDatabase(); + const staged = await beginSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite), + "IS-42", + "sandbox-1", + ownershipFence, + ); + + await assert.rejects( + finishSandboxCredentialPolicyRegistration( + sqliteRuntimeEnv(sqlite, { interruptAfterStatement: 1 }), + "IS-42", + "sandbox-1", + staged, + ownershipFence, + ), + /simulated batch interruption/, + ); + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => row.registration_generation), + ["generation:existing", "generation:existing"], + ); + assert.equal( + sqlite + .prepare(` + SELECT registration_generation + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get()?.registration_generation, + staged.generation, + ); +}); + +test("active credential-policy generation requires every exact lookup row", async () => { + const rows = [ + { + lookup_id: "sandbox-1", + state: "active", + registration_generation: "generation:test-1", + registration_claim: null, + }, + { + lookup_id: "do-1", + state: "active", + registration_generation: "generation:test-1", + registration_claim: null, + }, + ]; + const env = runtimeEnv(() => ({ results: rows }), undefined, "do-1"); + assert.equal( + await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), + "generation:test-1", + ); + + rows[0] = { ...rows[0]!, registration_generation: "legacy:test-1" }; + rows[1] = { ...rows[1]!, registration_generation: "legacy:test-1" }; + assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + + rows[0] = { ...rows[0]!, registration_generation: "generation:test-1" }; + rows[1] = { ...rows[1]!, registration_generation: "generation:test-1" }; + rows[1] = { ...rows[1]!, registration_claim: "stale" }; + assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + + rows[1] = { ...rows[1]!, registration_claim: null }; + rows.push({ + lookup_id: "do-obsolete", + state: "active", + registration_generation: "generation:test-1", + registration_claim: null, + }); + assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); +}); + +test("credential refresh repairs an incomplete legacy lookup set before rotation", async () => { + const sqlite = credentialPolicyDatabase(); + sqlite + .prepare(` + DELETE FROM interactive_session_credential_policies + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' AND lookup_id = 'do-1' + `) + .run(); + const now = Date.now(); + const policies = new Map([ + [ + "sandbox-1", + { + generation: "generation:existing", + registrationClaim: "registration:legacy", + registrationExpiresAt: now + 30_000, + policy: { + allowedHosts: [], + githubCredentialSource: "none", + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: "sandbox-1", + sessionId: "IS-42", + }, + }, + ], + ]); + const stub = { + async fetch(input: RequestInfo | URL, init?: RequestInit): Promise { + const url = new URL(String(input)); + const egress = url.pathname.match(/^\/api\/session-control\/egress\/([^/]+)$/); + if (egress && (!init?.method || init.method === "GET")) { + const current = policies.get(decodeURIComponent(egress[1] ?? "")); + return current + ? Response.json(current.policy, { + headers: { "x-crabfleet-policy-generation": current.generation }, + }) + : Response.json({ error: "not found" }, { status: 404 }); + } + if (url.pathname === "/api/session-control/register" && init?.method === "POST") { + const incoming = JSON.parse(String(init.body)) as StoredSandboxCredentialPolicy; + const current = policies.get(incoming.policy.sandboxId); + if (!credentialPolicyRegistrationAccepted(current, undefined, incoming, Date.now())) { + return Response.json({ error: "conflict" }, { status: 409 }); + } + policies.set(incoming.policy.sandboxId, incoming); + return Response.json({ ok: true }); + } + return Response.json({ error: "not found" }, { status: 404 }); + }, + }; + const env = sqliteRuntimeEnv(sqlite); + + assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + assert.equal( + await incompleteSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), + "generation:existing", + ); + await assert.rejects( + captureSandboxCredentialPolicyRollback(stub, sandboxLookupIds(env, "sandbox-1"), null, "IS-42"), + /no durable rollback owner/, + ); + + const registration = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const repairGeneration = await incompleteSandboxCredentialPolicyGeneration( + env, + "IS-42", + "sandbox-1", + ); + assert.equal(repairGeneration, "generation:existing"); + let registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + await stageSandboxCredentialPolicyReferenceRepair( + env, + "IS-42", + "sandbox-1", + registration, + repairGeneration, + ownershipFence, + ), + true, + ); + const legacyUpdate = sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'registering', + registration_claim = 'registration:legacy-race', + registration_claim_expires_at = ? + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(legacyUpdate.changes, 0); + const legacyInsert = sqlite + .prepare(` + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + created_at, + updated_at + ) VALUES ( + 'IS-42', + 'sandbox-1', + 'do-1', + 'registering', + 'generation:existing', + 'registration:legacy-race', + ?, + 1, + 1 + ) + `) + .run(Number.MAX_SAFE_INTEGER); + assert.equal(legacyInsert.changes, 0); + assert.equal( + ( + await stub.fetch("https://crabfleet.internal/api/session-control/register", { + method: "POST", + body: JSON.stringify({ + generation: repairGeneration, + registrationClaim: registration.claim, + registrationExpiresAt: registrationExpiresAt - 1, + policy: { ...policies.get("sandbox-1")!.policy, sandboxId: "do-1" }, + } satisfies StoredSandboxCredentialPolicy), + }) + ).ok, + true, + ); + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + await repairSandboxCredentialPolicyReferences( + env, + "IS-42", + "sandbox-1", + registration, + repairGeneration, + ownershipFence, + registrationExpiresAt, + ), + true, + ); + const rollback = await captureSandboxCredentialPolicyRollback( + stub, + registration.lookupIds, + repairGeneration, + "IS-42", + ); + assert.equal(rollback.length, 2); + assert.equal( + await recordSandboxCredentialPolicyRollback( + env, + "IS-42", + "sandbox-1", + registration, + rollback, + ownershipFence, + ), + true, + ); + for (const lookupId of registration.lookupIds) { + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + ( + await stub.fetch("https://crabfleet.internal/api/session-control/register", { + method: "POST", + body: JSON.stringify({ + generation: registration.generation, + registrationClaim: registration.claim, + registrationExpiresAt, + policy: { ...policies.get(lookupId)!.policy, sandboxId: lookupId }, + } satisfies StoredSandboxCredentialPolicy), + }) + ).ok, + true, + ); + } + assert.equal( + await finishSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ), + true, + ); + + const generation = await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"); + assert.match(generation ?? "", /^generation:/); + assert.notEqual(generation, "generation:existing"); + assert.deepEqual( + [...policies.entries()].map(([lookupId, policy]) => ({ + lookupId, + generation: policy.generation, + sandboxId: policy.policy.sandboxId, + })), + [ + { lookupId: "sandbox-1", generation, sandboxId: "sandbox-1" }, + { lookupId: "do-1", generation, sandboxId: "do-1" }, + ], + ); +}); + +test("credential refresh replaces an obsolete durable namespace without losing rollback coverage", async () => { + const sqlite = credentialPolicyDatabase(); + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET lookup_id = 'do-old' + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' AND lookup_id = 'do-1' + `) + .run(); + const now = Date.now(); + const policy = { + allowedHosts: [], + githubCredentialSource: "none" as const, + githubRepo: "openclaw/crabfleet", + owner: "operator", + sessionId: "IS-42", + }; + const policies = new Map( + ["sandbox-1", "do-old"].map((lookupId) => [ + lookupId, + { + generation: "generation:existing", + registrationClaim: "registration:legacy", + registrationExpiresAt: now + 1_000, + policy: { ...policy, sandboxId: lookupId }, + }, + ]), + ); + const tombstones = new Map(); + const stub = { + async fetch(input: RequestInfo | URL, init?: RequestInit): Promise { + const url = new URL(String(input)); + const egress = url.pathname.match(/^\/api\/session-control\/egress\/([^/]+)$/); + if (egress && (!init?.method || init.method === "GET")) { + const current = policies.get(decodeURIComponent(egress[1] ?? "")); + return current + ? Response.json(current.policy, { + headers: { "x-crabfleet-policy-generation": current.generation }, + }) + : Response.json({ error: "not found" }, { status: 404 }); + } + if (url.pathname === "/api/session-control/register" && init?.method === "POST") { + const incoming = JSON.parse(String(init.body)) as StoredSandboxCredentialPolicy; + const lookupId = incoming.policy.sandboxId; + if ( + tombstones.get(lookupId) === incoming.generation || + !credentialPolicyRegistrationAccepted( + policies.get(lookupId), + undefined, + incoming, + Date.now(), + ) + ) { + return Response.json({ error: "conflict" }, { status: 409 }); + } + policies.set(lookupId, incoming); + return Response.json({ ok: true }); + } + const removal = url.pathname.match(/^\/api\/session-control\/sandbox\/([^/]+)$/); + if (removal && init?.method === "DELETE") { + const lookupId = decodeURIComponent(removal[1] ?? ""); + const tombstone = JSON.parse(String(init.body)) as { + generation: string; + sessionId: string; + }; + tombstones.set(lookupId, tombstone.generation); + const current = policies.get(lookupId); + if ( + current?.generation === tombstone.generation && + current.policy.sessionId === tombstone.sessionId + ) { + policies.delete(lookupId); + } + return Response.json({ ok: true }); + } + return Response.json({ error: "not found" }, { status: 404 }); + }, + }; + const env = sqliteRuntimeEnv(sqlite); + + assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + assert.equal( + await incompleteSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), + "generation:existing", + ); + const registration = await beginSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + ownershipFence, + ); + const repairGeneration = await incompleteSandboxCredentialPolicyGeneration( + env, + "IS-42", + "sandbox-1", + ); + assert.equal(repairGeneration, "generation:existing"); + assert.deepEqual( + await sandboxCredentialPolicyLookupIdsForGeneration( + env, + "IS-42", + "sandbox-1", + repairGeneration, + ), + ["do-old", "sandbox-1"], + ); + let registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + await stageSandboxCredentialPolicyReferenceRepair( + env, + "IS-42", + "sandbox-1", + registration, + repairGeneration, + ownershipFence, + ), + true, + ); + registrationExpiresAt = await markSandboxCredentialPolicyRegistrationWriteStarted( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + ( + await stub.fetch("https://crabfleet.internal/api/session-control/register", { + method: "POST", + body: JSON.stringify({ + generation: repairGeneration, + registrationClaim: registration.claim, + registrationExpiresAt: registrationExpiresAt - 1, + policy: { ...policy, sandboxId: "do-1" }, + } satisfies StoredSandboxCredentialPolicy), + }) + ).ok, + true, + ); + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + await repairSandboxCredentialPolicyReferences( + env, + "IS-42", + "sandbox-1", + registration, + repairGeneration, + ownershipFence, + registrationExpiresAt, + ), + true, + ); + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + sqlite + .prepare(` + UPDATE interactive_session_credential_policies + SET + state = 'cleanup_pending', + cleanup_claim = 'registration:wrong', + cleanup_claim_expires_at = ?, + updated_at = ? + WHERE session_id = 'IS-42' + AND sandbox_id = 'sandbox-1' + AND lookup_id = 'do-old' + AND state = 'active' + `) + .run(registrationExpiresAt, Date.now()).changes, + 0, + ); + assert.deepEqual( + await claimObsoleteSandboxCredentialPolicyReferences( + env, + "IS-42", + "sandbox-1", + registration, + repairGeneration, + ["do-old"], + ownershipFence, + registrationExpiresAt, + ), + ["do-old"], + ); + assert.equal( + ( + await stub.fetch("https://crabfleet.internal/api/session-control/sandbox/do-old", { + method: "DELETE", + body: JSON.stringify({ + generation: repairGeneration, + sessionId: "IS-42", + tombstonedAt: Date.now(), + }), + }) + ).ok, + true, + ); + assert.equal( + await retireObsoleteSandboxCredentialPolicyReference( + env, + "IS-42", + "sandbox-1", + registration, + repairGeneration, + "do-old", + ownershipFence, + registrationExpiresAt, + ), + true, + ); + assert.equal( + await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), + repairGeneration, + ); + const rollback = await captureSandboxCredentialPolicyRollback( + stub, + registration.lookupIds, + repairGeneration, + "IS-42", + ); + assert.equal( + await recordSandboxCredentialPolicyRollback( + env, + "IS-42", + "sandbox-1", + registration, + rollback, + ownershipFence, + ), + true, + ); + for (const lookupId of registration.lookupIds) { + registrationExpiresAt = await renewSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ); + assert.ok(registrationExpiresAt); + assert.equal( + ( + await stub.fetch("https://crabfleet.internal/api/session-control/register", { + method: "POST", + body: JSON.stringify({ + generation: registration.generation, + registrationClaim: registration.claim, + registrationExpiresAt, + policy: { ...policy, sandboxId: lookupId }, + } satisfies StoredSandboxCredentialPolicy), + }) + ).ok, + true, + ); + } + assert.equal( + await finishSandboxCredentialPolicyRegistration( + env, + "IS-42", + "sandbox-1", + registration, + ownershipFence, + ), + true, + ); + + const generation = await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"); + assert.match(generation ?? "", /^generation:/); + assert.notEqual(generation, "generation:existing"); + assert.deepEqual( + activeCredentialPolicyRows(sqlite).map((row) => row.lookup_id), + ["do-1", "sandbox-1"], + ); + assert.deepEqual([...policies.keys()].sort(), ["do-1", "sandbox-1"]); + assert.equal(tombstones.get("do-old"), "generation:existing"); + assert.equal(policies.has("do-old"), false); + assert.equal( + sqlite + .prepare("SELECT count(*) AS count FROM interactive_session_credential_policy_registrations") + .get()?.count, + 0, + ); }); test("recording active policy refs promotes then upserts every lookup under one fence", async () => { diff --git a/tests/sandbox-credential-policy-rollback.test.ts b/tests/sandbox-credential-policy-rollback.test.ts new file mode 100644 index 00000000..2b68ec6a --- /dev/null +++ b/tests/sandbox-credential-policy-rollback.test.ts @@ -0,0 +1,121 @@ +import assert from "node:assert/strict"; +import test from "node:test"; + +import { credentialPolicyRegistrationAccepted } from "../src/credential-policy-fence.ts"; +import { + captureSandboxCredentialPolicyRollback, + restoreSandboxCredentialPolicyRollback, +} from "../src/worker/sandbox-credential-policy-rollback.ts"; +import type { + SandboxCredentialPolicyRegistration, + StoredSandboxCredentialPolicy, +} from "../src/worker/session-control-policy.ts"; + +function storedPolicy( + lookupId: string, + generation: string, + claim: string, + registrationExpiresAt: number, +): StoredSandboxCredentialPolicy { + return { + generation, + registrationClaim: claim, + registrationExpiresAt, + policy: { + allowedHosts: [], + githubCredentialSource: "none", + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: lookupId, + sessionId: "IS-42", + }, + }; +} + +function policyStub(policies: Map) { + return { + async fetch(input: RequestInfo | URL, init?: RequestInit): Promise { + const url = new URL(String(input)); + const egress = url.pathname.match(/^\/api\/session-control\/egress\/([^/]+)$/); + if (egress && (!init?.method || init.method === "GET")) { + const current = policies.get(decodeURIComponent(egress[1] ?? "")); + return current + ? Response.json(current.policy, { + headers: { "x-crabfleet-policy-generation": current.generation }, + }) + : Response.json({ error: "not found" }, { status: 404 }); + } + if (url.pathname === "/api/session-control/register" && init?.method === "POST") { + const incoming = JSON.parse(String(init.body)) as StoredSandboxCredentialPolicy; + const current = policies.get(incoming.policy.sandboxId); + if (!credentialPolicyRegistrationAccepted(current, undefined, incoming, Date.now())) { + return Response.json({ error: "conflict" }, { status: 409 }); + } + policies.set(incoming.policy.sandboxId, incoming); + return Response.json({ ok: true }); + } + return Response.json({ error: "not found" }, { status: 404 }); + }, + }; +} + +test("partial policy generation writes restore every prior live lookup", async () => { + const now = Date.now(); + const lookupIds = ["sandbox-1", "do-1"]; + const policies = new Map( + lookupIds.map((lookupId) => [ + lookupId, + storedPolicy(lookupId, "generation:prior", "registration:prior", now + 1_000), + ]), + ); + const stub = policyStub(policies); + const rollback = await captureSandboxCredentialPolicyRollback( + stub, + lookupIds, + "generation:prior", + "IS-42", + ); + const registration: SandboxCredentialPolicyRegistration = { + generation: "generation:replacement", + claim: "registration:replacement", + lookupIds, + }; + const replacementExpiresAt = now + 60_000; + + policies.set( + "sandbox-1", + storedPolicy("sandbox-1", registration.generation, registration.claim, replacementExpiresAt), + ); + await restoreSandboxCredentialPolicyRollback( + stub, + registration, + replacementExpiresAt, + JSON.stringify(rollback), + "IS-42", + ); + + for (const current of policies.values()) { + assert.equal(current.generation, "generation:prior"); + assert.match(current.registrationClaim, /^rollback:registration:replacement$/); + assert.ok(current.registrationExpiresAt > replacementExpiresAt); + } +}); + +test("rollback snapshots reject incomplete prior generations before replacement", async () => { + const policies = new Map([ + [ + "sandbox-1", + storedPolicy("sandbox-1", "generation:prior", "registration:prior", Date.now() + 1_000), + ], + ]); + + await assert.rejects( + captureSandboxCredentialPolicyRollback( + policyStub(policies), + ["sandbox-1", "do-1"], + "generation:prior", + "IS-42", + ), + /rollback generation is incomplete/, + ); +}); diff --git a/tests/sandbox-credential-policy-scanner.test.ts b/tests/sandbox-credential-policy-scanner.test.ts index 09e667c9..59bd7009 100644 --- a/tests/sandbox-credential-policy-scanner.test.ts +++ b/tests/sandbox-credential-policy-scanner.test.ts @@ -1,15 +1,121 @@ import assert from "node:assert/strict"; +import { readFile } from "node:fs/promises"; +import { DatabaseSync } from "node:sqlite"; import test from "node:test"; import { credentialPolicyProvisioningStaleMs, credentialPolicyScanOwnershipFence, credentialPolicyScanRequiresCleanup, + scanCredentialPolicyCleanupPage, type CredentialPolicyScanRow, } from "../src/worker/sandbox-credential-policy-scanner.ts"; +import type { RuntimeEnv } from "../src/worker/env.ts"; +import { activeSandboxCredentialPolicyGeneration } from "../src/worker/sandbox-credential-policy-repository.ts"; const now = 2_000_000; +type SqliteStatement = { + all(...parameters: unknown[]): Record[]; + run(...parameters: unknown[]): { changes: number | bigint; lastInsertRowid: number | bigint }; +}; + +function scannerDatabase(): DatabaseSync { + const db = new DatabaseSync(":memory:"); + db.exec(` + CREATE TABLE interactive_sessions ( + id TEXT PRIMARY KEY, + adapter TEXT, + status TEXT NOT NULL, + lease_id TEXT, + credential_cleanup_terminal_status TEXT, + sandbox_refresh_sandbox_id TEXT, + sandbox_refresh_claim TEXT, + sandbox_refresh_claim_expires_at INTEGER, + agent_token_hash TEXT, + updated_at INTEGER NOT NULL + ); + CREATE TABLE standalone_sandbox_provisions ( + id TEXT PRIMARY KEY, + sandbox_id TEXT NOT NULL, + state TEXT NOT NULL, + ownership_claim TEXT, + ownership_claim_expires_at INTEGER, + updated_at INTEGER NOT NULL + ); + CREATE TABLE interactive_session_credential_policies ( + session_id TEXT NOT NULL, + sandbox_id TEXT NOT NULL, + lookup_id TEXT NOT NULL, + state TEXT NOT NULL, + registration_generation TEXT NOT NULL, + registration_claim TEXT, + registration_claim_expires_at INTEGER, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (session_id, sandbox_id, lookup_id) + ); + CREATE TABLE interactive_session_credential_policy_registrations ( + session_id TEXT NOT NULL, + sandbox_id TEXT NOT NULL, + state TEXT NOT NULL, + registration_generation TEXT NOT NULL, + registration_claim TEXT, + registration_claim_expires_at INTEGER, + registration_write_started INTEGER NOT NULL DEFAULT 0, + lookup_ids_json TEXT, + rollback_policies_json TEXT, + last_error TEXT, + updated_at INTEGER NOT NULL, + PRIMARY KEY (session_id, sandbox_id) + ); + `); + return db; +} + +function scannerRuntimeEnv(sqlite: DatabaseSync): RuntimeEnv { + function execute(sql: string, parameters: unknown[]) { + const statement = sqlite.prepare(sql) as unknown as SqliteStatement; + if (/^\s*(?:select|pragma|with)\b|\breturning\b/i.test(sql)) { + const results = statement.all(...parameters).map((row) => ({ ...row })); + const changes = Number(sqlite.prepare("SELECT changes() AS changes").get()?.changes ?? 0); + return { results, success: true as const, meta: { changes } }; + } + const result = statement.run(...parameters); + return { + results: [], + success: true as const, + meta: { + changes: Number(result.changes), + last_row_id: Number(result.lastInsertRowid), + }, + }; + } + return { + DB: { + prepare(sql: string) { + return { + bind(...parameters: unknown[]) { + return { + async all() { + return execute(sql, parameters); + }, + async run() { + return execute(sql, parameters); + }, + }; + }, + }; + }, + } as unknown as D1Database, + SANDBOX: { + idFromName() { + return { toString: () => "do-current" }; + }, + } as unknown as DurableObjectNamespace, + } as RuntimeEnv; +} + function scanRow(values: Partial = {}): CredentialPolicyScanRow { return { scan_rowid: 1, @@ -122,6 +228,162 @@ test("credential-policy scan rejects incomplete, expired, and mismatched ownersh ); }); +test("staged recovery takes a fresh exclusive claim before promotion or rollback", async () => { + const source = await readFile( + new URL("../src/worker/sandbox-credential-policy-scanner.ts", import.meta.url), + "utf8", + ); + const start = source.indexOf("async function scanStagedCredentialPolicyRegistrations"); + const end = source.indexOf("async function readCredentialPolicyScanPage", start); + const stagedRecovery = source.slice(start, end); + + assert.ok( + stagedRecovery.indexOf("if (!ownershipFence)") < stagedRecovery.indexOf("restoreRollback({"), + ); + assert.ok( + stagedRecovery.indexOf("claimSandboxCredentialPolicyRegistrationRecovery(") < + stagedRecovery.indexOf("policyExists("), + ); + assert.ok( + stagedRecovery.indexOf("claimSandboxCredentialPolicyRegistrationRecovery(") < + stagedRecovery.indexOf("restoreRollback({"), + ); + assert.ok( + stagedRecovery.indexOf("claimSandboxCredentialPolicyRegistrationRecovery(") < + stagedRecovery.indexOf("sandboxCredentialPolicyRollbackIsSuperseded("), + ); + assert.ok( + stagedRecovery.indexOf("sandboxCredentialPolicyRollbackIsSuperseded(") < + stagedRecovery.indexOf("restoreRollback({"), + ); + assert.match(stagedRecovery, /sandboxCredentialPolicyRegistrationLookupIds/); + assert.match(stagedRecovery, /sandboxCredentialPolicyPersistedLookupIds/); + assert.match(stagedRecovery, /sandboxCredentialPolicyRollbackLookupIds/); + assert.match(stagedRecovery, /registration\.lookupIds/); + assert.doesNotMatch(stagedRecovery, /renewSandboxCredentialPolicyRegistration/); +}); + +test("staged recovery refuses rollback when one historical legacy lookup advances", async () => { + const sqlite = scannerDatabase(); + const env = scannerRuntimeEnv(sqlite); + const rollback = ["sandbox-1", "do-rollback"].map((lookupId) => ({ + generation: "generation:rollback", + policy: { + allowedHosts: [], + githubCredentialSource: "none", + githubRepo: "openclaw/crabfleet", + owner: "operator", + sandboxId: lookupId, + sessionId: "IS-42", + }, + })); + sqlite + .prepare(` + INSERT INTO interactive_sessions ( + id, + adapter, + status, + lease_id, + credential_cleanup_terminal_status, + sandbox_refresh_sandbox_id, + sandbox_refresh_claim, + sandbox_refresh_claim_expires_at, + agent_token_hash, + updated_at + ) VALUES (?, NULL, 'ready', ?, NULL, NULL, NULL, NULL, 'agent-token', ?) + `) + .run("IS-42", "sandbox:sandbox-1:terminal-1:autostart-v4", now); + sqlite + .prepare(` + INSERT INTO interactive_session_credential_policy_registrations ( + session_id, + sandbox_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + lookup_ids_json, + rollback_policies_json, + last_error, + updated_at + ) VALUES (?, ?, 'registering', ?, ?, ?, ?, ?, NULL, ?) + `) + .run( + "IS-42", + "sandbox-1", + "generation:stale", + "registration:stale", + now - 1, + JSON.stringify(["sandbox-1", "do-current"]), + JSON.stringify(rollback), + now - 1, + ); + const activeInsert = sqlite.prepare(` + INSERT INTO interactive_session_credential_policies ( + session_id, + sandbox_id, + lookup_id, + state, + registration_generation, + registration_claim, + registration_claim_expires_at, + created_at, + updated_at + ) VALUES (?, ?, ?, 'active', ?, NULL, NULL, ?, ?) + `); + activeInsert.run("IS-42", "sandbox-1", "sandbox-1", "generation:rollback", now, now); + activeInsert.run("IS-42", "sandbox-1", "do-legacy", "generation:advanced", now, now); + let rollbackCalls = 0; + + assert.equal(await activeSandboxCredentialPolicyGeneration(env, "IS-42", "sandbox-1"), null); + await scanCredentialPolicyCleanupPage( + env, + now, + async () => false, + "IS-42", + async () => { + rollbackCalls += 1; + }, + ); + + assert.equal(rollbackCalls, 0); + assert.deepEqual( + { + ...sqlite + .prepare(` + SELECT state, registration_claim, registration_claim_expires_at, last_error + FROM interactive_session_credential_policy_registrations + WHERE session_id = 'IS-42' AND sandbox_id = 'sandbox-1' + `) + .get(), + }, + { + state: "cleanup_pending", + registration_claim: null, + registration_claim_expires_at: null, + last_error: "sandbox credential policy generation advanced before rollback", + }, + ); +}); + +test("foreground rollback rechecks its exact claim before restoring policy", async () => { + const source = await readFile( + new URL("../src/worker/sandbox-credential-policy-registration-service.ts", import.meta.url), + "utf8", + ); + const start = source.indexOf( + "export async function restoreSandboxCredentialPolicyRollbackIfOwned", + ); + const end = source.indexOf("export async function registerSandboxCredentialPolicy", start); + const rollbackRecovery = source.slice(start, end); + + assert.ok( + rollbackRecovery.indexOf("renewSandboxCredentialPolicyRegistration(") < + rollbackRecovery.indexOf("restoreRollback("), + ); + assert.match(rollbackRecovery, /if \(!registrationExpiresAt\) return false/); +}); + test("credential-policy scan preserves live standalone and managed policies", () => { assert.equal( credentialPolicyScanRequiresCleanup( diff --git a/tests/sandbox-provisioning.test.ts b/tests/sandbox-provisioning.test.ts index ebeceb3e..6ad951b9 100644 --- a/tests/sandbox-provisioning.test.ts +++ b/tests/sandbox-provisioning.test.ts @@ -754,7 +754,7 @@ test("managed Sandbox commit fences activation and previous-policy cleanup", asy assert.ok(parameters.includes(claim.previousSandboxId)); }); -test("managed Sandbox refresh commit clears the claim before prior-policy cleanup", async () => { +test("managed Sandbox refresh commit preserves detached state before prior-policy cleanup", async () => { let statements: PreparedStatement[] = []; const expectedLeaseId = sandboxLeaseId(claim.lease); const env = runtimeEnv( @@ -764,7 +764,7 @@ test("managed Sandbox refresh commit clears the claim before prior-policy cleanu results: [ { lease_id: expectedLeaseId, - status: "ready", + status: "detached", credential_cleanup_terminal_status: null, agent_token_hash: claim.agentTokenHash, }, @@ -796,6 +796,8 @@ test("managed Sandbox refresh commit clears the claim before prior-policy cleanu assert.match(batchSql, /sandbox_refresh_claim_expires_at/); assert.match(batchSql, /agent_token_hash/); assert.match(batchSql, /interactive_session_credential_policies/); + assert.match(batchSql, /when \? = 'ready' and status in \('attached', 'detached'\)/i); + assert.match(batchSql, /then status/i); const parameters = statements.flatMap((statement) => statement.parameters); assert.ok(parameters.includes("cleanup_pending")); assert.ok(parameters.includes(claim.fence.claim)); diff --git a/tests/service-session-routes.test.ts b/tests/service-session-routes.test.ts index a86bbe9e..d28efb8a 100644 --- a/tests/service-session-routes.test.ts +++ b/tests/service-session-routes.test.ts @@ -304,3 +304,17 @@ test("service-session routes report missing reads and fall through on inexact re } assert.deepEqual(calls, []); }); + +test("service-session routes reject malformed encoded session ids with a client error", async () => { + const calls: string[] = []; + const value = request("POST", "/api/agent/interactive-sessions/%/events", {}); + + await assert.rejects(dispatch(value, calls), (error) => { + assert.equal( + typeof error === "object" && error && "status" in error ? error.status : undefined, + 400, + ); + return true; + }); + assert.deepEqual(calls, []); +}); diff --git a/tests/session-cleanup.test.ts b/tests/session-cleanup.test.ts index 770a0a11..8c9b9c51 100644 --- a/tests/session-cleanup.test.ts +++ b/tests/session-cleanup.test.ts @@ -118,6 +118,7 @@ test("cleanup candidates require finalized archives, no credentials, and no acti sessionReads += 1; assert.match(sql, /"terminal_finalize_pending" =/i); assert.match(sql, /interactive_session_credential_policies/i); + assert.match(sql, /interactive_session_credential_policy_registrations/i); assert.match(sql, /WITH RECURSIVE active_ancestor\(id\)/i); assert.match(sql, /archive\.session_updated_at = interactive_sessions\.updated_at/i); assert.match(sql, /events_key IS NOT NULL/i); @@ -163,6 +164,7 @@ test("finalized deletion claims and removes events, archive metadata, and sessio assert.equal(batch.length, 5); assert.match(batch[0]?.sql ?? "", /update "interactive_sessions"/i); assert.match(batch[0]?.sql ?? "", /interactive_session_credential_policies/i); + assert.match(batch[0]?.sql ?? "", /interactive_session_credential_policy_registrations/i); assert.match(batch[0]?.sql ?? "", /WITH RECURSIVE active_ancestor\(id\)/i); assert.match(batch[0]?.sql ?? "", /event_count =/i); assert.match(batch[0]?.sql ?? "", /events_key IS/i); diff --git a/tests/session-create-request.test.ts b/tests/session-create-request.test.ts index 8aa378c7..60041da0 100644 --- a/tests/session-create-request.test.ts +++ b/tests/session-create-request.test.ts @@ -76,6 +76,21 @@ test("session create requests enforce configured profiles and capability overlay }); test("session create requests fail before allocation when adapter routing is incomplete", () => { + assert.throws( + () => + resolveInteractiveSessionCreateRequest( + runtimeEnv({ + CRABBOX_RUNTIME_ADAPTER_URL: "https://adapter.example.test", + CRABBOX_RUNTIME_ADAPTER_URL_TEMPLATE: + "https://controller.example.test/adapters/{profile}", + CRABBOX_RUNTIME_ADAPTER_TOKEN: "adapter-token", + CRABBOX_RUNTIME_ADAPTER_NAMESPACE: "fleet", + }), + { repo: "openclaw/crabfleet" }, + { owner: "maintainer", createdBy: "maintainer" }, + ), + /runtime adapter URL or profile route template must be valid and unambiguous/, + ); assert.throws( () => resolveInteractiveSessionCreateRequest( diff --git a/tests/session-creation.test.ts b/tests/session-creation.test.ts index 95b67d7d..ca73e802 100644 --- a/tests/session-creation.test.ts +++ b/tests/session-creation.test.ts @@ -229,8 +229,8 @@ test("session creation owns normalization through decorated durable result", asy "insert", "supervise:IS-root:IS-2:100", "prepare", - "activate:IS-2:100:workspace-2", "request:IS-2:100", + "activate:IS-2:100:workspace-2", "provision:agent-token:unowned", "persist", "event", @@ -355,7 +355,7 @@ test("session creation returns a durable replay after a request reservation race assert.equal(provisioned, false); }); -test("session creation orders supervision, preparation, activation, evidence, and provisioning", async () => { +test("session creation records request evidence before activation and provisioning", async () => { const calls: string[] = []; const service = new InteractiveSessionCreationService( creationStore({ @@ -387,8 +387,8 @@ test("session creation orders supervision, preparation, activation, evidence, an assert.deepEqual(calls, [ "supervise:IS-1:IS-2:100", "prepare", - "activate:IS-2:100:workspace-2", "record:IS-2:100", + "activate:IS-2:100:workspace-2", "provision", ]); }); @@ -430,6 +430,65 @@ test("session creation rolls back failed preparation before returning the error" assert.deepEqual(calls, ["supervise", "prepare", "rollback:IS-2:100"]); }); +test("session creation rolls back when request recording fails", async () => { + const calls: string[] = []; + const failure = new Error("request event write failed"); + const service = new InteractiveSessionCreationService( + creationStore({ + enforceSupervision: async () => { + calls.push("supervise"); + }, + recordRequest: async () => { + calls.push("record"); + throw failure; + }, + rollbackReservation: async (sessionId, insertedAt) => { + calls.push(`rollback:${sessionId}:${insertedAt}`); + }, + activateReservation: async () => { + calls.push("activate"); + }, + }), + creationConfiguration, + ); + + await assert.rejects( + service.provision( + reservation, + async () => { + calls.push("prepare"); + }, + async () => { + calls.push("provision"); + }, + ), + failure, + ); + assert.deepEqual(calls, ["supervise", "prepare", "record", "rollback:IS-2:100"]); +}); + +test("session creation keeps ordinary reservations rollbackable through request recording", async () => { + const current = session({ id: "IS-2", status: "ready" }); + let preparationReservation: boolean | null = null; + let activated = false; + const service = new InteractiveSessionCreationService( + creationStore({ + insertReservation: async (input) => { + preparationReservation = input.preparationReservation; + }, + activateReservation: async () => { + activated = true; + }, + readSession: async () => current, + }), + creationConfiguration, + ); + + assert.equal(await service.create({ repo: "openclaw/crabfleet" }), current); + assert.equal(preparationReservation, true); + assert.equal(activated, true); +}); + test("session creation skips optional supervision, preparation, and activation", async () => { const calls: string[] = []; const service = new InteractiveSessionCreationService( @@ -623,6 +682,10 @@ test("session creation preserves the currently owned adapter provision", async ( { sessionId: "IS-2", adapterName: "runtime-v1", + adapterRegistration: { + profile: "default", + controlPlane: "https://controller.example", + }, sandboxLeasePrefix: "sandbox:", now: 150, }, @@ -652,8 +715,10 @@ test("session creation stops superseded adapter workspaces", async () => { const service = new InteractiveSessionCreationService( creationStore({ readSession: async () => current, - stopSupersededAdapter: async (sessionId, workspaceId, createPending, now) => { - calls.push(`stop:${sessionId}:${workspaceId}:${createPending}:${now}`); + stopSupersededAdapter: async (sessionId, workspaceId, registration, createPending, now) => { + calls.push( + `stop:${sessionId}:${workspaceId}:${registration?.profile}:${registration?.controlPlane}:${createPending}:${now}`, + ); }, }), creationConfiguration, @@ -664,6 +729,10 @@ test("session creation stops superseded adapter workspaces", async () => { { sessionId: "IS-2", adapterName: "runtime-v1", + adapterRegistration: { + profile: "default", + controlPlane: "https://controller.example", + }, sandboxLeasePrefix: "sandbox:", now: 150, }, @@ -680,7 +749,45 @@ test("session creation stops superseded adapter workspaces", async () => { ), current, ); - assert.deepEqual(calls, ["stop:IS-2:workspace-late:true:150"]); + assert.deepEqual(calls, ["stop:IS-2:workspace-late:default:https://controller.example:true:150"]); +}); + +test("session creation retains adapter registration when late provisioning loses ownership", async () => { + const current = session({ + id: "IS-2", + status: "ready", + adapter: "runtime-v1", + adapter_workspace_id: "workspace-current", + }); + const releases: unknown[][] = []; + const service = new InteractiveSessionCreationService( + creationStore({ + persistProvisionResult: async () => ({ + updated: false, + terminalStatus: null, + terminalAt: 101, + }), + readSession: async () => current, + stopSupersededAdapter: async (...args) => { + releases.push(args); + }, + }), + creationConfiguration, + ); + + assert.equal(await service.create({ repo: "openclaw/crabfleet" }), current); + assert.deepEqual(releases, [ + [ + "IS-2", + "workspace-2", + { + profile: "default", + controlPlane: "https://controller.example", + }, + false, + 100, + ], + ]); }); test("session creation cleans superseded sandbox ownership and rereads durability", async () => { @@ -706,6 +813,7 @@ test("session creation cleans superseded sandbox ownership and rereads durabilit { sessionId: "IS-2", adapterName: "runtime-v1", + adapterRegistration: null, sandboxLeasePrefix: "sandbox:", now: 150, }, @@ -730,6 +838,7 @@ test("session creation fails explicitly when durable ownership disappears", asyn { sessionId: "IS-2", adapterName: "runtime-v1", + adapterRegistration: null, sandboxLeasePrefix: "sandbox:", now: 150, }, diff --git a/tests/session-events.test.ts b/tests/session-events.test.ts index 777864f5..f53325fe 100644 --- a/tests/session-events.test.ts +++ b/tests/session-events.test.ts @@ -6,6 +6,7 @@ import { appendInteractiveSessionEventRecord, appendStructuredInteractiveSessionEventRecord, InteractiveSessionEventLedgerService, + persistInteractiveSessionEventRecord, structuredEventLedgerMaxBytes, structuredEventLedgerMaxCount, structuredEventPayloadMaxBytes, @@ -93,6 +94,24 @@ test("session event archive refresh remains best effort after durable persistenc assert.equal(persisted, true); }); +test("session event persistence can defer archive publication", async () => { + let statements: PreparedStatement[] = []; + const env = runtimeEnv((batch) => { + statements = batch; + }); + + await persistInteractiveSessionEventRecord(env, { + sessionId: "IS-1", + actor: "operator", + message: "interactive workspace requested", + now: 123, + }); + + assert.equal(statements.length, 2); + assert.match(statements[0]?.sql ?? "", /insert into "interactive_session_events"/i); + assert.match(statements[1]?.sql ?? "", /update "interactive_sessions"/i); +}); + test("structured session events canonicalize additive payloads and replay idempotently", async () => { let row: InteractiveSessionEventRow | undefined; const persistedPayloads: string[] = []; @@ -382,6 +401,26 @@ test("structured session events require bounded identifiers and a versioned obje { eventKey: "key", type: "action", message: "message", payload: [] }, { eventKey: "key", type: "action", message: "message", payload: {} }, { eventKey: "key", type: "action", message: "message", payload: { version: 0 } }, + { eventKey: "\ud800", type: "action", message: "message", payload: { version: 1 } }, + { eventKey: "key", type: "action", message: "\udfff", payload: { version: 1 } }, + { + eventKey: "key", + type: "action", + message: "message", + payload: { version: 1, value: "\ud800" }, + }, + { + eventKey: "key", + type: "action", + message: "message", + payload: { version: 1, value: Number.MAX_SAFE_INTEGER + 1 }, + }, + { + eventKey: "key", + type: "action", + message: "message", + payload: { version: 1, value: -0 }, + }, ]; for (const input of cases) { await assert.rejects( @@ -403,6 +442,41 @@ test("structured session events require bounded identifiers and a versioned obje assert.equal(persisted, false); }); +test("structured session events retain valid supplementary Unicode", async () => { + let payloadJson = ""; + const service = new InteractiveSessionEventLedgerService({ + async persist(event) { + payloadJson = event.payloadJson; + return { + inserted: true, + refreshArchive: false, + row: { + id: 1, + session_id: event.sessionId, + actor: event.actor, + event_key: event.eventKey, + event_type: event.type, + message: event.message, + payload_json: event.payloadJson, + created_at: event.now, + }, + }; + }, + async archive() {}, + }); + + await service.append({ + sessionId: "IS-1", + actor: "operator", + eventKey: "run:\u{1f980}", + type: "action", + message: "valid \u{1f980}", + payload: { version: 1, value: "\u{1f980}" }, + now: 123, + }); + assert.equal(payloadJson, '{"value":"\u{1f980}","version":1}'); +}); + test("structured session event payload budgets fail with controlled client errors", async () => { let persisted = false; const service = new InteractiveSessionEventLedgerService({ diff --git a/tests/session-grant-repository.test.ts b/tests/session-grant-repository.test.ts index a50f796a..0e3d8028 100644 --- a/tests/session-grant-repository.test.ts +++ b/tests/session-grant-repository.test.ts @@ -17,6 +17,7 @@ function runtimeEnv(options: { batches?: PreparedStatement[][]; mutations?: PreparedStatement[]; mutationChanges?: number; + batchChanges?: number[]; env?: Partial; }): RuntimeEnv { return { @@ -52,7 +53,9 @@ function runtimeEnv(options: { }, async batch(statements: unknown[]) { options.batches?.push(statements as PreparedStatement[]); - return []; + return statements.map((_, index) => ({ + meta: { changes: options.batchChanges?.[index] ?? 0 }, + })); }, } as unknown as D1Database, } as RuntimeEnv; @@ -206,23 +209,18 @@ test("grant upsert is atomically fenced to the live session revision", async () test("grant revocation atomically removes access and delegated control", async () => { const batches: PreparedStatement[][] = []; - const mutations: PreparedStatement[] = []; const repository = new InteractiveSessionGrantRepository( runtimeEnv({ - mutations, - mutationChanges: 1, batches, + batchChanges: [1, 1, 1, 1], }), ); assert.equal(await repository.revoke("IS-42", "proxy:collaborator@example.test", 2_000), true); - assert.equal(mutations.length, 1); - assert.match(mutations[0]?.sql ?? "", /^delete from "interactive_session_grants"/i); - assert.ok(mutations[0]?.parameters.includes("IS-42")); - assert.ok(mutations[0]?.parameters.includes("proxy:collaborator@example.test")); assert.equal(batches.length, 1); - assert.equal(batches[0]?.length, 3); + assert.equal(batches[0]?.length, 4); assert.match(batches[0]?.[0]?.sql ?? "", /^update "interactive_sessions"/i); + assert.match(batches[0]?.[0]?.sql ?? "", /exists/i); assert.match(batches[0]?.[0]?.sql ?? "", /"control_requested_by_subject" = \?/i); assert.doesNotMatch(batches[0]?.[0]?.sql ?? "", /"control_requested_by" in/i); assert.doesNotMatch(batches[0]?.[0]?.sql ?? "", /where[^]*"controller_subject" = \?/i); @@ -231,19 +229,21 @@ test("grant revocation atomically removes access and delegated control", async ( assert.doesNotMatch(batches[0]?.[1]?.sql ?? "", /"controller" in/i); assert.doesNotMatch(batches[0]?.[1]?.sql ?? "", /where[^]*"control_requested_by_subject" = \?/i); assert.match(batches[0]?.[2]?.sql ?? "", /"updated_at" = MAX\(updated_at \+ 1, \?\)/i); + assert.match(batches[0]?.[3]?.sql ?? "", /^delete from "interactive_session_grants"/i); assert.ok(batches[0]?.[0]?.parameters.includes("proxy:collaborator@example.test")); assert.ok(batches[0]?.[1]?.parameters.includes("proxy:collaborator@example.test")); assert.ok(batches[0]?.[2]?.parameters.includes(2_000)); + assert.ok(batches[0]?.[3]?.parameters.includes("IS-42")); + assert.ok(batches[0]?.[3]?.parameters.includes("proxy:collaborator@example.test")); }); test("grant revocation leaves sessions untouched when no grant exists", async () => { const batches: PreparedStatement[][] = []; - const mutations: PreparedStatement[] = []; const repository = new InteractiveSessionGrantRepository( - runtimeEnv({ mutations, mutationChanges: 0, batches }), + runtimeEnv({ batches, batchChanges: [0, 0, 0, 0] }), ); assert.equal(await repository.revoke("IS-42", "proxy:missing@example.test"), false); - assert.equal(mutations.length, 1); - assert.deepEqual(batches, []); + assert.equal(batches.length, 1); + assert.equal(batches[0]?.length, 4); }); diff --git a/tests/session-reconciliation.test.ts b/tests/session-reconciliation.test.ts index a2182d0b..28cb8e92 100644 --- a/tests/session-reconciliation.test.ts +++ b/tests/session-reconciliation.test.ts @@ -333,14 +333,18 @@ test("lost reconciliation ownership finalizes terminal rereads or releases super [], ); }, - async stopSuperseded(sessionId, workspaceId, createPending, now) { - released.push(`${sessionId}:${workspaceId}:${createPending}:${now}`); + async stopSuperseded(sessionId, workspaceId, registration, createPending, now) { + released.push( + `${sessionId}:${workspaceId}:${registration?.profile}:${registration?.controlPlane}:${createPending}:${now}`, + ); }, }), "runtime-adapter", ); await supersededService.reconcile(row, 150); - assert.deepEqual(released, [`${row.id}:workspace-1:false:200`]); + assert.deepEqual(released, [ + `${row.id}:workspace-1:cloudflare-sandbox:https://adapter.example:false:200`, + ]); }); test("reconciliation failures retain the claimed lifecycle fence", async () => { diff --git a/tests/session-terminal-finalization.test.ts b/tests/session-terminal-finalization.test.ts index 9a76588f..0eca38a8 100644 --- a/tests/session-terminal-finalization.test.ts +++ b/tests/session-terminal-finalization.test.ts @@ -1,7 +1,12 @@ import assert from "node:assert/strict"; import test from "node:test"; -import { terminalInteractiveSessionFinalizationMessage } from "../src/worker/session-terminal-finalization.ts"; +import { database } from "../src/worker/database.ts"; +import type { RuntimeEnv } from "../src/worker/env.ts"; +import { + terminalFinalizationClearPendingQuery, + terminalInteractiveSessionFinalizationMessage, +} from "../src/worker/session-terminal-finalization.ts"; test("terminal finalization messages preserve lifecycle and failure evidence", () => { assert.equal( @@ -29,3 +34,12 @@ test("terminal finalization messages preserve lifecycle and failure evidence", ( "interactive workspace failed after release", ); }); + +test("terminal finalization remains pending while any credential lifecycle row exists", () => { + const db = database({ DB: {} as D1Database } as RuntimeEnv); + const compiled = terminalFinalizationClearPendingQuery("IS-42", "stopped", true).compile(db); + + assert.match(compiled.sql, /interactive_session_credential_policies/i); + assert.match(compiled.sql, /interactive_session_credential_policy_registrations/i); + assert.match(compiled.sql, /NOT EXISTS/i); +}); diff --git a/tests/terminal-hub.test.ts b/tests/terminal-hub.test.ts index 7ed93128..3c620fd0 100644 --- a/tests/terminal-hub.test.ts +++ b/tests/terminal-hub.test.ts @@ -10,6 +10,15 @@ import { encodeSubscribePayload, encodeTerminalFrame, } from "@openclaw/libterminal/protocol"; +import { + attachGitHubActionsViewerProtocol, + encodeGitHubActionsRelayInputAcknowledgement, + encodeGitHubActionsRelayOutput, + githubActionsFramedRunnerCapability, + githubActionsGenerationFencedCapability, + notifyGitHubActionsViewers, + parseGitHubActionsRelayInput, +} from "../src/github-actions-runtime.ts"; import type { User } from "../src/worker/models.ts"; import { containerCapabilities, interactiveSession } from "../src/worker/session-model.ts"; import { TerminalHub, type TerminalHubDependencies } from "../src/worker/terminal-hub.ts"; @@ -22,6 +31,7 @@ class TestSocket { readonly sent: Array = []; readonly closed: Array<{ code?: number; reason?: string }> = []; accepted = false; + private attachment: unknown; private readonly listeners = new Map(); accept(): void { @@ -43,6 +53,14 @@ class TestSocket { this.readyState = WebSocket.CLOSED; } + serializeAttachment(attachment: unknown): void { + this.attachment = attachment; + } + + deserializeAttachment(): unknown { + return this.attachment; + } + emit(type: string, values: Record = {}): void { for (const listener of this.listeners.get(type) ?? []) { listener(Object.assign(new Event(type), values)); @@ -61,10 +79,54 @@ function frame(value: string | ArrayBuffer | ArrayBufferView | Blob) { return decoded; } +function relayInput(value: string | ArrayBuffer | ArrayBufferView | Blob) { + assert.ok(value instanceof ArrayBuffer); + const input = parseGitHubActionsRelayInput(value); + assert.ok(input); + return { + inputId: input.inputId, + generation: input.generation, + text: new TextDecoder().decode(input.payload), + }; +} + +function emitRelayAcknowledgement( + upstream: TestSocket, + inputId: string, + accepted: boolean, + generation?: string, +): void { + upstream.emit("message", { + data: encodeGitHubActionsRelayInputAcknowledgement({ + inputId, + accepted, + ...(generation ? { generation } : {}), + }), + }); +} + +function emitRelayEvent( + upstream: TestSocket, + type: "runner_connected" | "runner_disconnected" | "runner_waiting", + generation?: string, +): void { + const source = socket(); + attachGitHubActionsViewerProtocol( + source, + generation ? githubActionsGenerationFencedCapability : githubActionsFramedRunnerCapability, + ); + notifyGitHubActionsViewers([source], type, generation); + upstream.emit("message", { data: source.sent[0] }); +} + async function flushQueues(): Promise { await new Promise((resolve) => setImmediate(resolve)); } +async function waitForInputPayloads(): Promise { + await new Promise((resolve) => setTimeout(resolve, 100)); +} + const user: User = { subject: "github:42", login: "operator", @@ -84,6 +146,16 @@ const session = interactiveSession( }), ); +const githubActionsSession = interactiveSession( + sessionRow({ + adapter: null, + adapter_workspace_id: null, + capabilities_json: JSON.stringify(containerCapabilities), + runtime: "github_actions", + status: "ready", + }), +); + function dependencies( client: WebSocket, server: WebSocket, @@ -276,6 +348,7 @@ test("terminal hub routes multiplex frames and explicit output acknowledgements" ok: true, version: 2, multiplex: true, + inputAcknowledgements: true, }); server.emit("message", { @@ -337,6 +410,9 @@ test("terminal hub routes multiplex frames and explicit output acknowledgements" }); await flushQueues(); assert.deepEqual(new Uint8Array(upstream.sent.at(-1) as Uint8Array), inputPayload); + const inputAccepted = frame(server.sent.at(-1)!); + assert.equal(inputAccepted.type, TerminalMessageType.Event); + assert.deepEqual(decodeJsonPayload(inputAccepted.payload), { type: "input-accepted" }); server.emit("message", { data: encodeTerminalFrame({ @@ -350,6 +426,51 @@ test("terminal hub routes multiplex frames and explicit output acknowledgements" server.emit("close"); }); +test("duplicate subscriptions preserve the current input capability", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let upstreamOpens = 0; + const hub = new TerminalHub( + dependencies(client, server, upstream, { + inputGrant: () => async () => false, + async openUpstream() { + upstreamOpens += 1; + return { + socket: upstream, + outputAcknowledgements: true, + async markConnected() {}, + }; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + const subscribe = encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: session.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }); + + server.emit("message", { data: subscribe }); + await flushQueues(); + await flushQueues(); + server.emit("message", { data: subscribe }); + await flushQueues(); + await flushQueues(); + + assert.equal(upstreamOpens, 1); + assert.deepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "subscribed", + canInput: false, + }); + server.emit("close"); +}); + test("terminal hub publishes live controller downgrades and promotions", async () => { const client = socket(); const server = socket(); @@ -385,7 +506,11 @@ test("terminal hub publishes live controller downgrades and promotions", async ( }), }); await flushQueues(); - assert.equal(frame(server.sent.at(-1)!).type, TerminalMessageType.ControlRevoked); + assert.equal(frame(server.sent.at(-2)!).type, TerminalMessageType.ControlRevoked); + assert.deepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "input-rejected", + error: "terminal control is not granted", + }); assert.deepEqual(upstream.sent, []); canInput = true; @@ -397,8 +522,1477 @@ test("terminal hub publishes live controller downgrades and promotions", async ( }), }); await flushQueues(); - assert.equal(frame(server.sent.at(-1)!).type, TerminalMessageType.ControlGranted); + assert.equal( + server.sent.map((payload) => frame(payload).type).at(-2), + TerminalMessageType.ControlGranted, + ); assert.equal(new TextDecoder().decode(upstream.sent.at(-1) as Uint8Array), "allowed"); + assert.deepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "input-accepted", + }); + server.emit("close"); +}); + +test("terminal hub never acknowledges input after its upstream closes", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let resolvePayloads: ((payloads: Uint8Array[]) => void) | undefined; + const payloads = new Promise((resolve) => { + resolvePayloads = resolve; + }); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async inputPayloads() { + return payloads; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: session.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: session.id, + payload: new TextEncoder().encode("dropped"), + }), + }); + await flushQueues(); + upstream.close(1011, "upstream failed"); + resolvePayloads?.([new TextEncoder().encode("dropped")]); + await flushQueues(); + + assert.deepEqual(upstream.sent, []); + const messages = server.sent.map((payload) => frame(payload)); + assert.equal( + messages.some( + (message) => + message.type === TerminalMessageType.Event && + (decodeJsonPayload(message.payload) as { type?: string }).type === "input-accepted", + ), + false, + ); + assert.deepEqual(decodeJsonPayload(messages.at(-1)!.payload), { + error: "terminal upstream is not open", + }); + server.emit("close"); +}); + +test("GitHub Actions input waits for the correlated runner acknowledgement", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("steer\r"), + }), + }); + await flushQueues(); + + const input = relayInput(upstream.sent.at(-1)!); + assert.equal(input.text, "steer\r"); + assert.equal( + server.sent.some( + (payload) => + frame(payload).type === TerminalMessageType.Event && + (decodeJsonPayload(frame(payload).payload) as { type?: string }).type === "input-accepted", + ), + false, + ); + + emitRelayAcknowledgement(upstream, input.inputId, true); + await flushQueues(); + await flushQueues(); + + const accepted = frame(server.sent.at(-1)!); + assert.equal(accepted.type, TerminalMessageType.Event); + assert.deepEqual(decodeJsonPayload(accepted.payload), { type: "input-accepted" }); + server.emit("close"); +}); + +test("GitHub Actions acknowledgement timeout reports an ambiguous delivery outcome", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const detachedSessions: string[] = []; + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async markDetached(_user, sessionId) { + detachedSessions.push(sessionId); + }, + inputAcknowledgementTimeoutMs: 1, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("possibly-delivered"), + }), + }); + await new Promise((resolve) => setTimeout(resolve, 10)); + await flushQueues(); + + assert.deepEqual(upstream.closed, [{ code: 1011, reason: "input acknowledgement timed out" }]); + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload)) + .filter( + (message) => + typeof message === "object" && + message !== null && + "type" in message && + String(message.type).startsWith("input-"), + ); + assert.deepEqual(completions, [ + { + type: "input-delivery-unknown", + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ]); + upstream.emit("close", { code: 1011, reason: "input acknowledgement timed out" }); + await flushQueues(); + assert.deepEqual(detachedSessions, []); + server.emit("close"); +}); + +test("GitHub Actions falls back to raw relay input when viewer negotiation is absent", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async openUpstream() { + return { + socket: upstream, + inputAcknowledgements: false, + outputAcknowledgements: false, + async markConnected() {}, + }; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("legacy"), + }), + }); + await flushQueues(); + await flushQueues(); + + assert.equal(new TextDecoder().decode(upstream.sent.at(-1) as Uint8Array), "legacy"); + assert.deepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "input-accepted", + }); + server.emit("close"); +}); + +test("GitHub Actions relay rejection is request-scoped", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("dropped"), + }), + }); + await flushQueues(); + + const input = relayInput(upstream.sent.at(-1)!); + emitRelayAcknowledgement(upstream, input.inputId, false); + await flushQueues(); + await flushQueues(); + + const messages = server.sent.map((payload) => frame(payload)); + assert.equal( + messages.some( + (message) => + message.type === TerminalMessageType.Event && + (decodeJsonPayload(message.payload) as { type?: string }).type === "input-accepted", + ), + false, + ); + const rejected = messages.at(-1)!; + assert.equal(rejected.type, TerminalMessageType.Event); + assert.deepEqual(decodeJsonPayload(rejected.payload), { + type: "input-rejected", + error: "GitHub Actions runner did not accept terminal input", + }); + assert.equal(upstream.closed.length, 0); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("retry"), + }), + }); + await flushQueues(); + const retry = relayInput(upstream.sent.at(-1)!); + emitRelayAcknowledgement(upstream, retry.inputId, true); + await flushQueues(); + await flushQueues(); + + assert.deepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "input-accepted", + }); + server.emit("close"); +}); + +test("GitHub Actions input acknowledgements correlate overlapping payloads out of order", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async inputPayloads() { + return [new TextEncoder().encode("first"), new TextEncoder().encode("second")]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("input"), + }), + }); + await waitForInputPayloads(); + + const inputs = upstream.sent.map(relayInput); + assert.deepEqual( + inputs.map((input) => input.text), + ["first", "second"], + ); + assert.notEqual(inputs[0]!.inputId, inputs[1]!.inputId); + + emitRelayAcknowledgement(upstream, inputs[1]!.inputId, true); + emitRelayAcknowledgement(upstream, "stale-input-id", true); + await flushQueues(); + assert.equal( + server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Output).length, + 0, + ); + assert.notDeepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "input-accepted", + }); + + emitRelayAcknowledgement(upstream, inputs[0]!.inputId, true); + await flushQueues(); + await flushQueues(); + + assert.deepEqual(decodeJsonPayload(frame(server.sent.at(-1)!).payload), { + type: "input-accepted", + }); + server.emit("close"); +}); + +test("GitHub Actions serializes completion events for overlapping client inputs", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + for (const text of ["first", "second"]) { + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode(text), + }), + }); + } + await flushQueues(); + await flushQueues(); + + assert.equal(upstream.sent.length, 1); + const first = relayInput(upstream.sent[0]!); + assert.equal(first.text, "first"); + emitRelayAcknowledgement(upstream, first.inputId, false); + await flushQueues(); + await flushQueues(); + + assert.equal(upstream.sent.length, 2); + const second = relayInput(upstream.sent[1]!); + assert.equal(second.text, "second"); + emitRelayAcknowledgement(upstream, second.inputId, true); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string }) + .filter((message) => message.type === "input-accepted" || message.type === "input-rejected"); + assert.deepEqual( + completions.map((message) => message.type), + ["input-rejected", "input-accepted"], + ); + server.emit("close"); +}); + +test("terminal input queue rejects excess frames after earlier completions", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let releaseFirstInput: (() => void) | undefined; + const firstInput = new Promise((resolve) => { + releaseFirstInput = resolve; + }); + let inputCalls = 0; + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async inputPayloads(_subscription, _user, payload) { + inputCalls += 1; + if (inputCalls === 1) await firstInput; + return [payload]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: session.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + for (let index = 0; index < 128; index += 1) { + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: session.id, + payload: new Uint8Array([index]), + }), + }); + } + await flushQueues(); + await flushQueues(); + + assert.equal(inputCalls, 1); + assert.equal(upstream.sent.length, 0); + assert.equal( + server.sent.some( + (payload) => + (decodeJsonPayload(frame(payload).payload) as { error?: string }).error === + "terminal input backlog exceeded", + ), + false, + ); + + releaseFirstInput?.(); + await flushQueues(); + await flushQueues(); + await flushQueues(); + + assert.equal(inputCalls, 32); + assert.equal(upstream.sent.length, 32); + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type === "input-accepted" || message.type === "input-rejected"); + assert.equal(completions.length, 128); + assert.deepEqual(completions.slice(0, 32), Array(32).fill({ type: "input-accepted" })); + assert.deepEqual( + completions.slice(32), + Array(96).fill({ + type: "input-rejected", + error: "terminal input backlog exceeded", + }), + ); + server.emit("close"); +}); + +test("terminal input queue enforces the protocol-sized byte budget", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let releaseFirstInput: (() => void) | undefined; + const firstInput = new Promise((resolve) => { + releaseFirstInput = resolve; + }); + let inputCalls = 0; + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async inputPayloads(_subscription, _user, payload) { + inputCalls += 1; + if (inputCalls === 1) await firstInput; + return [payload]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: session.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + for (const payload of [new Uint8Array(9 * 1024 * 1024), new Uint8Array(8 * 1024 * 1024)]) { + server.emit("message", { + data: encodeTerminalFrame( + { + type: TerminalMessageType.Input, + sessionId: session.id, + payload, + }, + { maxFrameBytes: 16 * 1024 * 1024 }, + ), + }); + } + await flushQueues(); + await flushQueues(); + assert.equal(inputCalls, 1); + + releaseFirstInput?.(); + await flushQueues(); + await flushQueues(); + + assert.equal(inputCalls, 1); + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type === "input-accepted" || message.type === "input-rejected"); + assert.deepEqual(completions, [ + { type: "input-accepted" }, + { type: "input-rejected", error: "terminal input backlog exceeded" }, + ]); + server.emit("close"); +}); + +test("GitHub Actions framed output preserves control-shaped terminal bytes", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + const collision = encodeGitHubActionsRelayInputAcknowledgement({ + inputId: "stale-input-id", + accepted: true, + }); + upstream.emit("message", { data: encodeGitHubActionsRelayOutput(collision) }); + await flushQueues(); + + const output = frame(server.sent.at(-1)!); + assert.equal(output.type, TerminalMessageType.Output); + assert.deepEqual(output.payload, new Uint8Array(collision)); + server.emit("close"); +}); + +test("a stalled GitHub Actions acknowledgement does not block other sessions or ping", async () => { + const client = socket(); + const server = socket(); + const firstUpstream = socket(); + const secondUpstream = socket(); + const secondSession = interactiveSession( + sessionRow({ + id: "IS-actions-second", + adapter: null, + adapter_workspace_id: null, + capabilities_json: JSON.stringify(containerCapabilities), + runtime: "github_actions", + status: "ready", + }), + ); + const hub = new TerminalHub( + dependencies(client, server, firstUpstream, { + async readSession(_request, _user, id) { + return id === secondSession.id ? secondSession : githubActionsSession; + }, + async openUpstream(_request, _user, selectedSession) { + return { + socket: selectedSession.id === secondSession.id ? secondUpstream : firstUpstream, + outputAcknowledgements: true, + async markConnected() {}, + }; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + for (const selectedSession of [githubActionsSession, secondSession]) { + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: selectedSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + } + await flushQueues(); + await flushQueues(); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("stalled"), + }), + }); + await flushQueues(); + assert.equal(relayInput(firstUpstream.sent[0]!).text, "stalled"); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Ping, + sessionId: "", + payload: new TextEncoder().encode("still-live"), + }), + }); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: secondSession.id, + payload: new TextEncoder().encode("independent"), + }), + }); + await flushQueues(); + await flushQueues(); + + assert.equal(relayInput(secondUpstream.sent[0]!).text, "independent"); + assert.equal( + server.sent.some((payload) => { + const message = frame(payload); + return ( + message.type === TerminalMessageType.Pong && + new TextDecoder().decode(message.payload) === "still-live" + ); + }), + true, + ); + + const secondInput = relayInput(secondUpstream.sent[0]!); + emitRelayAcknowledgement(secondUpstream, secondInput.inputId, true); + await flushQueues(); + assert.equal( + server.sent.some((payload) => { + const message = frame(payload); + return ( + message.type === TerminalMessageType.Event && + message.sessionId === secondSession.id && + (decodeJsonPayload(message.payload) as { type?: string }).type === "input-accepted" + ); + }), + true, + ); + + server.emit("close"); +}); + +test("GitHub Actions send failure removes only its own acknowledgement waiter", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const send = upstream.send.bind(upstream); + let sendCount = 0; + upstream.send = (payload) => { + sendCount += 1; + if (sendCount === 2) throw new Error("runner disconnected"); + send(payload); + }; + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async inputPayloads() { + return [new TextEncoder().encode("delivered"), new TextEncoder().encode("failed")]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("input"), + }), + }); + await waitForInputPayloads(); + const input = relayInput(upstream.sent[0]!); + assert.equal(input.text, "delivered"); + emitRelayAcknowledgement(upstream, input.inputId, true); + await flushQueues(); + await flushQueues(); + + const rejected = frame(server.sent.at(-1)!); + assert.equal(rejected.type, TerminalMessageType.Event); + assert.deepEqual(decodeJsonPayload(rejected.payload), { + type: "input-rejected", + error: "terminal upstream send failed", + }); + server.emit("close"); +}); + +test("generation-fenced send failure completes its matching acknowledgement", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async openUpstream() { + return { + socket: upstream, + inputAcknowledgements: true, + inputGenerations: true, + outputAcknowledgements: false, + async markConnected() {}, + }; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + emitRelayEvent(upstream, "runner_connected", "generation-one"); + await flushQueues(); + await flushQueues(); + + upstream.send = () => { + throw new Error("runner disconnected"); + }; + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("input"), + }), + }); + await waitForInputPayloads(); + await flushQueues(); + + const rejected = frame(server.sent.at(-1)!); + assert.equal(rejected.type, TerminalMessageType.Event); + assert.deepEqual(decodeJsonPayload(rejected.payload), { + type: "input-rejected", + error: "terminal upstream send failed", + }); + assert.deepEqual(upstream.closed, []); + server.emit("close"); +}); + +test("GitHub Actions close reports forwarded pending input as delivery unknown", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async inputPayloads() { + return [new TextEncoder().encode("first"), new TextEncoder().encode("second")]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("input"), + }), + }); + await waitForInputPayloads(); + assert.equal(upstream.sent.length, 2); + upstream.emit("close", { code: 1011, reason: "runner disconnected" }); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type?.startsWith("input-")); + assert.deepEqual(completions, [ + { + type: "input-delivery-unknown", + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ]); + server.emit("close"); +}); + +test("GitHub Actions error reports forwarded pending input as delivery unknown", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("input"), + }), + }); + await waitForInputPayloads(); + assert.equal(upstream.sent.length, 1); + upstream.emit("error"); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type?.startsWith("input-")); + assert.deepEqual(completions, [ + { + type: "input-delivery-unknown", + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ]); + server.emit("close"); +}); + +test("GitHub Actions runner disconnect reports pending input as delivery unknown", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("pending"), + }), + }); + await flushQueues(); + + emitRelayEvent(upstream, "runner_disconnected"); + await flushQueues(); + await flushQueues(); + await flushQueues(); + + assert.deepEqual(upstream.closed, []); + const disconnectEvents = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload)); + assert.equal( + disconnectEvents.some( + (event) => + (event as { type?: string }).type === "input-delivery-unknown" && + (event as { error?: string }).error === + "terminal input delivery outcome is unknown; the runner may still complete it", + ), + true, + JSON.stringify(disconnectEvents), + ); + server.emit("close"); +}); + +test("generation-fenced runner disconnect marks only matching input as delivery unknown", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async openUpstream() { + return { + socket: upstream, + inputAcknowledgements: true, + inputGenerations: true, + initialRunnerGeneration: "generation-current", + outputAcknowledgements: false, + async markConnected() {}, + }; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("possibly delivered"), + }), + }); + await flushQueues(); + assert.equal(relayInput(upstream.sent.at(-1)!).generation, "generation-current"); + + emitRelayEvent(upstream, "runner_disconnected", "generation-stale"); + await flushQueues(); + await flushQueues(); + assert.equal( + server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string }) + .some((message) => message.type?.startsWith("input-")), + false, + ); + + emitRelayEvent(upstream, "runner_disconnected", "generation-current"); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type?.startsWith("input-")); + assert.deepEqual(completions, [ + { + type: "input-delivery-unknown", + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ]); + assert.deepEqual(upstream.closed, []); + server.emit("close"); +}); + +test("GitHub Actions runner replacement reports old input as unknown and accepts new input", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("old runner"), + }), + }); + await flushQueues(); + const oldInput = relayInput(upstream.sent.at(-1)!); + + emitRelayEvent(upstream, "runner_connected"); + await flushQueues(); + await flushQueues(); + await flushQueues(); + assert.deepEqual(upstream.closed, []); + const replacementEvents = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload)); + assert.equal( + replacementEvents.some( + (event) => + (event as { type?: string }).type === "input-delivery-unknown" && + (event as { error?: string }).error === + "terminal input delivery outcome is unknown; the runner may still complete it", + ), + true, + JSON.stringify(replacementEvents), + ); + + emitRelayAcknowledgement(upstream, oldInput.inputId, true); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("new runner"), + }), + }); + await flushQueues(); + await flushQueues(); + const newInput = relayInput(upstream.sent.at(-1)!); + assert.equal(newInput.text, "new runner"); + emitRelayAcknowledgement(upstream, newInput.inputId, true); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string }) + .filter( + (message) => message.type === "input-accepted" || message.type === "input-delivery-unknown", + ); + assert.deepEqual( + completions.map((message) => message.type), + ["input-delivery-unknown", "input-accepted"], + ); + assert.deepEqual(upstream.closed, []); + server.emit("close"); +}); + +test("queued runner replacement marks only old-generation acknowledgements unknown", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let releaseBlockedOutput: ((output: ArrayBuffer) => void) | undefined; + const blockedOutput = new Promise((resolve) => { + releaseBlockedOutput = resolve; + }); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async inputPayloads() { + return [new TextEncoder().encode("old runner"), new TextEncoder().encode("new runner")]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + + upstream.emit("message", { data: { arrayBuffer: () => blockedOutput } }); + const send = upstream.send.bind(upstream); + upstream.send = (data) => { + send(data); + if (upstream.sent.length === 1) emitRelayEvent(upstream, "runner_connected"); + }; + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("split input"), + }), + }); + await waitForInputPayloads(); + + assert.equal(upstream.sent.length, 2); + const oldInput = relayInput(upstream.sent[0]!); + const replacementInput = relayInput(upstream.sent[1]!); + releaseBlockedOutput?.(encodeGitHubActionsRelayOutput("blocked output")); + await flushQueues(); + await flushQueues(); + await flushQueues(); + + emitRelayAcknowledgement(upstream, oldInput.inputId, true); + await flushQueues(); + assert.equal( + server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string }) + .some((message) => message.type === "input-accepted" || message.type === "input-rejected"), + false, + ); + + emitRelayAcknowledgement(upstream, replacementInput.inputId, true); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type === "input-delivery-unknown"); + assert.deepEqual(completions, [ + { + type: "input-delivery-unknown", + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ]); + assert.deepEqual(upstream.closed, []); + server.emit("close"); +}); + +test("relay generations bind interleaved replacement input before lifecycle processing", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let releaseBlockedOutput: ((output: ArrayBuffer) => void) | undefined; + const blockedOutput = new Promise((resolve) => { + releaseBlockedOutput = resolve; + }); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + async readSession() { + return githubActionsSession; + }, + async openUpstream() { + return { + socket: upstream, + inputAcknowledgements: true, + inputGenerations: true, + outputAcknowledgements: false, + async markConnected() {}, + }; + }, + async inputPayloads() { + return [new TextEncoder().encode("old runner"), new TextEncoder().encode("new runner")]; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + await flushQueues(); + emitRelayEvent(upstream, "runner_connected", "generation-old"); + await flushQueues(); + await flushQueues(); + + upstream.emit("message", { data: { arrayBuffer: () => blockedOutput } }); + const send = upstream.send.bind(upstream); + upstream.send = (data) => { + send(data); + if (upstream.sent.length === 1) { + emitRelayEvent(upstream, "runner_connected", "generation-new"); + } + }; + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("split input"), + }), + }); + await waitForInputPayloads(); + + assert.equal(upstream.sent.length, 2); + const oldInput = relayInput(upstream.sent[0]!); + const replacementInput = relayInput(upstream.sent[1]!); + assert.equal(oldInput.generation, "generation-old"); + assert.equal(replacementInput.generation, "generation-new"); + + releaseBlockedOutput?.(encodeGitHubActionsRelayOutput("blocked output")); + await flushQueues(); + await flushQueues(); + await flushQueues(); + assert.equal( + server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string }) + .some((message) => message.type === "input-accepted" || message.type === "input-rejected"), + false, + ); + + emitRelayAcknowledgement(upstream, replacementInput.inputId, true, "generation-new"); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string; error?: string }) + .filter((message) => message.type === "input-delivery-unknown"); + assert.deepEqual(completions, [ + { + type: "input-delivery-unknown", + error: "terminal input delivery outcome is unknown; the runner may still complete it", + }, + ]); + assert.deepEqual(upstream.closed, []); + server.emit("close"); +}); + +test("viewer captures runner replacement during authorization setup", async () => { + const client = socket(); + const server = socket(); + const upstream = socket(); + let releaseView!: (allowed: boolean) => void; + const viewAllowed = new Promise((resolve) => { + releaseView = resolve; + }); + const hub = new TerminalHub( + dependencies(client, server, upstream, { + viewGrant: () => () => viewAllowed, + async openUpstream() { + return { + socket: upstream, + inputAcknowledgements: true, + inputGenerations: true, + initialRunnerGeneration: "generation-initial", + outputAcknowledgements: false, + async markConnected() {}, + }; + }, + }), + ); + await hub.open( + new Request("https://fleet.example/api/terminal/ws", { + headers: { upgrade: "websocket" }, + }), + user, + ); + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Subscribe, + sessionId: githubActionsSession.id, + payload: encodeSubscribePayload({ flags: 0, columns: 120, rows: 34 }), + }), + }); + await flushQueues(); + emitRelayEvent(upstream, "runner_connected", "generation-replacement"); + upstream.emit("message", { + data: encodeGitHubActionsRelayOutput("output-before-authorization"), + }); + releaseView(true); + await flushQueues(); + await flushQueues(); + + const messages = server.sent.map((payload) => frame(payload)); + assert.deepEqual( + messages + .filter((message) => message.type === TerminalMessageType.Event) + .map( + (message) => decodeJsonPayload(message.payload) as { type?: string; generation?: string }, + ) + .filter((message) => message.type === "runner_connected"), + [{ type: "runner_connected", generation: "generation-replacement" }], + ); + assert.equal( + messages.some((message) => message.type === TerminalMessageType.Output), + false, + ); + + server.emit("message", { + data: encodeTerminalFrame({ + type: TerminalMessageType.Input, + sessionId: githubActionsSession.id, + payload: new TextEncoder().encode("first input"), + }), + }); + await waitForInputPayloads(); + + const input = relayInput(upstream.sent.at(-1)!); + assert.equal(input.generation, "generation-replacement"); + emitRelayAcknowledgement(upstream, input.inputId, true, "generation-replacement"); + await flushQueues(); + await flushQueues(); + + const completions = server.sent + .map((payload) => frame(payload)) + .filter((message) => message.type === TerminalMessageType.Event) + .map((message) => decodeJsonPayload(message.payload) as { type?: string }) + .filter((message) => message.type === "input-accepted" || message.type === "input-rejected"); + assert.deepEqual( + completions.map((message) => message.type), + ["input-accepted"], + ); server.emit("close"); }); diff --git a/tests/terminal-multiplayer.test.ts b/tests/terminal-multiplayer.test.ts index f1609cd8..2c79abed 100644 --- a/tests/terminal-multiplayer.test.ts +++ b/tests/terminal-multiplayer.test.ts @@ -37,10 +37,14 @@ test("multiplayer input tracks interleaved writers on the shared session line", [encoder.encode("world")], ); - assert.equal( - text(multiplayerTerminalInputPayloadsForMode(state, secondUser, encoder.encode("\r"), true)), - '\x15 hello world\r', + const attributed = multiplayerTerminalInputPayloadsForMode( + state, + secondUser, + encoder.encode("\r"), + true, ); + assert.equal(attributed.length, 1); + assert.equal(text(attributed), '\x15 hello world\r'); }); test("multiplayer input attributes a final text fragment batched with enter", () => { @@ -51,10 +55,26 @@ test("multiplayer input attributes a final text fragment batched with enter", () [encoder.encode("hel")], ); - assert.equal( - text(multiplayerTerminalInputPayloadsForMode(state, user, encoder.encode("lo\r"), true)), - '\x15 hello\r', + const attributed = multiplayerTerminalInputPayloadsForMode( + state, + user, + encoder.encode("lo\r"), + true, ); + assert.equal(attributed.length, 1); + assert.equal(text(attributed), '\x15 hello\r'); +}); + +test("multiplayer input emits a complete attributed command atomically", () => { + const attributed = multiplayerTerminalInputPayloadsForMode( + newTerminalInputState(), + user, + encoder.encode("hello\r"), + true, + ); + + assert.equal(attributed.length, 1); + assert.equal(text(attributed), ' hello\r'); }); test("multiplayer input does not attribute text while a control sequence is pending", () => {