diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 620753d9a..d6f21aa0c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -311,6 +311,11 @@ jobs: TEST_SSH_USER: testuser TEST_SSH_KEY: ${{ github.workspace }}/testdata/ssh/test_key + - name: Run MinIO object-store integration test + run: make test-minio + env: + CGO_ENABLED: "1" + e2e: runs-on: ubuntu-latest steps: diff --git a/AGENTS.md b/AGENTS.md index a669d9271..f88ffef2b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -45,6 +45,9 @@ Instructions for autonomous coding agents working in this repository. details, and absolute user paths out of code, tests, fixtures, docs, commit messages, and pull request text. Run the private-data scrub before publishing. +- Keep agent-authored working specs and implementation plans under the ignored + `.superpowers/` directory. Never add them to tracked `docs/superpowers/` or + ship them in pull requests. - Keep pull request titles and descriptions synchronized with the current diff. - Do not post pull request or issue comments unless explicitly requested. diff --git a/Makefile b/Makefile index 5611bcf79..6ddac8a6a 100644 --- a/Makefile +++ b/Makefile @@ -36,7 +36,7 @@ AIR_BIN := $(shell if command -v air >/dev/null 2>&1; then command -v air; \ elif [ -x "$(GOPATH_FIRST)/bin/air" ]; then printf "%s" "$(GOPATH_FIRST)/bin/air"; \ fi) -.PHONY: build build-release install frontend frontend-dev dev check-air air-install desktop-dev desktop-build desktop-macos-app desktop-macos-dmg desktop-windows-installer desktop-linux-appimage desktop-app docs-install docs-build docs-serve docs-check docs-screenshots docs-assets-branch docs-generated-assets-branch docs-deploy-staging docs-deploy test test-short test-evalingest bench-backends bench-gate bench-gate-config test-postgres test-postgres-ci test-s3 postgres-up postgres-down test-ssh test-ssh-ci ssh-up ssh-down e2e e2e-duckdb vet lint lint-ci lint-golangci lint-golangci-ci nilaway nilaway-golangci-build lint-tools tidy clean release release-darwin-arm64 release-darwin-amd64 release-linux-amd64 install-hooks ensure-embed-dir pricing-snapshot sqlite-vec-header dev-snapshot help +.PHONY: build build-release install frontend frontend-dev dev check-air air-install desktop-dev desktop-build desktop-macos-app desktop-macos-dmg desktop-windows-installer desktop-linux-appimage desktop-app docs-install docs-build docs-serve docs-check docs-screenshots docs-assets-branch docs-generated-assets-branch docs-deploy-staging docs-deploy test test-short test-evalingest bench-backends bench-gate bench-gate-config test-postgres test-postgres-ci test-s3 test-minio postgres-up postgres-down test-ssh test-ssh-ci ssh-up ssh-down e2e e2e-duckdb vet lint lint-ci lint-golangci lint-golangci-ci nilaway nilaway-golangci-build lint-tools tidy clean release release-darwin-arm64 release-darwin-amd64 release-linux-amd64 install-hooks ensure-embed-dir pricing-snapshot sqlite-vec-header dev-snapshot help # Ensure go:embed has at least one file (no-op if frontend is built) ensure-embed-dir: @@ -351,6 +351,11 @@ test-postgres-ci: pricing-snapshot ensure-embed-dir test-s3: pricing-snapshot ensure-embed-dir CGO_ENABLED=1 go test -tags "fts5,s3test" -v ./internal/sync/... -run TestS3 -count=1 +# MinIO/S3 object-store integration test. testcontainers starts and tears down +# the MinIO container automatically, so this just needs a working Docker daemon. +test-minio: ensure-embed-dir + CGO_ENABLED=1 go test -tags "fts5,miniotest" -v ./internal/artifact/... -run MinIO -count=1 + # Start test SSH container ssh-up: docker compose -f docker-compose.test.yml up -d --build --wait sshd @@ -371,8 +376,9 @@ test-ssh: pricing-snapshot ensure-embed-dir ssh-up test-ssh-ci: pricing-snapshot ensure-embed-dir CGO_ENABLED=1 go test -tags "fts5,sshtest" -v ./internal/ssh/... -count=1 -# Run Playwright E2E tests -e2e: +# Run artifact sync and Playwright E2E tests +e2e: ensure-embed-dir + CGO_ENABLED=1 go test -tags "fts5,e2e" ./internal/e2e -v -count=1 cd frontend && npx playwright test # Run focused Playwright smoke tests against duckdb serve. @@ -547,6 +553,7 @@ help: @echo " bench-gate - Run the hot-path benchmarks CI gates PRs on" @echo " test-postgres - Run PostgreSQL integration tests" @echo " test-s3 - Run S3 discovery integration tests (Docker)" + @echo " test-minio - Run MinIO/S3 object-store integration test (needs Docker)" @echo " postgres-up - Start test PostgreSQL container" @echo " postgres-down - Stop test PostgreSQL container" @echo " test-ssh - Run SSH integration tests" diff --git a/README.md b/README.md index 97035495a..3c2042991 100644 --- a/README.md +++ b/README.md @@ -103,6 +103,27 @@ Use `--public-origin` (repeatable or comma-separated) to trust additional browser origins. If you expose the UI beyond loopback, also enable `--require-auth`. +## Local-First Artifact Sync + +Artifact sync is for a fully trusted personal fleet. It exchanges immutable +artifacts through a dedicated folder, HTTP peer, or S3-compatible target and +imports them into each machine's local SQLite archive: + +```bash +agentsview sync --init /path/to/agentsview-artifacts +agentsview sync /path/to/agentsview-artifacts +agentsview sync --watch /path/to/agentsview-artifacts +``` + +Do not sync the live SQLite database, `AGENTSVIEW_DATA_DIR`, or whole raw agent +directories. SQLite remains the authoritative local archive; the +`AGENTSVIEW_DATA_DIR/artifacts` directory is a private AgentsView-owned Docbank +vault, not an external sync target. Two intermittent machines only sync when +both can reach the same transport unless you provide a rendezvous such as a NAS +folder, cloud-synced folder, S3-compatible bucket, or always-on peer. See +[Trusted-Fleet Artifact Sync](docs/artifact-sync.md) for setup and safety +boundaries. + ## Docker The container image defaults to local `agentsview serve`. Set `PG_SERVE=1` to @@ -312,56 +333,56 @@ support is deprecated because current Amp releases may store threads server-side and leave only local stubs; agentsview can still parse historical local Amp thread JSON files. -| Agent | Session Directory | -| --------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| Aider | `/.aider.chat.history.md` (per repo; opt in with `AIDER_DIR` or `aider_dirs`) | -| Amp (deprecated) | `~/.local/share/amp/threads/` (historical local thread JSON only) | -| Antigravity | `~/.gemini/antigravity/` | -| Antigravity CLI | `~/.gemini/antigravity-cli/` (see note below) | -| Claude Code | `~/.claude/projects/` | -| OpenClaude | `~/.openclaude/projects/` | -| Claude Cowork | `~/Library/Application Support/Claude/local-agent-mode-sessions/` (macOS) | -| Codex | `~/.codex/sessions/` | -| Copilot CLI | `~/.copilot/` | -| Devin CLI | `~/.local/share/devin/` (Linux), `~/Library/Application Support/devin/` (macOS); point `DEVIN_DIR` / `devin_dirs` at the root that contains `cli/` | -| Cortex Code | `~/.snowflake/cortex/conversations/` | -| Cursor | `~/.cursor/projects/` | -| DeepSeek TUI | `~/.codewhale/sessions/`, `~/.deepseek/sessions/` | -| Forge | `~/.forge/` | -| Gemini CLI | `~/.gemini/` | -| gptme | `~/.local/share/gptme/logs/` | -| Grok | `~/.grok/sessions/` | -| Hermes Agent | `~/.hermes/sessions/` | -| iFlow | `~/.iflow/projects/` | -| Kilo | `~/.local/share/kilo/` | -| Kilo (legacy) | `~/Library/Application Support/Code/User/globalStorage/kilocode.kilo-code/` (macOS), `~/.config/Code/User/globalStorage/kilocode.kilo-code/` (Linux) | -| Kimi | `~/.kimi/sessions/` | -| Kiro CLI | `~/.kiro/sessions/cli/`, `~/.local/share/kiro-cli/` | -| Kiro IDE | `~/Library/Application Support/Kiro/` (macOS) | -| MiMoCode | `~/.local/share/mimocode/` | -| Mistral Vibe | `~/.vibe/logs/session/` | -| OpenClaw | `~/.openclaw/agents/` | -| OpenCode | `~/.local/share/opencode/` | -| OpenHands CLI | `~/.openhands/conversations/` | -| OhMyPi | `~/.omp/agent/sessions/` | -| Pi | `~/.pi/agent/sessions/` | -| Piebald | `~/.local/share/piebald/` | -| Posit Assistant | `~/.posit/assistant/workspaces/` | -| Positron Assistant | `~/Library/Application Support/Positron/User/` (macOS) | -| QClaw | `~/.qclaw/agents/` | -| Qoder | `~/.qoder/projects/`, `~/.qoderwork/projects/` | -| Qwen Code | `~/.qwen/projects/` | -| QwenPaw | `~/.copaw/workspaces/`, `~/.qwenpaw/workspaces/` | -| Reasonix | `~/.reasonix/`, `%APPDATA%\\reasonix\\` (Windows) | +| Agent | Session Directory | +| --------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| Aider | `/.aider.chat.history.md` (per repo; opt in with `AIDER_DIR` or `aider_dirs`) | +| Amp (deprecated) | `~/.local/share/amp/threads/` (historical local thread JSON only) | +| Antigravity | `~/.gemini/antigravity/` | +| Antigravity CLI | `~/.gemini/antigravity-cli/` (see note below) | +| Claude Code | `~/.claude/projects/` | +| OpenClaude | `~/.openclaude/projects/` | +| Claude Cowork | `~/Library/Application Support/Claude/local-agent-mode-sessions/` (macOS) | +| Codex | `~/.codex/sessions/` | +| Copilot CLI | `~/.copilot/` | +| Devin CLI | `~/.local/share/devin/` (Linux), `~/Library/Application Support/devin/` (macOS); point `DEVIN_DIR` / `devin_dirs` at the root that contains `cli/` | +| Cortex Code | `~/.snowflake/cortex/conversations/` | +| Cursor | `~/.cursor/projects/` | +| DeepSeek TUI | `~/.codewhale/sessions/`, `~/.deepseek/sessions/` | +| Forge | `~/.forge/` | +| Gemini CLI | `~/.gemini/` | +| gptme | `~/.local/share/gptme/logs/` | +| Grok | `~/.grok/sessions/` | +| Hermes Agent | `~/.hermes/sessions/` | +| iFlow | `~/.iflow/projects/` | +| Kilo | `~/.local/share/kilo/` | +| Kilo (legacy) | `~/Library/Application Support/Code/User/globalStorage/kilocode.kilo-code/` (macOS), `~/.config/Code/User/globalStorage/kilocode.kilo-code/` (Linux) | +| Kimi | `~/.kimi/sessions/` | +| Kiro CLI | `~/.kiro/sessions/cli/`, `~/.local/share/kiro-cli/` | +| Kiro IDE | `~/Library/Application Support/Kiro/` (macOS) | +| MiMoCode | `~/.local/share/mimocode/` | +| Mistral Vibe | `~/.vibe/logs/session/` | +| OpenClaw | `~/.openclaw/agents/` | +| OpenCode | `~/.local/share/opencode/` | +| OpenHands CLI | `~/.openhands/conversations/` | +| OhMyPi | `~/.omp/agent/sessions/` | +| Pi | `~/.pi/agent/sessions/` | +| Piebald | `~/.local/share/piebald/` | +| Posit Assistant | `~/.posit/assistant/workspaces/` | +| Positron Assistant | `~/Library/Application Support/Positron/User/` (macOS) | +| QClaw | `~/.qclaw/agents/` | +| Qoder | `~/.qoder/projects/`, `~/.qoderwork/projects/` | +| Qwen Code | `~/.qwen/projects/` | +| QwenPaw | `~/.copaw/workspaces/`, `~/.qwenpaw/workspaces/` | +| Reasonix | `~/.reasonix/`, `%APPDATA%\\reasonix\\` (Windows) | | RooCode | `~/Library/Application Support/Code/User/globalStorage/rooveterinaryinc.roo-cline/` (macOS), `~/.config/Code/User/globalStorage/rooveterinaryinc.roo-cline/` (Linux), `%APPDATA%\\Code\\User\\globalStorage\\rooveterinaryinc.roo-cline\\` (Windows) | -| VSCode Copilot | `~/Library/Application Support/Code/User/` (macOS) | -| Visual Studio Copilot | `%LOCALAPPDATA%\\Temp\\VSGitHubCopilotLogs\\traces\\` (Windows), `~/Library/Caches/VSGitHubCopilotLogs/traces/` (macOS), `~/.cache/VSGitHubCopilotLogs/traces/` (Linux) | -| Windsurf | `~/Library/Application Support/Windsurf/User/` (macOS), `~/.config/Windsurf/User/` (Linux), `%APPDATA%\\Windsurf\\User\\` (Windows) | -| Warp | `~/.warp/` (platform-dependent) | -| WorkBuddy | `~/.workbuddy/projects/` | -| ZCode | `~/.zcode/cli/db/`, `~/.zcode/cli/` | -| Zed | `~/Library/Application Support/Zed/` (macOS) | -| Zencoder | `~/.zencoder/sessions/` | +| VSCode Copilot | `~/Library/Application Support/Code/User/` (macOS) | +| Visual Studio Copilot | `%LOCALAPPDATA%\\Temp\\VSGitHubCopilotLogs\\traces\\` (Windows), `~/Library/Caches/VSGitHubCopilotLogs/traces/` (macOS), `~/.cache/VSGitHubCopilotLogs/traces/` (Linux) | +| Windsurf | `~/Library/Application Support/Windsurf/User/` (macOS), `~/.config/Windsurf/User/` (Linux), `%APPDATA%\\Windsurf\\User\\` (Windows) | +| Warp | `~/.warp/` (platform-dependent) | +| WorkBuddy | `~/.workbuddy/projects/` | +| ZCode | `~/.zcode/cli/db/`, `~/.zcode/cli/` | +| Zed | `~/Library/Application Support/Zed/` (macOS) | +| Zencoder | `~/.zencoder/sessions/` | Grok sessions are read from `summary.json` (title, timestamps, project), optional `signals.json` (token counters), and `chat_history.jsonl` when present @@ -467,10 +488,11 @@ or read them, and treats sidecars as untrusted structured input -- see ### Kilo vs Kilo (legacy). -*Kilo* is the OpenCode-based CLI/editor core (`~/.local/share/kilo/kilo.db`); -it covers both the Kilo CLI and the rebuilt Kilo Code VS Code extension (after -March 2026), which shares that same database. *Kilo (legacy)* is the legacy RooCode-derived -VS Code extension that wrote per-task JSON under `kilocode.kilo-code/tasks/`. +*Kilo* is the OpenCode-based CLI/editor core (`~/.local/share/kilo/kilo.db`); it +covers both the Kilo CLI and the rebuilt Kilo Code VS Code extension (after +March 2026), which shares that same database. *Kilo (legacy)* is the legacy +RooCode-derived VS Code extension that wrote per-task JSON under +`kilocode.kilo-code/tasks/`. ## PostgreSQL Sync diff --git a/SECURITY.md b/SECURITY.md index 3dc153e7e..b46d4396a 100644 --- a/SECURITY.md +++ b/SECURITY.md @@ -30,11 +30,11 @@ risk via documented flags and config. attacker. - Parser crashes, excessive resource use, or active-content injection triggered by content inside supported session files. Session files often contain - web/tool output that agentsview did not author, and defensive parsing of that - content is a security-relevant concern. + web/tool output that agentsview did not author, and defensive parsing of + that content is a security-relevant concern. - Inadvertent exposure of secrets that appear in transcripts. agentsview ships a - best-effort secret detector and redacts findings in the UI and CLI by default - (see [Secrets subsystem](#secrets-subsystem)). + best-effort secret detector and redacts findings in the UI and CLI by + default (see [Secrets subsystem](#secrets-subsystem)). ### Explicitly out of scope (today) @@ -63,6 +63,7 @@ risk via documented flags and config. | HTTP server → caller | Loopback-trusted; bearer-gated for `/api/` when `--require-auth` | Static assets are not gated. | | Browser → HTTP server | Host-header allowlist + CORS + CSP + X-Frame-Options enforced | DNS-rebinding, framing, and cross-origin defenses. | | agentsview → PostgreSQL (pg push) | TLS required for non-loopback hosts | Plaintext rejected unless `allow_insecure = true` is set explicitly. | +| agentsview → artifact HTTP peer | HTTPS required for non-loopback peers | Plaintext rejected unless `sync --allow-insecure` is explicit. | | agentsview → update endpoint | One-way outbound, opt-out | Disable with `--no-update-check`. | | agentsview → LiteLLM pricing | One-way outbound, on-demand | Public JSON fetched from GitHub raw; no session data sent. | @@ -70,8 +71,8 @@ risk via documented flags and config. - The local archive (SQLite + FTS5 index) stores indexed session data in plaintext. This includes assistant responses, user prompts, tool arguments, - command output, file contents fetched by agents, and any secrets that may have - been pasted into an agent session. + command output, file contents fetched by agents, and any secrets that may + have been pasted into an agent session. - File permissions follow the user's umask. agentsview does not chmod the data directory and does not encrypt at rest. - Treat the agentsview data directory with the same care you would treat your @@ -89,10 +90,10 @@ explicitly because "data stays on your machine" is the default but is not a complete description of the system once optional features are in use. - **Local UI / API.** The HTTP server binds to `127.0.0.1` by default. When - exposed beyond loopback, `--require-auth` should be enabled. Authentication is - a bearer token applied to `/api/` routes only; static assets remain ungated. - Browser-facing defenses (Host-header allowlist, CORS restrictions, CSP, - `X-Frame-Options: DENY`) are always on. See the CLI reference for token + exposed beyond loopback, `--require-auth` should be enabled. Authentication + is a bearer token applied to `/api/` routes only; static assets remain + ungated. Browser-facing defenses (Host-header allowlist, CORS restrictions, + CSP, `X-Frame-Options: DENY`) are always on. See the CLI reference for token configuration. - **PostgreSQL sync.** `agentsview pg push` exports the local archive to a user-supplied PostgreSQL instance. Non-loopback DSNs are rejected unless TLS @@ -102,8 +103,14 @@ complete description of the system once optional features are in use. its access controls. - **SSH remote sync.** agentsview can pull session archives from another machine over SSH. Authentication is whatever the user's SSH configuration provides. - Pulled files are parsed locally as untrusted data and merged into the unified - archive. + Pulled files are parsed locally as untrusted data and merged into the + unified archive. +- **Artifact HTTP peer sync.** `agentsview sync http(s)://...` exchanges bearer + credentials and archive artifacts with a user-supplied peer. Non-loopback + peers require HTTPS by default; loopback HTTP is allowed, while remote + plaintext requires the deliberate `--allow-insecure` opt-in and emits a + warning. Redirects are rejected. Received artifacts remain untrusted + structured input even when the peer is part of a trusted personal fleet. - **Imports.** Imported archives from other agentsview instances or third-party exports are treated as untrusted structured data, no different from session files written by local agents. @@ -125,9 +132,9 @@ The HTTP server applies the following defenses unconditionally: - **CORS restrictions.** Cross-origin API requests must come from an allowed origin or carry the bearer token; preflight handling is explicit. - **Content-Security-Policy.** A policy pinning the server's exact origin for - script/style/image/font/default-src is set on non-API responses. `connect-src` - is intentionally widened to allow the remote-server feature in the SPA; this - is a documented tradeoff. + script/style/image/font/default-src is set on non-API responses. + `connect-src` is intentionally widened to allow the remote-server feature in + the SPA; this is a documented tradeoff. - **X-Frame-Options: DENY.** Framing is disallowed on non-API responses. ## Secrets subsystem @@ -181,11 +188,11 @@ its own proposal. 1. **Multi-user machine support.** Is agentsview ever meant to run on a shared host, and if so what are the minimum hardening steps? 1. **`allow_insecure` UX.** Should setting `[pg] allow_insecure = true` require - an additional confirmation (e.g., a `--yes-really` flag) on first use, given - that it disables the only protection against plaintext PG egress? + an additional confirmation (e.g., a `--yes-really` flag) on first use, + given that it disables the only protection against plaintext PG egress? 1. **Deletion guarantees.** Should "permanent delete" grow into a stronger - erasure path (VACUUM, WAL checkpoint + truncate, mirror propagation to PG/SSH - targets), or should the docs simply make the current limits clearer? + erasure path (VACUUM, WAL checkpoint + truncate, mirror propagation to + PG/SSH targets), or should the docs simply make the current limits clearer? 1. **Secret detection scope.** Should the detector expand (more patterns, structured-secret types), should redacted-by-default extend to exports, and should there be a "scrub-on-import" pass? diff --git a/cmd/agentsview/cli.go b/cmd/agentsview/cli.go index caceeaf2d..d3ddefd09 100644 --- a/cmd/agentsview/cli.go +++ b/cmd/agentsview/cli.go @@ -298,7 +298,7 @@ func newOpenAPICommand() *cobra.Command { func newSyncCommand() *cobra.Command { var cfg SyncConfig cmd := &cobra.Command{ - Use: "sync", + Use: "sync [artifact-folder]", Short: "Sync session data without serving", Long: "Sync session data into the local database without starting the\n" + "HTTP server.\n\n" + @@ -309,12 +309,33 @@ func newSyncCommand() *cobra.Command { "exits non-zero if any configured host failed.\n\n" + "With --host, syncs only that host. A running local daemon may use a\n" + "matching configured remote_hosts entry and transport; otherwise,\n" + - "ad hoc --host sync uses your existing SSH configuration and requires\n" + + "ad hoc --host sync falls back to SSH.\n\n" + + "With an artifact-folder argument or --artifact-folder, sync also\n" + + "exchanges local-first immutable artifacts with that folder target.\n" + + "Artifact sync v1 is for a fully trusted personal fleet. Use a\n" + + "dedicated artifact share folder; do not point this at the\n" + + "agentsview data directory, raw agent directories, or the live\n" + + "SQLite database file and its WAL/SHM files.\n\n" + + "Use --init with an artifact folder on first setup to generate and\n" + + "persist this machine's artifact origin, backfill existing local\n" + + "sessions into the artifact store, exchange with the folder target,\n" + + "and import any peer artifacts already present. Two intermittent\n" + + "machines only sync while both can reach the same transport; use a\n" + + "NAS, cloud folder, object store, or always-on peer as a rendezvous\n" + + "when asynchronous convergence matters.\n\n" + + "Use --watch with an artifact folder to keep syncing. Watch mode\n" + + "runs an initial local sync and artifact exchange, coalesces file\n" + + "changes with --debounce, retries failed exchanges on later\n" + + "changes or --interval ticks, and performs a final best-effort\n" + + "exchange on shutdown. Combining --init with --watch publishes\n" + + "the first-run baseline on the first successful exchange, then\n" + + "keeps watching.\n\n" + + "Remote sync uses your existing SSH configuration and requires\n" + "key-based (passwordless) auth; it never prompts for a password.", GroupID: groupCore, SilenceUsage: true, - Args: cobra.NoArgs, - PreRunE: func(cmd *cobra.Command, _ []string) error { + Args: cobra.MaximumNArgs(1), + PreRunE: func(cmd *cobra.Command, args []string) error { if cfg.Host == "" { if cmd.Flags().Changed("user") || cmd.Flags().Changed("port") { @@ -323,9 +344,23 @@ func newSyncCommand() *cobra.Command { ) } } + if err := applySyncArtifactTarget(&cfg, args, cmd.Flags().Changed("artifact-folder")); err != nil { + return err + } + if err := validateSyncConfig(cfg); err != nil { + return err + } return nil }, Run: func(cmd *cobra.Command, args []string) { + if cfg.Watch { + runSyncWatch(cfg) + return + } + if cmd.Flags().Changed("debounce") || cmd.Flags().Changed("interval") { + fmt.Fprintln(os.Stderr, + "warning: --debounce and --interval have no effect without --watch") + } runSync(cfg) }, } @@ -333,6 +368,10 @@ func newSyncCommand() *cobra.Command { &cfg.Full, "full", false, "Force a full resync regardless of data version", ) + cmd.Flags().BoolVar( + &cfg.Init, "init", false, + "Initialize artifact sync with the folder target", + ) cmd.Flags().StringVar( &cfg.Host, "host", "", "SSH hostname for deprecated remote sync", @@ -345,6 +384,30 @@ func newSyncCommand() *cobra.Command { &cfg.Port, "port", 0, "SSH port for deprecated remote sync (default: 22)", ) + cmd.Flags().StringVar( + &cfg.ArtifactFolder, "artifact-folder", "", + "Exchange local-first sync artifacts with a folder, http(s) peer, or s3:// target", + ) + cmd.Flags().StringVar( + &cfg.Token, "token", "", + "Bearer token for an http(s) artifact peer target", + ) + cmd.Flags().BoolVar( + &cfg.AllowInsecure, "allow-insecure", false, + "Allow plaintext HTTP to a non-loopback artifact peer", + ) + cmd.Flags().BoolVar( + &cfg.Watch, "watch", false, + "Run artifact folder sync continuously, syncing on change plus a periodic floor", + ) + cmd.Flags().DurationVar( + &cfg.Debounce, "debounce", defaultWatchDebounce, + "Coalesce window after a change before artifact sync (--watch only)", + ) + cmd.Flags().DurationVar( + &cfg.Interval, "interval", defaultWatchInterval, + "Periodic floor artifact sync interval (--watch only)", + ) cmd.Flags().StringVar( &cfg.CPUProfile, "cpuprofile", "", "Write CPU profile to file (developer use)", @@ -362,6 +425,8 @@ func newSyncCommand() *cobra.Command { panic(err) } } + cmd.AddCommand(newSyncGCCommand()) + cmd.AddCommand(newSyncArtifactResetCommand()) return cmd } diff --git a/cmd/agentsview/cli_test.go b/cmd/agentsview/cli_test.go index 334a5bafd..6e312dec0 100644 --- a/cmd/agentsview/cli_test.go +++ b/cmd/agentsview/cli_test.go @@ -103,6 +103,13 @@ func TestDuckDBPushHelpShowsProjectFlags(t *testing.T) { } } +func TestSyncHelpShowsArtifactTransportSafetyFlag(t *testing.T) { + help, err := executeCommand(newRootCommand(), "sync", "--help") + require.NoError(t, err, "Execute") + assert.Contains(t, help, "--allow-insecure") + assert.Contains(t, help, "non-loopback artifact peer") +} + func TestPGStatusHelpShowsProjectFlags(t *testing.T) { help, err := executeCommand(newRootCommand(), "pg", "status", "--help") require.NoError(t, err, "Execute") @@ -144,6 +151,8 @@ func TestOpenAPICommandEmitsSpec(t *testing.T) { assert.Contains(t, spec.Paths["/api/v1/sessions"], "get") require.Contains(t, spec.Paths, "/api/v1/sessions/{id}/rename") assert.Contains(t, spec.Paths["/api/v1/sessions/{id}/rename"], "patch") + require.Contains(t, spec.Paths, "/api/v1/artifacts/peers") + assert.Contains(t, spec.Paths["/api/v1/artifacts/peers"], "get") } func TestServeCheckDataVersionRejectsNewerDatabase(t *testing.T) { @@ -362,7 +371,7 @@ func TestRootHelpDocumentsRemoteHosts(t *testing.T) { func TestSyncHelpMentionsConfiguredHosts(t *testing.T) { help, err := executeCommand(newRootCommand(), "sync", "--help") require.NoError(t, err, "Execute") - for _, want := range []string{"remote_hosts", "--host", "passwordless"} { + for _, want := range []string{"remote_hosts", "--host", "passwordless", "trusted personal fleet", "rendezvous"} { assert.Contains(t, help, want, "sync help missing %q", want) } } diff --git a/cmd/agentsview/main.go b/cmd/agentsview/main.go index 9e91bde83..340d785a0 100644 --- a/cmd/agentsview/main.go +++ b/cmd/agentsview/main.go @@ -20,6 +20,7 @@ import ( _ "time/tzdata" "github.com/spf13/cobra" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/config" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/parser" @@ -460,6 +461,25 @@ func runServe(cfg config.Config, opts serveOptions) { } cfg = preparedCfg + // Reconcile an already-adopted artifact origin so every origin lookup + // (recorder, peer import, folder sync) agrees: the config.toml origin is + // authoritative and overwrites a divergent DB sync-state value. Serve + // never creates an origin — a machine opts into artifact sync only via + // `sync --init`, a sync run, or an incoming peer exchange, and until then + // curation stays local with no metadata ledger writes. + if cfg.DataDir != "" && !database.ReadOnly() && cfg.ArtifactOriginID != "" { + if err := artifact.AdoptOrigin(database, cfg.ArtifactOriginID); err != nil { + fatal("reconcile artifact origin id: %v", err) + } + } + artifactRepository, err := openServeArtifactStore(ctx, cfg.DataDir) + if err != nil { + fatal("open artifact store: %v", err) + } + if err := recoverServeArtifactRepository(ctx, database, artifactRepository); err != nil { + fatal("recover artifact repository reset: %v", errors.Join(err, artifactRepository.Close())) + } + srvOpts := []server.Option{ server.WithVersion(server.VersionInfo{ Version: version, @@ -472,6 +492,7 @@ func runServe(cfg config.Config, opts serveOptions) { server.WithIdleTracker(idleTracker), server.WithHTTPRemoteCleanupRegistry(httpRemoteCleanupRegistry), server.WithPprof(opts.Pprof), + server.WithArtifactRepository(artifactRepository), } srvOpts = append(srvOpts, vectorServe.ServerOpts...) if src := newVectorPushSource(cfg); src != nil { @@ -588,6 +609,41 @@ func runServe(cfg config.Config, opts serveOptions) { } } +func openServeArtifactStore(ctx context.Context, dataDir string) (*artifact.Repository, error) { + repository, err := artifact.OpenRepository(ctx, dataDir) + if err != nil { + return nil, err + } + if err := repository.RecoverPacking(ctx); err != nil { + return nil, errors.Join(err, repository.Close()) + } + return repository, nil +} + +func recoverServeArtifactRepository( + ctx context.Context, database *db.DB, repository *artifact.Repository, +) error { + if database == nil || database.ReadOnly() { + return nil + } + origin, err := artifact.StoredOrigin(database) + if err != nil || origin == "" { + return err + } + _, recovered, err := artifact.RecoverRepositoryResetRepublish(ctx, database, repository, origin) + if err == nil && recovered { + repository.NotifyBatch(ctx) + } + if err != nil { + return err + } + coordinator := artifact.NewStoreImportCoordinator( + database, repository.Content(), origin, + ) + _, err = coordinator.Finalize(ctx) + return err +} + func runDeferredStartupSyncFallback( ctx context.Context, cfg config.Config, diff --git a/cmd/agentsview/main_test.go b/cmd/agentsview/main_test.go index 8c0b79e73..87e3de15d 100644 --- a/cmd/agentsview/main_test.go +++ b/cmd/agentsview/main_test.go @@ -22,6 +22,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/config" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/dbtest" @@ -48,6 +49,35 @@ func TestRuntimeWarningHelper(t *testing.T) { assert.Contains(t, logOutput.String(), "could not write daemon runtime record") } +func TestStartupBacklogServeOpensDocbankRepository(t *testing.T) { + dataDir := t.TempDir() + store, err := openServeArtifactStore(t.Context(), dataDir) + require.NoError(t, err) + require.NoError(t, store.Close()) + + artifactDir := filepath.Join(dataDir, "artifacts") + assert.FileExists(t, filepath.Join(artifactDir, "docbank.db")) +} + +func TestServeStartupRecoversPendingArtifactRepositoryReset(t *testing.T) { + dataDir := t.TempDir() + database := dbtest.OpenTestDBAt(t, filepath.Join(dataDir, "sessions.db")) + origin := "desktop-d4e5f6" + require.NoError(t, artifact.AdoptOrigin(database, origin)) + _, err := artifact.PrepareRepositoryResetRepublish( + t.Context(), database, dataDir, origin, + ) + require.NoError(t, err) + store, err := openServeArtifactStore(t.Context(), dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + + require.NoError(t, recoverServeArtifactRepository(t.Context(), database, store)) + _, pending, err := database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.False(t, pending) +} + func TestServeRuntimeRecordWriteFailureWarnsVisibleAfterSlowStartup(t *testing.T) { out, err := runServeRuntimeWarningHelper(t, true, 1200*time.Millisecond) require.NoError(t, err, string(out)) diff --git a/cmd/agentsview/pg_watch_loop.go b/cmd/agentsview/pg_watch_loop.go index 8d30a2a83..10875c535 100644 --- a/cmd/agentsview/pg_watch_loop.go +++ b/cmd/agentsview/pg_watch_loop.go @@ -5,6 +5,9 @@ import ( "log" "sync" "time" + + "go.kenn.io/agentsview/internal/config" + syncpkg "go.kenn.io/agentsview/internal/sync" ) // pushReason labels why a push was triggered, for logging. @@ -58,15 +61,23 @@ func newPushLoopWithLabel( label string, debounce, interval time.Duration, push func(context.Context, pushReason) error, +) (*pushLoop, *time.Ticker) { + return newNamedPushLoop(label, debounce, interval, push) +} + +func newNamedPushLoop( + label string, + debounce, interval time.Duration, + push func(context.Context, pushReason) error, ) (*pushLoop, *time.Ticker) { ticker := time.NewTicker(interval) return &pushLoop{ + label: label, debounce: debounce, dirty: make(chan struct{}, 1), floor: ticker.C, after: time.After, push: push, - label: label, flushTimeout: defaultFlushTimeout, }, ticker } @@ -169,3 +180,44 @@ func (l *pushLoop) restorePending(waiters []chan error) { l.pendingMu.Unlock() l.signalDirty() } + +type watchedSinkConfig struct { + AppConfig config.Config + Engine *syncpkg.Engine + Debounce time.Duration + Interval time.Duration + LogPrefix string + Push func(context.Context, pushReason) error +} + +func runWatchedSink(ctx context.Context, cfg watchedSinkConfig) { + loop, ticker := newNamedPushLoop( + cfg.LogPrefix, cfg.Debounce, cfg.Interval, cfg.Push, + ) + defer ticker.Stop() + + stopWatcher, openDispatch, unwatchedDirs, _ := startFileWatcher( + cfg.AppConfig, cfg.Engine, + func(callbackCtx context.Context, batch syncpkg.WatchBatch) error { + scope := func() watchRecoveryScope { + return probeWatchRecoveryScope(cfg.AppConfig) + } + if err := syncWatchBatch(callbackCtx, cfg.Engine, batch, scope); err != nil { + return err + } + loop.NotifyDirty() + return nil + }, + syncpkg.WatcherOptions{OnCoverageDegraded: loop.NotifyCoverageDegraded}, + ) + defer stopWatcher() + openDispatch() + if len(unwatchedDirs) > 0 { + log.Printf( + "%s: %d root(s) not watched; relying on the %s floor for coverage", + cfg.LogPrefix, len(unwatchedDirs), cfg.Interval, + ) + } + + loop.Run(ctx) +} diff --git a/cmd/agentsview/sync.go b/cmd/agentsview/sync.go index 6e75cebf7..687a29b18 100644 --- a/cmd/agentsview/sync.go +++ b/cmd/agentsview/sync.go @@ -12,11 +12,13 @@ import ( "io" "log" "net/http" + "net/url" "os" "strings" stdsync "sync" "time" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/config" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/parser" @@ -28,10 +30,19 @@ import ( // SyncConfig holds parsed CLI options for the sync command. type SyncConfig struct { - Full bool - Host string - User string - Port int + Full bool + Init bool + Watch bool + Debounce time.Duration + Interval time.Duration + Host string + User string + Port int + ArtifactFolder string + // Token is the bearer token used for an http(s):// artifact peer target. + Token string + // AllowInsecure permits plaintext HTTP to a non-loopback artifact peer. + AllowInsecure bool // CPUProfile, MemProfile, and Trace are hidden flags that capture a // pprof CPU profile, allocation snapshot, and runtime trace for the // sync pass. Empty strings disable each independently. @@ -40,17 +51,54 @@ type SyncConfig struct { Trace string } +func applySyncArtifactTarget(cfg *SyncConfig, args []string, flagChanged bool) error { + if len(args) == 0 { + return nil + } + if flagChanged { + return errors.New("artifact folder target cannot be provided both as an argument and --artifact-folder") + } + cfg.ArtifactFolder = args[0] + return nil +} + +func validateSyncConfig(cfg SyncConfig) error { + if cfg.Init && cfg.Host != "" { + return errors.New("--init cannot be combined with --host") + } + if cfg.Init && cfg.ArtifactFolder == "" { + return errors.New( + "--init requires an artifact folder target", + ) + } + if cfg.Watch && cfg.Host != "" { + return errors.New("--watch cannot be combined with --host") + } + if cfg.Watch && cfg.ArtifactFolder == "" { + return errors.New("--watch requires an artifact folder target") + } + if cfg.Host != "" && cfg.ArtifactFolder != "" { + // SSH remote sync (--host) returns before artifact sync runs, so a + // combined invocation would silently ignore the artifact target. + return errors.New("--host cannot be combined with an artifact target") + } + return nil +} + func runSync(cfg SyncConfig) { - if doSync(cfg) { + hadRemoteFailures, err := doSync(cfg) + if err != nil { + fatal("sync: %v", err) + } + if hadRemoteFailures { os.Exit(1) } } -// doSync performs the sync run and reports whether any configured -// remote host failed. It owns the deferred cleanup (profile stop, -// db close) so runSync can translate the result into a non-zero -// exit code without skipping that cleanup. -func doSync(cfg SyncConfig) (hadRemoteFailures bool) { +// doSync performs the sync run and reports whether any configured remote host +// failed. It owns the deferred cleanup (profile stop, db close) so runSync can +// translate the result into a non-zero exit code without skipping that cleanup. +func doSync(cfg SyncConfig) (hadRemoteFailures bool, err error) { appCfg, err := config.LoadMinimal() if err != nil { log.Fatalf("loading config: %v", err) @@ -93,9 +141,7 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { // includeLocal is true. A remote-only request still needs a newly // launched daemon to populate its existing local archive at startup. appCfg.SkipInitialSync = includeLocal - tr, err := ensureTransport( - &appCfg, transportIntentArchiveWrite, 0, - ) + tr, err := syncTransport(&appCfg, cfg) if err != nil { fatal("detecting daemon: %v", err) } @@ -113,7 +159,14 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { if err != nil { fatal("daemon remote sync: %v", err) } - return len(failures) > 0 + if cfg.ArtifactFolder != "" { + if _, err := runDaemonArtifactExchange( + context.Background(), tr, appCfg.AuthToken, cfg, + ); err != nil { + return false, fmt.Errorf("daemon artifact exchange: %w", err) + } + } + return len(failures) > 0, nil } if useDaemon { start := time.Now() @@ -154,7 +207,14 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { fatal("daemon sync: %v", err) } printSyncSummary(stats, start) - return false + if cfg.ArtifactFolder != "" { + if _, err := runDaemonArtifactExchange( + context.Background(), tr, appCfg.AuthToken, cfg, + ); err != nil { + return false, fmt.Errorf("daemon artifact exchange: %w", err) + } + } + return false, nil } // Read-only mirror daemons do not own the local SQLite // archive. Remote sync can still proceed through the direct @@ -162,7 +222,7 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { // writing imported remote sessions. } if tr.DirectReadOnly { - fatal( + return false, errors.New( "local daemon owns the SQLite archive but is not " + "responding; refusing to sync directly", ) @@ -177,12 +237,15 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { if cfg.Host != "" { runRemoteSync(appCfg, database, cfg) - return false + return false, nil } if len(appCfg.RemoteHosts) == 0 { runLocalSync(context.Background(), appCfg, database, cfg.Full) - return false + if cfg.ArtifactFolder != "" { + runArtifactFolderSync(appCfg, database, cfg.ArtifactFolder, cfg) + } + return false, nil } progress := newRemoteProgressPrinter(os.Stdout, time.Now) _, failures, blocked := runConfiguredLocalAndRemotesCLI( @@ -190,6 +253,9 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { cfg.Full, progress.Print, ) progress.Finish() + if cfg.ArtifactFolder != "" { + runArtifactFolderSync(appCfg, database, cfg.ArtifactFolder, cfg) + } reportRemoteFailures(failures) if blocked != nil { var pending *remotesync.PendingCleanupError @@ -199,11 +265,18 @@ func doSync(cfg SyncConfig) (hadRemoteFailures bool) { "sync: remote HTTP cleanup remains pending: %s\n", remotesync.FailureSummary(blocked), ) - return true + return true, nil } fatal("local sync: %v", blocked) } - return len(failures) > 0 + return len(failures) > 0, nil +} + +func syncTransport(appCfg *config.Config, cfg SyncConfig) (transport, error) { + if cfg.ArtifactFolder != "" { + return detectTransport(appCfg.DataDir, appCfg.AuthToken, 0) + } + return ensureTransport(appCfg, transportIntentArchiveWrite, 0) } func useDaemonForSync(tr transport) bool { @@ -317,6 +390,106 @@ func (p *remoteProgressPrinter) finishCurrent() { p.inPlace = false } +// resolveArtifactOrigin returns this machine's artifact origin. A config +// origin wins; otherwise an origin already stored in database sync state -- +// for example one minted by an incoming peer exchange on serve before the +// config ever initialized an origin -- is promoted into the config; otherwise +// a new origin is generated and persisted. Without the promotion step, CLI +// sync would generate a second origin that serve later adopts as +// authoritative, stranding metadata events published under the DB origin. +// The resolved config authority is reconciled back into database sync state so +// DB-only consumers such as direct PG push use the same canonical identity. +func resolveArtifactOrigin( + appCfg config.Config, database *db.DB, +) (string, error) { + var origin string + var err error + if appCfg.ArtifactOriginID == "" { + stored, readErr := artifact.StoredOrigin(database) + if readErr != nil { + return "", readErr + } + if stored != "" { + origin, err = appCfg.AdoptArtifactOriginID(stored) + } else { + origin, err = appCfg.EnsureArtifactOriginID() + } + } else { + origin, err = appCfg.EnsureArtifactOriginID() + } + if err != nil { + return "", err + } + if err := artifact.AdoptOrigin(database, origin); err != nil { + return "", fmt.Errorf("reconciling artifact origin in database: %w", err) + } + return origin, nil +} + +func runArtifactFolderSync( + appCfg config.Config, database *db.DB, target string, cfg SyncConfig, +) { + origin, err := resolveArtifactOrigin(appCfg, database) + if err != nil { + fatal("artifact sync origin: %v", err) + } + ctx := context.Background() + res, err := syncArtifactFolder( + ctx, appCfg, database, target, origin, artifactPeerToken(cfg), + cfg.AllowInsecure, cfg.Init, nil, + ) + if err != nil { + fatal("artifact sync: %v", err) + } + printArtifactSyncSummary(res, cfg.Init) +} + +// artifactPeerToken resolves the bearer token for an HTTP peer target. Tokens +// are never inferred from local server auth because an explicit peer URL may +// point at an untrusted endpoint. +func artifactPeerToken(cfg SyncConfig) string { + return cfg.Token +} + +func syncArtifactFolder( + ctx context.Context, + appCfg config.Config, + database *db.DB, + target string, + origin string, + token string, + allowInsecure bool, + baselineMetadata bool, + onDataChanged func(), +) (artifact.SyncResult, error) { + if !artifact.IsFolderTarget(target) && !artifact.IsHTTPTarget(target) && !artifact.IsObjectTarget(target) { + return artifact.SyncResult{}, fmt.Errorf( + "artifact sync supports local folder, http(s) peer, or s3:// object-store targets: %s", + target, + ) + } + return artifact.Sync(ctx, database, artifact.SyncOptions{ + DataDir: appCfg.DataDir, + Target: target, + Origin: origin, + Token: token, + AllowInsecure: allowInsecure, + BaselineMetadata: baselineMetadata, + OnDataChanged: onDataChanged, + }) +} + +func printArtifactSyncSummary(res artifact.SyncResult, init bool) { + label := "Artifact sync" + if init { + label = "Artifact sync initialized" + } + fmt.Printf( + "%s (%s): exported %d sessions, imported %d sessions / %d messages / %d metadata events\n", + label, res.Origin, res.ExportedSessions, res.ImportedSessions, res.ImportedMessages, res.ImportedMetadata, + ) +} + // syncLocalAndRemotes runs the local sync, then the configured // remote hosts. A local resync (forced via --full or an automatic // data-version resync) forces every remote sync full as well, so @@ -989,6 +1162,92 @@ func runDaemonSync( return parseDaemonSyncSSE(resp.Body, onProgress) } +func runDaemonArtifactExchange( + ctx context.Context, + tr transport, + authToken string, + cfg SyncConfig, +) (artifact.SyncResult, error) { + if err := artifact.ValidateSyncTarget(cfg.ArtifactFolder); err != nil { + return artifact.SyncResult{}, errors.New("invalid artifact exchange target") + } + endpoint, origin, err := daemonArtifactExchangeEndpoint(tr) + if err != nil { + return artifact.SyncResult{}, err + } + body, err := json.Marshal(struct { + Target string `json:"target"` + Token string `json:"token,omitempty"` + AllowInsecure bool `json:"allow_insecure,omitempty"` + BaselineMetadata bool `json:"baseline_metadata,omitempty"` + }{ + Target: cfg.ArtifactFolder, Token: cfg.Token, + AllowInsecure: cfg.AllowInsecure, BaselineMetadata: cfg.Init, + }) + if err != nil { + return artifact.SyncResult{}, err + } + req, err := http.NewRequestWithContext( + ctx, http.MethodPost, endpoint, bytes.NewReader(body), + ) + if err != nil { + return artifact.SyncResult{}, err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", origin) + if authToken != "" { + req.Header.Set("Authorization", "Bearer "+authToken) + } + client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }} + resp, err := client.Do(req) + if err != nil { + return artifact.SyncResult{}, fmt.Errorf("daemon artifact exchange request failed: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return artifact.SyncResult{}, fmt.Errorf("daemon artifact exchange returned HTTP %d", + resp.StatusCode) + } + var response struct { + Origin string `json:"origin"` + ExportedSessions int `json:"exported_sessions"` + ImportedSessions int `json:"imported_sessions"` + ImportedMessages int `json:"imported_messages"` + ImportedMetadata int `json:"imported_metadata"` + } + if err := json.NewDecoder(resp.Body).Decode(&response); err != nil { + return artifact.SyncResult{}, errors.New("decoding daemon artifact exchange response") + } + return artifact.SyncResult{ + Origin: response.Origin, ExportedSessions: response.ExportedSessions, + ImportedSessions: response.ImportedSessions, ImportedMessages: response.ImportedMessages, + ImportedMetadata: response.ImportedMetadata, + }, nil +} + +func daemonArtifactExchangeEndpoint(tr transport) (endpoint, origin string, err error) { + if tr.Runtime != nil && tr.Runtime.Host != "" && !isLoopbackHost(tr.Runtime.Host) { + return "", "", errors.New( + "artifact exchange requires a loopback-bound writable daemon; restart agentsview on localhost", + ) + } + base, err := url.Parse(tr.URL) + if err != nil || base == nil || base.Host == "" || + (base.Scheme != "http" && base.Scheme != "https") { + return "", "", errors.New("invalid local daemon endpoint") + } + if base.User != nil || base.RawQuery != "" || base.Fragment != "" || + !isLoopbackHost(base.Hostname()) { + return "", "", errors.New("artifact exchange requires a credential-free loopback daemon endpoint") + } + base.Path = strings.TrimRight(base.Path, "/") + "/api/v1/artifacts/exchange" + base.RawPath = "" + origin = (&url.URL{Scheme: base.Scheme, Host: base.Host}).String() + return base.String(), origin, nil +} + func runDaemonRemoteSync( ctx context.Context, tr transport, diff --git a/cmd/agentsview/sync_artifact_reset.go b/cmd/agentsview/sync_artifact_reset.go new file mode 100644 index 000000000..b1f423926 --- /dev/null +++ b/cmd/agentsview/sync_artifact_reset.go @@ -0,0 +1,176 @@ +package main + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + + "github.com/spf13/cobra" + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" +) + +type syncArtifactResetResponse struct { + artifact.RepositoryResetResult + ManualCleanup string `json:"manual_cleanup"` + ForeignArtifacts string `json:"foreign_artifacts"` +} + +type syncArtifactResetDependencies struct { + findDaemon func(string, ...string) *DaemonRuntime + localDaemonActive func(string, ...string) bool + openDirect func(context.Context, config.Config) (*db.DB, func(), error) + resetRepository func(context.Context, string, *db.DB, string, *artifact.Repository) (*artifact.Repository, artifact.RepositoryResetResult, error) +} + +func newSyncArtifactResetCommand() *cobra.Command { + return &cobra.Command{ + Use: "artifact-reset", + Short: "Move aside a failed artifact vault and rebuild it from SQLite", + SilenceUsage: true, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + return runSyncArtifactReset(cmd) + }, + } +} + +func runSyncArtifactReset(cmd *cobra.Command) error { + cfg, err := config.LoadMinimal() + if err != nil { + return fmt.Errorf("loading config: %w", err) + } + return runSyncArtifactResetWith(cmd, cfg, syncArtifactResetDependencies{}) +} + +func runSyncArtifactResetWith( + cmd *cobra.Command, + cfg config.Config, + deps syncArtifactResetDependencies, +) (retErr error) { + ctx := cmd.Context() + if ctx == nil { + ctx = context.Background() + } + if deps.findDaemon == nil { + deps.findDaemon = FindDaemonRuntime + } + if deps.localDaemonActive == nil { + deps.localDaemonActive = IsLocalDaemonActive + } + if deps.openDirect == nil { + deps.openDirect = openArtifactResetDirect + } + if deps.resetRepository == nil { + deps.resetRepository = artifact.ResetRepository + } + + if runtime := deps.findDaemon(cfg.DataDir, cfg.AuthToken); runtime != nil { + if runtime.ReadOnly { + return errors.New("a read-only daemon owns artifact access; stop it before reset") + } + response, err := runSyncArtifactResetDaemon(ctx, runtime, cfg.AuthToken) + if err != nil { + return err + } + printSyncArtifactReset(cmd.OutOrStdout(), response) + return nil + } + if deps.localDaemonActive(cfg.DataDir, cfg.AuthToken) { + return errors.New("the writable daemon owns the artifact vault but is not responding; refusing direct reset") + } + + database, cleanup, err := deps.openDirect(ctx, cfg) + if err != nil { + return fmt.Errorf("acquiring direct artifact reset ownership: %w", err) + } + if cleanup != nil { + defer cleanup() + } + origin := cfg.ArtifactOriginID + if origin == "" { + origin, err = artifact.StoredOrigin(database) + if err != nil { + return err + } + } + fresh, result, err := deps.resetRepository( + ctx, cfg.DataDir, database, origin, nil, + ) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, fresh.Close()) }() + printSyncArtifactReset(cmd.OutOrStdout(), syncArtifactResetResponse{ + RepositoryResetResult: result, + ManualCleanup: artifact.ArtifactResetManualCleanupWarning, + ForeignArtifacts: artifact.ArtifactResetForeignRelayWarning, + }) + return nil +} + +func openArtifactResetDirect( + ctx context.Context, cfg config.Config, +) (*db.DB, func(), error) { + database, lock, err := openWriteDB(ctx, cfg) + if err != nil { + return nil, nil, err + } + return database, func() { closeWriteDB(database, lock) }, nil +} + +func runSyncArtifactResetDaemon( + ctx context.Context, runtime *DaemonRuntime, authToken string, +) (syncArtifactResetResponse, error) { + endpoint, origin, err := daemonArtifactExchangeEndpoint(transportFromRuntime(runtime)) + if err != nil { + return syncArtifactResetResponse{}, err + } + endpoint = strings.TrimSuffix(endpoint, "/exchange") + "/reset" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, http.NoBody) + if err != nil { + return syncArtifactResetResponse{}, err + } + req.Header.Set("Origin", origin) + if authToken != "" { + req.Header.Set("Authorization", "Bearer "+authToken) + } + client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }} + resp, err := client.Do(req) + if err != nil { + return syncArtifactResetResponse{}, fmt.Errorf("daemon artifact reset request failed: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10)) + message := strings.TrimSpace(string(body)) + if message == "" { + message = http.StatusText(resp.StatusCode) + } + return syncArtifactResetResponse{}, fmt.Errorf( + "daemon artifact reset returned HTTP %d: %s", resp.StatusCode, message, + ) + } + var response syncArtifactResetResponse + if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&response); err != nil { + return syncArtifactResetResponse{}, errors.New("decoding daemon artifact reset response") + } + return response, nil +} + +func printSyncArtifactReset(w io.Writer, response syncArtifactResetResponse) { + fmt.Fprintln(w, "Artifact vault reset complete") + fmt.Fprintf(w, "Moved-aside diagnostic vault: %s\n", response.DiagnosticRoot) + fmt.Fprintf(w, "Fresh artifact vault: %s\n", response.VaultRoot) + fmt.Fprintf(w, "Republished %d local session(s); checkpoint sequence %d\n", + response.Export.ExportedSessions, response.Export.CheckpointSequence) + fmt.Fprintln(w, response.ManualCleanup) + fmt.Fprintln(w, response.ForeignArtifacts) +} diff --git a/cmd/agentsview/sync_artifact_reset_test.go b/cmd/agentsview/sync_artifact_reset_test.go new file mode 100644 index 000000000..1c4bb0714 --- /dev/null +++ b/cmd/agentsview/sync_artifact_reset_test.go @@ -0,0 +1,163 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "testing" + + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" +) + +func TestArtifactResetCommandIsExplicitAndTakesNoTarget(t *testing.T) { + cmd := newSyncArtifactResetCommand() + assert.Equal(t, "artifact-reset", cmd.Use) + require.Error(t, cmd.Args(cmd, []string{"some-vault"})) +} + +func TestArtifactResetCLIUsesAuthenticatedDaemonRoute(t *testing.T) { + var called bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true + assert.Equal(t, http.MethodPost, r.Method) + assert.Equal(t, "/api/v1/artifacts/reset", r.URL.Path) + assert.Equal(t, "Bearer daemon-secret", r.Header.Get("Authorization")) + w.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(w).Encode(syncArtifactResetResponse{ + RepositoryResetResult: artifact.RepositoryResetResult{ + VaultRoot: "/tmp/fresh", DiagnosticRoot: "/tmp/moved", + Export: artifact.ExportResult{ExportedSessions: 3, CheckpointSequence: 7}, + }, + ManualCleanup: artifact.ArtifactResetManualCleanupWarning, + ForeignArtifacts: artifact.ArtifactResetForeignRelayWarning, + })) + })) + defer server.Close() + cmd := &cobra.Command{} + var output bytes.Buffer + cmd.SetOut(&output) + direct := false + + err := runSyncArtifactResetWith(cmd, config.Config{ + DataDir: t.TempDir(), AuthToken: "daemon-secret", + }, syncArtifactResetDependencies{ + findDaemon: func(string, ...string) *DaemonRuntime { + return daemonRuntimeFromTestURL(t, server.URL) + }, + localDaemonActive: func(string, ...string) bool { return true }, + openDirect: func(context.Context, config.Config) (*db.DB, func(), error) { + direct = true + return nil, nil, errors.New("direct reset must not run") + }, + }) + + require.NoError(t, err) + assert.True(t, called) + assert.False(t, direct) + assert.Contains(t, output.String(), "/tmp/moved") + assert.Contains(t, output.String(), artifact.ArtifactResetManualCleanupWarning) + assert.Contains(t, output.String(), artifact.ArtifactResetForeignRelayWarning) +} + +func TestArtifactResetCLINeverFallsBackFromDaemonOwner(t *testing.T) { + for _, tt := range []struct { + name string + runtime *DaemonRuntime + active bool + }{ + {name: "read only daemon", runtime: &DaemonRuntime{ReadOnly: true}}, + {name: "unreachable writable daemon", active: true}, + } { + t.Run(tt.name, func(t *testing.T) { + direct := false + cmd := &cobra.Command{} + err := runSyncArtifactResetWith(cmd, config.Config{DataDir: t.TempDir()}, + syncArtifactResetDependencies{ + findDaemon: func(string, ...string) *DaemonRuntime { return tt.runtime }, + localDaemonActive: func(string, ...string) bool { return tt.active }, + openDirect: func(context.Context, config.Config) (*db.DB, func(), error) { + direct = true + return nil, nil, errors.New("unexpected direct reset") + }, + }) + require.Error(t, err) + assert.False(t, direct) + }) + } +} + +func TestArtifactResetCLIDirectModeUsesSQLiteWriteOwnership(t *testing.T) { + dataDir := t.TempDir() + cfg := config.Config{DataDir: dataDir, DBPath: filepath.Join(dataDir, "sessions.db")} + database, err := db.Open(cfg.DBPath) + require.NoError(t, err) + require.NoError(t, artifact.AdoptOrigin(database, "desktop-d4e5f6")) + startedAt := "2026-06-14T01:02:03Z" + require.NoError(t, database.UpsertSession(db.Session{ + ID: "local-session", Machine: "local", Agent: "codex", Project: "project-a", + StartedAt: &startedAt, CreatedAt: startedAt, + })) + database.Close() + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + require.NoError(t, repository.Close()) + cmd := &cobra.Command{} + var output bytes.Buffer + cmd.SetOut(&output) + + err = runSyncArtifactResetWith(cmd, cfg, syncArtifactResetDependencies{ + findDaemon: func(string, ...string) *DaemonRuntime { return nil }, + localDaemonActive: func(string, ...string) bool { return false }, + openDirect: openArtifactResetDirect, + resetRepository: artifact.ResetRepository, + }) + + require.NoError(t, err) + assert.Contains(t, output.String(), "Artifact vault reset complete") + assert.Contains(t, output.String(), artifact.ArtifactResetManualCleanupWarning) + assert.Contains(t, output.String(), artifact.ArtifactResetForeignRelayWarning) + assert.DirExists(t, filepath.Join(dataDir, "artifacts")) + matches, err := filepath.Glob(filepath.Join(dataDir, "artifacts.reset-*")) + require.NoError(t, err) + assert.Len(t, matches, 1) + reopened, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + originIterator, err := reopened.Content().Origins(t.Context()) + require.NoError(t, err) + origins, nextErr := originIterator.Next(t.Context(), 10) + require.ErrorIs(t, nextErr, io.EOF) + require.NoError(t, originIterator.Close()) + assert.Equal(t, []string{"desktop-d4e5f6"}, origins) + require.NoError(t, reopened.Close()) +} + +func TestArtifactResetCLIDirectLockFailureDoesNotTouchVault(t *testing.T) { + lockErr := errors.New("SQLite write ownership is held") + resetCalled := false + cmd := &cobra.Command{} + err := runSyncArtifactResetWith(cmd, config.Config{DataDir: t.TempDir()}, + syncArtifactResetDependencies{ + findDaemon: func(string, ...string) *DaemonRuntime { return nil }, + localDaemonActive: func(string, ...string) bool { return false }, + openDirect: func(context.Context, config.Config) (*db.DB, func(), error) { + return nil, nil, lockErr + }, + resetRepository: func(context.Context, string, *db.DB, string, *artifact.Repository) (*artifact.Repository, artifact.RepositoryResetResult, error) { + resetCalled = true + return nil, artifact.RepositoryResetResult{}, nil + }, + }) + + require.ErrorIs(t, err, lockErr) + assert.False(t, resetCalled) +} diff --git a/cmd/agentsview/sync_gc.go b/cmd/agentsview/sync_gc.go new file mode 100644 index 000000000..dcfb2ed97 --- /dev/null +++ b/cmd/agentsview/sync_gc.go @@ -0,0 +1,295 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strconv" + "strings" + "time" + + "github.com/spf13/cobra" + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" +) + +const ( + defaultArtifactGCGrace = 7 * 24 * time.Hour + defaultArtifactGCMaxObjects = 1024 + defaultArtifactGCMaxBytes = int64(256 << 20) +) + +// SyncGCConfig holds parsed logical-retention and physical-maintenance limits. +type SyncGCConfig struct { + Grace time.Duration + QuarantineGrace time.Duration + DryRun bool + MaxObjects int + MaxBytes int64 + TrashCursor string + GCCursor string + RepackCursor string +} + +type syncGCDependencies struct { + findDaemon func(string, ...string) *DaemonRuntime + localDaemonActive func(string, ...string) bool + openRepository func(context.Context, string) (*artifact.Repository, error) +} + +func newSyncGCCommand() *cobra.Command { + cfg := SyncGCConfig{ + Grace: defaultArtifactGCGrace, + QuarantineGrace: defaultArtifactGCGrace, + MaxObjects: defaultArtifactGCMaxObjects, + MaxBytes: defaultArtifactGCMaxBytes, + } + cmd := &cobra.Command{ + Use: "gc", + Short: "Retain logical artifacts and reclaim vault storage", + SilenceUsage: true, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + return runSyncGC(cmd, cfg) + }, + } + cmd.Flags().DurationVar(&cfg.Grace, "grace", defaultArtifactGCGrace, + "Minimum age before unreachable logical artifacts enter trash") + cmd.Flags().DurationVar(&cfg.QuarantineGrace, "quarantine-grace", defaultArtifactGCGrace, + "Minimum diagnostic retention for quarantined artifacts") + cmd.Flags().BoolVar(&cfg.DryRun, "dry-run", false, + "Report logical artifacts without trashing or physical reclamation") + cmd.Flags().IntVar(&cfg.MaxObjects, "max-objects", defaultArtifactGCMaxObjects, + "Maximum objects processed by each physical maintenance stage") + cmd.Flags().Int64Var(&cfg.MaxBytes, "max-bytes", defaultArtifactGCMaxBytes, + "Soft byte budget for blob garbage collection and repacking") + cmd.Flags().StringVar(&cfg.TrashCursor, "trash-cursor", "", + "Resume physical trash emptying from this cursor") + cmd.Flags().StringVar(&cfg.GCCursor, "gc-cursor", "", + "Resume physical garbage collection from this cursor") + cmd.Flags().StringVar(&cfg.RepackCursor, "repack-cursor", "", + "Resume physical repacking from this cursor") + return cmd +} + +func runSyncGC(cmd *cobra.Command, cfg SyncGCConfig) error { + appCfg, err := config.LoadMinimal() + if err != nil { + return fmt.Errorf("loading config: %w", err) + } + return runSyncGCWith(cmd, appCfg, cfg, syncGCDependencies{ + findDaemon: FindDaemonRuntime, + localDaemonActive: IsLocalDaemonActive, + openRepository: artifact.OpenRepository, + }) +} + +func runSyncGCWith( + cmd *cobra.Command, appCfg config.Config, cfg SyncGCConfig, deps syncGCDependencies, +) (retErr error) { + ctx := cmd.Context() + if ctx == nil { + ctx = context.Background() + } + if cfg.Grace < 0 || cfg.QuarantineGrace < 0 || + cfg.MaxObjects < 0 || cfg.MaxBytes < 0 { + return errors.New("artifact maintenance limits must not be negative") + } + if deps.findDaemon == nil { + deps.findDaemon = FindDaemonRuntime + } + if deps.localDaemonActive == nil { + deps.localDaemonActive = IsLocalDaemonActive + } + if deps.openRepository == nil { + deps.openRepository = artifact.OpenRepository + } + if runtime := deps.findDaemon(appCfg.DataDir, appCfg.AuthToken); runtime != nil { + if runtime.ReadOnly { + return errors.New("a read-only daemon owns artifact access; stop it before maintenance") + } + return runSyncGCDaemon(ctx, cmd, runtime, appCfg.AuthToken, cfg) + } + if deps.localDaemonActive(appCfg.DataDir, appCfg.AuthToken) { + return errors.New("the writable daemon owns the artifact vault but is not responding; refusing direct maintenance") + } + repository, err := deps.openRepository(ctx, appCfg.DataDir) + if err != nil { + return fmt.Errorf("opening artifact repository: %w", err) + } + defer func() { + retErr = errors.Join(retErr, repository.Close()) + }() + response, err := runSyncGCDirect(ctx, repository, cfg) + if err != nil { + return err + } + printArtifactMaintenanceSummary(cmd.OutOrStdout(), response, cfg) + return nil +} + +type syncGCMaintenanceResponse struct { + Logical artifact.GCResult `json:"logical"` + Physical struct { + Supported bool `json:"supported"` + Result artifact.PhysicalMaintenanceResult `json:"result"` + } `json:"physical"` +} + +func runSyncGCDirect( + ctx context.Context, repository *artifact.Repository, cfg SyncGCConfig, +) (syncGCMaintenanceResponse, error) { + var response syncGCMaintenanceResponse + maintenanceOpts := artifactMaintenanceOptions(cfg) + if err := artifact.ValidateArtifactMaintenanceOptions(maintenanceOpts); err != nil { + return response, fmt.Errorf("artifact physical maintenance: %w", err) + } + logical, err := artifact.GarbageCollect(ctx, artifact.GCOptions{ + Store: repository.Content(), Grace: cfg.Grace, QuarantineGrace: cfg.QuarantineGrace, + DryRun: cfg.DryRun, + }) + if err != nil { + return response, fmt.Errorf("artifact retention: %w", err) + } + response.Logical = logical + if cfg.DryRun { + return response, nil + } + response.Physical.Supported = true + response.Physical.Result, err = repository.RunMaintenance(ctx, maintenanceOpts) + if err != nil { + return response, fmt.Errorf("artifact physical maintenance: %w", err) + } + return response, nil +} + +func runSyncGCDaemon( + ctx context.Context, + cmd *cobra.Command, + runtime *DaemonRuntime, + authToken string, + cfg SyncGCConfig, +) error { + endpoint, origin, err := daemonArtifactExchangeEndpoint(transportFromRuntime(runtime)) + if err != nil { + return err + } + endpoint = strings.TrimSuffix(endpoint, "/exchange") + "/maintenance" + body, err := json.Marshal(struct { + Grace string `json:"grace"` + QuarantineGrace string `json:"quarantine_grace"` + MaxObjects int `json:"max_objects"` + MaxBytes int64 `json:"max_bytes"` + DryRun bool `json:"dry_run"` + TrashCursor string `json:"trash_cursor,omitempty"` + GCCursor string `json:"gc_cursor,omitempty"` + RepackCursor string `json:"repack_cursor,omitempty"` + }{ + Grace: cfg.Grace.String(), QuarantineGrace: cfg.QuarantineGrace.String(), + MaxObjects: cfg.MaxObjects, MaxBytes: cfg.MaxBytes, DryRun: cfg.DryRun, + TrashCursor: cfg.TrashCursor, GCCursor: cfg.GCCursor, RepackCursor: cfg.RepackCursor, + }) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", origin) + if authToken != "" { + req.Header.Set("Authorization", "Bearer "+authToken) + } + client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }} + resp, err := client.Do(req) + if err != nil { + return fmt.Errorf("daemon artifact maintenance request failed: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 64<<10)) + return fmt.Errorf("daemon artifact maintenance returned HTTP %d", resp.StatusCode) + } + var response syncGCMaintenanceResponse + if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&response); err != nil { + return errors.New("decoding daemon artifact maintenance response") + } + printArtifactMaintenanceSummary(cmd.OutOrStdout(), response, cfg) + return nil +} + +func artifactMaintenanceOptions(cfg SyncGCConfig) artifact.ArtifactMaintenanceOptions { + return artifact.ArtifactMaintenanceOptions{ + TrashGrace: cfg.Grace, + EmptyTrash: artifact.WorkBudget{MaxObjects: cfg.MaxObjects, Cursor: cfg.TrashCursor}, + GC: artifact.WorkBudget{ + MaxObjects: cfg.MaxObjects, MaxBytes: cfg.MaxBytes, Cursor: cfg.GCCursor, + }, + Repack: artifact.WorkBudget{ + MaxObjects: cfg.MaxObjects, MaxBytes: cfg.MaxBytes, Cursor: cfg.RepackCursor, + }, + } +} + +func printArtifactMaintenanceSummary( + w io.Writer, response syncGCMaintenanceResponse, cfg SyncGCConfig, +) { + result := response.Logical + action := "trashed" + count := result.Deleted + bytes := result.BytesDeleted + if result.DryRun { + action = "would trash" + count = result.Eligible + bytes = result.BytesEligible + } + fmt.Fprintf(w, + "Artifact retention: scanned %d origin(s), skipped %d unsafe origin(s), %s %d artifact(s) (%s)\n", + result.Origins, result.SkippedOrigins, action, count, formatBytes(bytes)) + if result.QuarantineSkipped { + fmt.Fprintln(w, "Artifact quarantine retention: unsupported by this store") + } + if response.Physical.Supported { + physical := response.Physical.Result + if !physical.EmptyTrash.More && !physical.GarbageCollect.More && !physical.Repack.More { + fmt.Fprintln(w, "Artifact physical maintenance complete") + return + } + fmt.Fprintln(w, "Artifact physical maintenance: more work remains") + fmt.Fprint(w, " agentsview sync gc") + printArtifactMaintenanceFlag(w, "grace", cfg.Grace.String()) + printArtifactMaintenanceFlag(w, "quarantine-grace", cfg.QuarantineGrace.String()) + printArtifactMaintenanceFlag(w, "max-objects", strconv.Itoa(cfg.MaxObjects)) + printArtifactMaintenanceFlag(w, "max-bytes", strconv.FormatInt(cfg.MaxBytes, 10)) + if cfg.DryRun { + fmt.Fprint(w, " --dry-run") + } + printArtifactMaintenanceResume(w, "trash", physical.EmptyTrash) + printArtifactMaintenanceResume(w, "gc", physical.GarbageCollect) + printArtifactMaintenanceResume(w, "repack", physical.Repack) + fmt.Fprintln(w) + } +} + +func printArtifactMaintenanceFlag(w io.Writer, name, value string) { + fmt.Fprintf(w, " --%s %s", name, shellQuoteArtifactValue(value)) +} + +func printArtifactMaintenanceResume(w io.Writer, stage string, result artifact.MaintenanceResult) { + if !result.More || result.NextCursor == "" { + return + } + fmt.Fprintf(w, " --%s-cursor %s", + stage, shellQuoteArtifactValue(result.NextCursor)) +} + +func shellQuoteArtifactValue(value string) string { + return "'" + strings.ReplaceAll(value, "'", "'\"'\"'") + "'" +} diff --git a/cmd/agentsview/sync_gc_test.go b/cmd/agentsview/sync_gc_test.go new file mode 100644 index 000000000..c5a97ad00 --- /dev/null +++ b/cmd/agentsview/sync_gc_test.go @@ -0,0 +1,190 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/docbank" +) + +func TestSyncGCCommandIsVaultMaintenanceNotFolderDeletion(t *testing.T) { + cmd := newSyncGCCommand() + assert.Equal(t, "gc", cmd.Use) + assert.Error(t, cmd.Args(cmd, []string{"shared-folder"}), + "maintenance must not accept an external folder path") + assert.NotNil(t, cmd.Flags().Lookup("grace")) + assert.NotNil(t, cmd.Flags().Lookup("quarantine-grace")) + assert.NotNil(t, cmd.Flags().Lookup("max-objects")) + maxBytes := cmd.Flags().Lookup("max-bytes") + require.NotNil(t, maxBytes) + assert.Contains(t, maxBytes.Usage, "garbage collection") + assert.Contains(t, maxBytes.Usage, "repacking") + assert.NotNil(t, cmd.Flags().Lookup("trash-cursor")) + assert.NotNil(t, cmd.Flags().Lookup("gc-cursor")) + assert.NotNil(t, cmd.Flags().Lookup("repack-cursor")) +} + +func TestRunSyncGCDaemonUsesAuthenticatedLoopbackRoute(t *testing.T) { + var requestBody map[string]any + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/api/v1/artifacts/maintenance", r.URL.Path) + assert.Equal(t, "Bearer secret", r.Header.Get("Authorization")) + require.NoError(t, json.NewDecoder(r.Body).Decode(&requestBody)) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"logical":{"origins":2,"deleted":3},"physical":{"supported":true,"result":{}}}`)) + })) + defer server.Close() + cmd := &cobra.Command{} + var output bytes.Buffer + cmd.SetOut(&output) + + err := runSyncGCDaemon(t.Context(), cmd, daemonRuntimeFromTestURL(t, server.URL), + "secret", SyncGCConfig{ + Grace: time.Hour + 1500*time.Nanosecond, + QuarantineGrace: 2*time.Hour + 2500*time.Nanosecond, + MaxObjects: 7, MaxBytes: 8, DryRun: true, + TrashCursor: "trash-next", GCCursor: "gc-next", RepackCursor: "repack-next", + }) + require.NoError(t, err) + assert.Equal(t, "1h0m0.0000015s", requestBody["grace"]) + assert.Equal(t, "2h0m0.0000025s", requestBody["quarantine_grace"]) + assert.NotContains(t, requestBody, "grace_seconds") + assert.NotContains(t, requestBody, "quarantine_grace_seconds") + assert.Equal(t, float64(7), requestBody["max_objects"]) + assert.Equal(t, float64(8), requestBody["max_bytes"]) + assert.Equal(t, true, requestBody["dry_run"]) + assert.Equal(t, "trash-next", requestBody["trash_cursor"]) + assert.Equal(t, "gc-next", requestBody["gc_cursor"]) + assert.Equal(t, "repack-next", requestBody["repack_cursor"]) + assert.Contains(t, output.String(), "scanned 2 origin(s)") +} + +func TestPrintArtifactMaintenanceSummaryReportsResumableStages(t *testing.T) { + response := syncGCMaintenanceResponse{} + response.Physical.Supported = true + response.Physical.Result.EmptyTrash = artifact.MaintenanceResult{ + More: true, NextCursor: "trash'next", + } + response.Physical.Result.GarbageCollect = artifact.MaintenanceResult{ + More: true, NextCursor: "gc-next", + } + response.Physical.Result.Repack = artifact.MaintenanceResult{ + More: true, NextCursor: "repack-next", + } + var output bytes.Buffer + + printArtifactMaintenanceSummary(&output, response, SyncGCConfig{}) + + assert.NotContains(t, output.String(), "physical maintenance complete") + assert.Contains(t, output.String(), "--trash-cursor 'trash'\"'\"'next'") + assert.Contains(t, output.String(), "--gc-cursor 'gc-next'") + assert.Contains(t, output.String(), "--repack-cursor 'repack-next'") +} + +func TestRunSyncGCDaemonResumeCommandPreservesEffectivePolicy(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{ + "logical":{}, + "physical":{"supported":true,"result":{ + "EmptyTrash":{"More":true,"NextCursor":"trash'next"}, + "GarbageCollect":{"More":true,"NextCursor":"gc-next"}, + "Repack":{"More":true,"NextCursor":"repack-next"} + }} + }`)) + })) + defer server.Close() + cmd := &cobra.Command{} + var output bytes.Buffer + cmd.SetOut(&output) + cfg := SyncGCConfig{ + Grace: time.Hour + 1500*time.Nanosecond, + QuarantineGrace: 2*time.Hour + 2500*time.Nanosecond, + MaxObjects: 7, MaxBytes: 8, + } + + err := runSyncGCDaemon(t.Context(), cmd, daemonRuntimeFromTestURL(t, server.URL), "", cfg) + + require.NoError(t, err) + assert.Contains(t, output.String(), + " agentsview sync gc --grace '1h0m0.0000015s'"+ + " --quarantine-grace '2h0m0.0000025s'"+ + " --max-objects '7' --max-bytes '8'"+ + " --trash-cursor 'trash'\"'\"'next'"+ + " --gc-cursor 'gc-next' --repack-cursor 'repack-next'\n") +} + +func TestRunSyncGCWithOwnershipNeverFallsBackFromDaemonOwner(t *testing.T) { + cases := []struct { + name string + runtime *DaemonRuntime + active bool + }{ + {name: "read-only owner", runtime: &DaemonRuntime{ReadOnly: true}}, + {name: "unreachable writable owner", active: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + opened := false + deps := syncGCDependencies{ + findDaemon: func(string, ...string) *DaemonRuntime { return tc.runtime }, + localDaemonActive: func(string, ...string) bool { return tc.active }, + openRepository: func(context.Context, string) (*artifact.Repository, error) { + opened = true + return nil, nil + }, + } + cmd := &cobra.Command{} + err := runSyncGCWith(cmd, config.Config{DataDir: t.TempDir()}, SyncGCConfig{}, deps) + require.Error(t, err) + assert.False(t, opened, "daemon ownership must prohibit direct vault fallback") + }) + } +} + +func TestRunSyncGCDirectUsesRepositoryLogicalAndPhysicalMaintenance(t *testing.T) { + dataDir := t.TempDir() + cmd := &cobra.Command{} + var output bytes.Buffer + cmd.SetOut(&output) + err := runSyncGCWith(cmd, config.Config{DataDir: dataDir}, SyncGCConfig{ + Grace: time.Hour, QuarantineGrace: time.Hour, + MaxObjects: 8, MaxBytes: 1 << 20, + }, syncGCDependencies{ + findDaemon: func(string, ...string) *DaemonRuntime { return nil }, + localDaemonActive: func(string, ...string) bool { return false }, + openRepository: artifact.OpenRepository, + }) + require.NoError(t, err) + assert.Contains(t, output.String(), "scanned 0 origin(s)") + assert.Contains(t, output.String(), "physical maintenance complete") +} + +func TestRunSyncGCDirectRejectsOversizedBudgetBeforeLogicalRetention(t *testing.T) { + _, err := runSyncGCDirect(t.Context(), nil, SyncGCConfig{ + MaxObjects: docbank.MaxMaintenanceObjects + 1, + }) + + assert.ErrorIs(t, err, artifact.ErrArtifactInvalid) +} + +func TestRunSyncGCDirectPreservesExplicitZeroBudgets(t *testing.T) { + repository, err := artifact.OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + response, err := runSyncGCDirect(t.Context(), repository, SyncGCConfig{}) + + require.NoError(t, err) + assert.True(t, response.Physical.Supported) +} diff --git a/cmd/agentsview/sync_test.go b/cmd/agentsview/sync_test.go index b7f7f3328..9c1eda3eb 100644 --- a/cmd/agentsview/sync_test.go +++ b/cmd/agentsview/sync_test.go @@ -780,8 +780,9 @@ token = "remote-token" } t.Cleanup(func() { prepareHTTPRebuildCLI = originalPrepare }) - hadRemoteFailures := doSync(SyncConfig{Full: true}) + hadRemoteFailures, err := doSync(SyncConfig{Full: true}) + require.NoError(t, err) assert.True(t, hadRemoteFailures) } @@ -815,8 +816,9 @@ token = "remote-token" httpRemoteCleanupRegistry = new(remotesync.CleanupRegistry) t.Cleanup(func() { httpRemoteCleanupRegistry = originalRegistry }) - hadRemoteFailures := doSync(SyncConfig{Full: true}) + hadRemoteFailures, err := doSync(SyncConfig{Full: true}) + require.NoError(t, err) assert.True(t, hadRemoteFailures) assert.Equal(t, 1, prepared.closed) assert.True(t, prepared.closeReleased) @@ -895,8 +897,9 @@ func TestDoSyncSingleHostFullStaysOnActiveArchivePath(t *testing.T) { } t.Cleanup(func() { runSSHRemoteSync = originalSSH }) - hadRemoteFailures := doSync(SyncConfig{Host: "one-box", Full: true}) + hadRemoteFailures, err := doSync(SyncConfig{Host: "one-box", Full: true}) + require.NoError(t, err) assert.False(t, hadRemoteFailures) assert.Equal(t, 0, prepareCalls) assert.Equal(t, 1, sshCalls) @@ -1530,12 +1533,81 @@ func TestDoSyncUsesDaemonRouteWhenWritableDaemonRunning(t *testing.T) { registerSyncRouteTestRuntime(t, env.DataDir, ts.URL) - hadFailures := doSync(SyncConfig{}) + hadFailures, err := doSync(SyncConfig{}) + require.NoError(t, err) require.False(t, hadFailures) assert.True(t, syncCalled) env.assertNoLocalDB(t) } +func TestDoSyncRoutesArtifactTargetThroughWritableDaemon(t *testing.T) { + env := newSyncCLIEnv(t) + target := t.TempDir() + + var syncCalled bool + var exchangeCalled bool + ts := daemonRouteTestServer(t, map[string]http.HandlerFunc{ + "/api/v1/sync": func(w http.ResponseWriter, r *http.Request) { + syncCalled = true + writeDoneSSE(t, w, agentsync.SyncStats{Synced: 7}) + }, + "/api/v1/artifacts/exchange": func(w http.ResponseWriter, r *http.Request) { + exchangeCalled = true + var request struct { + Target string `json:"target"` + } + require.NoError(t, json.NewDecoder(r.Body).Decode(&request)) + assert.Equal(t, target, request.Target) + w.Header().Set("Content-Type", "application/json") + _, err := io.WriteString(w, `{"origin":"daemon-a1b2c3","exported_sessions":3}`) + require.NoError(t, err) + }, + }) + + registerSyncRouteTestRuntime(t, env.DataDir, ts.URL) + + hadFailures, err := doSync(SyncConfig{ArtifactFolder: target}) + require.NoError(t, err) + assert.False(t, hadFailures) + assert.True(t, syncCalled) + assert.True(t, exchangeCalled) + env.assertNoLocalDB(t) +} + +func TestSyncTransportArtifactTargetUsesDirectDBWithoutAutoStartingDaemon(t *testing.T) { + env := newSyncCLIEnv(t) + var autostartCalled bool + stubStartBackgroundServeForTransport(t, func( + context.Context, *config.Config, time.Duration, + ) (*DaemonRuntime, error) { + autostartCalled = true + return nil, errors.New("unexpected daemon autostart") + }) + + appCfg := config.Config{ + DataDir: env.DataDir, + } + tr, err := syncTransport(&appCfg, SyncConfig{ + ArtifactFolder: t.TempDir(), + }) + require.NoError(t, err) + assert.Equal(t, transportDirect, tr.Mode) + assert.False(t, autostartCalled) +} + +func TestDoSyncRefusesDirectArtifactOwnershipWhenDaemonIsUnreachable(t *testing.T) { + env := newSyncCLIEnv(t) + writeUnreachableDaemonRuntime(t, env.DataDir, false) + + hadFailures, err := doSync(SyncConfig{ArtifactFolder: t.TempDir()}) + + require.Error(t, err) + assert.False(t, hadFailures) + assert.Contains(t, err.Error(), "daemon owns the SQLite archive") + assert.Contains(t, err.Error(), "refusing to sync directly") + env.assertNoLocalDB(t) +} + func TestDoSyncFullUsesDaemonResyncRoute(t *testing.T) { env := newSyncCLIEnv(t) @@ -1557,7 +1629,9 @@ func TestDoSyncFullUsesDaemonResyncRoute(t *testing.T) { var hadFailures bool out := captureStdout(t, func() { - hadFailures = doSync(SyncConfig{Full: true}) + var err error + hadFailures, err = doSync(SyncConfig{Full: true}) + require.NoError(t, err) }) require.False(t, hadFailures) assert.True(t, resyncCalled) @@ -1599,9 +1673,14 @@ func TestDoSyncPrintsStatusBeforeWaitingForDaemonStartup(t *testing.T) { _ = outFile.Close() }) - done := make(chan bool, 1) + type syncResult struct { + hadFailures bool + err error + } + done := make(chan syncResult, 1) go func() { - done <- doSync(SyncConfig{Full: true}) + hadFailures, err := doSync(SyncConfig{Full: true}) + done <- syncResult{hadFailures: hadFailures, err: err} }() <-startupEntered require.NoError(t, outFile.Sync()) @@ -1610,7 +1689,9 @@ func TestDoSyncPrintsStatusBeforeWaitingForDaemonStartup(t *testing.T) { assert.Contains(t, string(output), "Preparing full sync...") close(releaseStartup) - assert.False(t, <-done) + result := <-done + require.NoError(t, result.err) + assert.False(t, result.hadFailures) require.NoError(t, outFile.Close()) os.Stdout = oldStdout env.assertNoLocalDB(t) @@ -1631,8 +1712,9 @@ func TestDoSyncFullSkipsRedundantDaemonInitialSync(t *testing.T) { return &DaemonRuntime{Host: endpoint.Host, Port: endpoint.Port}, nil }) - hadFailures := doSync(SyncConfig{Full: true}) + hadFailures, err := doSync(SyncConfig{Full: true}) + require.NoError(t, err) assert.False(t, hadFailures) assert.True(t, skipInitialSync) env.assertNoLocalDB(t) @@ -1653,8 +1735,9 @@ func TestDoSyncSkipsRedundantDaemonInitialSync(t *testing.T) { return &DaemonRuntime{Host: endpoint.Host, Port: endpoint.Port}, nil }) - hadFailures := doSync(SyncConfig{}) + hadFailures, err := doSync(SyncConfig{}) + require.NoError(t, err) assert.False(t, hadFailures) assert.True(t, skipInitialSync) env.assertNoLocalDB(t) @@ -1673,8 +1756,9 @@ func TestDoSyncRemoteHostKeepsDaemonInitialLocalSync(t *testing.T) { return &DaemonRuntime{Host: endpoint.Host, Port: endpoint.Port}, nil }) - hadFailures := doSync(SyncConfig{Host: "host-a.example"}) + hadFailures, err := doSync(SyncConfig{Host: "host-a.example"}) + require.NoError(t, err) assert.False(t, hadFailures) assert.False(t, skipInitialSync, "remote-only request needs the daemon startup local sync") @@ -1740,6 +1824,57 @@ func TestRunDaemonSyncDetectsResyncRequired(t *testing.T) { } } +func TestRunDaemonArtifactExchangeDoesNotFollowRedirects(t *testing.T) { + var redirectedRequests atomic.Int32 + receiver := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + redirectedRequests.Add(1) + _, _ = io.Copy(io.Discard, r.Body) + w.WriteHeader(http.StatusOK) + })) + defer receiver.Close() + redirector := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, receiver.URL+"/captured", http.StatusTemporaryRedirect) + })) + defer redirector.Close() + + _, err := runDaemonArtifactExchange(t.Context(), transport{ + Mode: transportHTTP, URL: redirector.URL, + Runtime: &DaemonRuntime{Host: "127.0.0.1"}, + }, "daemon-secret", SyncConfig{ArtifactFolder: t.TempDir(), Token: "peer-secret"}) + require.Error(t, err) + assert.Contains(t, err.Error(), "HTTP 307") + assert.Zero(t, redirectedRequests.Load(), + "redirect receiver must get neither request body nor credentials") +} + +func TestRunDaemonArtifactExchangeRejectsLANRuntimeBeforeRequest(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + requests.Add(1) + })) + defer server.Close() + + _, err := runDaemonArtifactExchange(t.Context(), transport{ + Mode: transportHTTP, URL: server.URL, + Runtime: &DaemonRuntime{Host: "192.0.2.44"}, + }, "", SyncConfig{ArtifactFolder: t.TempDir()}) + require.ErrorContains(t, err, "loopback-bound writable daemon") + assert.Zero(t, requests.Load(), "LAN refusal must happen before any HTTP request") +} + +func TestRunDaemonArtifactExchangeRejectsSecretTargetWithoutDisclosure(t *testing.T) { + const secret = "do-not-disclose" + _, err := runDaemonArtifactExchange(t.Context(), transport{ + Mode: transportHTTP, URL: "http://127.0.0.1:1", + Runtime: &DaemonRuntime{Host: "127.0.0.1"}, + }, "daemon-"+secret, SyncConfig{ + ArtifactFolder: "https://user:" + secret + "@example.invalid/archive?token=" + secret + "#" + secret, + Token: "peer-" + secret, + }) + require.Error(t, err) + assert.NotContains(t, err.Error(), secret) +} + func TestDoSyncRemoteHostUsesDaemonRouteWhenWritableDaemonRunning(t *testing.T) { env := newSyncCLIEnv(t) @@ -1747,13 +1882,14 @@ func TestDoSyncRemoteHostUsesDaemonRouteWhenWritableDaemonRunning(t *testing.T) ts := remoteSyncRouteTestServer(t, handler) registerSyncRouteTestRuntime(t, env.DataDir, ts.URL) - hadFailures := doSync(SyncConfig{ + hadFailures, err := doSync(SyncConfig{ Host: "devbox", User: "alice", Port: 2222, Full: true, }) + require.NoError(t, err) require.False(t, hadFailures) assert.False(t, got.IncludeLocal) assert.True(t, got.Full) @@ -1784,10 +1920,12 @@ func TestDoSyncRemoteHostPrintsDaemonProgress(t *testing.T) { registerSyncRouteTestRuntime(t, env.DataDir, ts.URL) var hadFailures bool + var err error out := captureStdout(t, func() { - hadFailures = doSync(SyncConfig{Host: "devbox"}) + hadFailures, err = doSync(SyncConfig{Host: "devbox"}) }) + require.NoError(t, err) require.False(t, hadFailures) assert.Contains(t, out, "Running sync with remotes via daemon...") assert.Contains(t, out, "Resolving agent directories on devbox") @@ -2006,8 +2144,9 @@ user = "robot" ts := remoteSyncRouteTestServer(t, handler) registerSyncRouteTestRuntime(t, env.DataDir, ts.URL) - hadFailures := doSync(SyncConfig{}) + hadFailures, err := doSync(SyncConfig{}) + require.NoError(t, err) require.False(t, hadFailures) assert.True(t, got.IncludeLocal) require.Len(t, got.Hosts, 1) @@ -2228,7 +2367,6 @@ func TestRemoteFailureDisplaySanitizesHTTPErrors(t *testing.T) { }) } } - func TestRunHTTPRemoteSyncReachesMirrorPath(t *testing.T) { manifestRequests := 0 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -2263,3 +2401,236 @@ func TestRunHTTPRemoteSyncReachesMirrorPath(t *testing.T) { assert.Equal(t, 1, manifestRequests, "configured DataDir must route HTTP sync through the manifest/mirror path") } + +func TestNewSyncCommandRegistersArtifactFolderFlag(t *testing.T) { + cmd := newSyncCommand() + flag := cmd.Flags().Lookup("artifact-folder") + require.NotNil(t, flag) + assert.Equal(t, "", flag.DefValue) + assert.Contains(t, flag.Usage, "local-first sync artifacts") + initFlag := cmd.Flags().Lookup("init") + require.NotNil(t, initFlag) + assert.Equal(t, "false", initFlag.DefValue) + assert.Contains(t, initFlag.Usage, "Initialize artifact sync") + assert.Contains(t, cmd.Long, "do not point this at the") + assert.Contains(t, cmd.Long, "Use --init with an artifact folder") + watchFlag := cmd.Flags().Lookup("watch") + require.NotNil(t, watchFlag) + assert.Equal(t, "false", watchFlag.DefValue) + assert.Contains(t, watchFlag.Usage, "Run artifact folder sync continuously") + assert.Equal(t, defaultWatchDebounce.String(), cmd.Flags().Lookup("debounce").DefValue) + assert.Equal(t, defaultWatchInterval.String(), cmd.Flags().Lookup("interval").DefValue) + assert.Contains(t, cmd.Long, "Use --watch with an artifact folder") +} + +func TestApplySyncArtifactTargetUsesPositionalArgument(t *testing.T) { + cfg := SyncConfig{} + err := applySyncArtifactTarget(&cfg, []string{"/tmp/agentsview-share"}, false) + require.NoError(t, err) + assert.Equal(t, "/tmp/agentsview-share", cfg.ArtifactFolder) +} + +func TestApplySyncArtifactTargetRejectsArgumentAndFlag(t *testing.T) { + cfg := SyncConfig{ArtifactFolder: "/tmp/from-flag"} + err := applySyncArtifactTarget(&cfg, []string{"/tmp/from-arg"}, true) + require.Error(t, err) + assert.Contains(t, err.Error(), "both as an argument") + assert.Equal(t, "/tmp/from-flag", cfg.ArtifactFolder) +} + +func TestValidateSyncConfigInitRequiresArtifactFolder(t *testing.T) { + err := validateSyncConfig(SyncConfig{Init: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), "--init requires") +} + +func TestValidateSyncConfigInitRejectsHost(t *testing.T) { + err := validateSyncConfig(SyncConfig{ + Init: true, + Host: "remote", + ArtifactFolder: "/tmp/agentsview-share", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot be combined with --host") +} + +func TestValidateSyncConfigInitAllowsArtifactFolder(t *testing.T) { + err := validateSyncConfig(SyncConfig{ + Init: true, + ArtifactFolder: "/tmp/agentsview-share", + }) + require.NoError(t, err) +} + +func TestValidateSyncConfigWatchRequiresArtifactFolder(t *testing.T) { + err := validateSyncConfig(SyncConfig{Watch: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), "--watch requires") +} + +func TestValidateSyncConfigWatchRejectsHost(t *testing.T) { + err := validateSyncConfig(SyncConfig{ + Watch: true, + Host: "remote", + ArtifactFolder: "/tmp/agentsview-share", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot be combined with --host") +} + +func TestValidateSyncConfigWatchAllowsArtifactFolder(t *testing.T) { + err := validateSyncConfig(SyncConfig{ + Watch: true, + ArtifactFolder: "/tmp/agentsview-share", + Debounce: time.Second, + Interval: time.Minute, + }) + require.NoError(t, err) +} + +func TestValidateSyncConfigRejectsHostWithArtifactTarget(t *testing.T) { + err := validateSyncConfig(SyncConfig{ + Host: "remote", + ArtifactFolder: "/tmp/agentsview-share", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "--host cannot be combined with an artifact target") +} + +func TestValidateSyncConfigAllowsHostWithoutArtifactTarget(t *testing.T) { + require.NoError(t, validateSyncConfig(SyncConfig{Host: "remote"})) +} + +func TestArtifactPeerTokenDoesNotReuseLocalAuthToken(t *testing.T) { + got := artifactPeerToken(SyncConfig{ + ArtifactFolder: "https://peer.example.test", + }) + assert.Empty(t, got) + + got = artifactPeerToken(SyncConfig{ + ArtifactFolder: "https://peer.example.test", + Token: "peer-secret", + }) + assert.Equal(t, "peer-secret", got) +} + +func TestSyncArtifactFolderPlumbsInsecurePeerOptIn(t *testing.T) { + target, requests := insecureArtifactPeerTarget(t) + dataDir := t.TempDir() + database, err := db.Open(filepath.Join(dataDir, "sessions.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, database.Close()) }) + + _, err = syncArtifactFolder( + context.Background(), config.Config{DataDir: dataDir}, database, + target, "desk-a1b2c3", "", true, false, nil, + ) + + require.NoError(t, err) + assert.Positive(t, requests.Load(), + "one-shot sync must reach an explicitly allowed plaintext peer") +} + +func insecureArtifactPeerTarget(t *testing.T) (string, *atomic.Int32) { + t.Helper() + var host net.IP + addrs, err := net.InterfaceAddrs() + require.NoError(t, err) + for _, addr := range addrs { + ip, _, parseErr := net.ParseCIDR(addr.String()) + if parseErr == nil && ip.To4() != nil && !ip.IsLoopback() && !ip.IsUnspecified() { + host = ip + break + } + } + if host == nil { + t.Skip("no non-loopback IPv4 interface available for plaintext peer plumbing test") + } + listener, err := net.Listen("tcp4", "0.0.0.0:0") + require.NoError(t, err) + var requests atomic.Int32 + peer := &httptest.Server{ + Listener: listener, + Config: &http.Server{Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + w.Header().Set("Content-Type", "application/json") + switch { + case r.Method == http.MethodGet && strings.HasSuffix(r.URL.Path, "/origins"): + _, _ = io.WriteString(w, `{"origins":[]}`) + case r.Method == http.MethodGet && strings.HasSuffix(r.URL.Path, "/index"): + _, _ = io.WriteString(w, `{"origin":"desk-a1b2c3"}`) + case r.Method == http.MethodPost: + w.WriteHeader(http.StatusCreated) + default: + http.NotFound(w, r) + } + })}, + } + peer.Start() + t.Cleanup(peer.Close) + t.Setenv("HTTP_PROXY", "") + t.Setenv("HTTPS_PROXY", "") + t.Setenv("NO_PROXY", "*") + port := listener.Addr().(*net.TCPAddr).Port + return "http://" + net.JoinHostPort(host.String(), strconv.Itoa(port)), &requests +} + +func TestSyncGCHelpDocumentsMaintenanceFlags(t *testing.T) { + help, err := executeCommand(newRootCommand(), "sync", "gc", "--help") + require.NoError(t, err) + for _, want := range []string{ + "Retain logical artifacts", "--grace", "--quarantine-grace", + "--max-objects", "--max-bytes", "--dry-run", + } { + assert.Contains(t, help, want) + } +} + +func TestSyncGCRejectsExternalFolderTarget(t *testing.T) { + _, err := executeCommand(newRootCommand(), "sync", "gc", "https://example.test/artifacts") + require.Error(t, err) + assert.Contains(t, err.Error(), "unknown command") +} + +func TestSyncGCDryRunPrintsSummary(t *testing.T) { + t.Setenv("AGENTSVIEW_DATA_DIR", t.TempDir()) + out, err := executeCommand(newRootCommand(), "sync", "gc", "--dry-run") + require.NoError(t, err) + assert.Contains(t, out, "Artifact retention:") + assert.Contains(t, out, "would trash 0 artifact(s)") +} + +func TestArtifactFolderPusherPushExportsToTarget(t *testing.T) { + dataDir := t.TempDir() + target := t.TempDir() + database, err := db.Open(filepath.Join(dataDir, "sessions.db")) + require.NoError(t, err) + t.Cleanup(func() { database.Close() }) + + dbtest.SeedSession(t, database, "sess-1", "alpha", func(s *db.Session) { + s.MessageCount = 1 + s.UserMessageCount = 1 + }) + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + })) + + pusher := &artifactFolderPusher{ + appCfg: config.Config{DataDir: dataDir}, + database: database, + target: target, + origin: "desk-a1b2c3", + } + require.NoError(t, pusher.push(context.Background(), reasonChange)) + + checkpoints, err := filepath.Glob( + filepath.Join(target, "desk-a1b2c3", "checkpoints", "*.json"), + ) + require.NoError(t, err) + require.Len(t, checkpoints, 1) + manifests, err := filepath.Glob( + filepath.Join(target, "desk-a1b2c3", "manifests", "*.json.zst"), + ) + require.NoError(t, err) + assert.Len(t, manifests, 1) +} diff --git a/cmd/agentsview/sync_watch.go b/cmd/agentsview/sync_watch.go new file mode 100644 index 000000000..350207b43 --- /dev/null +++ b/cmd/agentsview/sync_watch.go @@ -0,0 +1,203 @@ +package main + +import ( + "context" + "fmt" + "log" + "os" + "os/signal" + "syscall" + + "github.com/gofrs/flock" + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/parser" + syncpkg "go.kenn.io/agentsview/internal/sync" + "go.kenn.io/kit/daemon" +) + +type artifactFolderPusher struct { + appCfg config.Config + database *db.DB + engine artifactWatchSyncer + target string + origin string + token string + allowInsecure bool + // baseline publishes first-run curation metadata (--init) on the next + // push. It stays set until a push succeeds so a failed initial exchange + // retries the baseline; AppendBaselineSnapshot skips already-covered + // fields, making the retry idempotent. Pushes run on a single loop + // goroutine, so no locking is needed. + baseline bool + onDataChanged func() +} + +type artifactWatchSyncer interface { + SyncAll(context.Context, syncpkg.ProgressFunc) syncpkg.SyncStats + FlushSignals() +} + +func newArtifactWatchEngine( + database *db.DB, appCfg config.Config, +) *syncpkg.Engine { + return syncpkg.NewEngine(database, syncpkg.EngineConfig{ + AgentDirs: appCfg.AgentDirs, + IncludeCwdPrefixes: appCfg.SyncIncludeCwdPrefixes, + Machine: "local", + BlockedResultCategories: appCfg.ResultContentBlockedCategories, + }) +} + +func (p *artifactFolderPusher) push( + ctx context.Context, reason pushReason, +) error { + if reason == reasonShutdown { + ctx = artifact.SuppressArtifactMaintenance(ctx) + } + if p.engine != nil { + // Startup already performed a full sync, and watcher change bursts have + // already applied their targeted paths. The periodic floor covers roots + // that could not be watched, while shutdown discovers events still held + // in the watcher's batching window before the final export. + if reason == reasonInterval || reason == reasonShutdown { + p.engine.SyncAll(ctx, nil) + } + // Export reads session rows outside a sync operation; flush + // debounced signal recomputes so manifests carry current signals. + p.engine.FlushSignals() + } + res, err := syncArtifactFolder( + ctx, p.appCfg, p.database, p.target, p.origin, p.token, + p.allowInsecure, p.baseline, p.onDataChanged, + ) + if err != nil { + return err + } + p.baseline = false + log.Printf( + "artifact watch: exported %d sessions, imported %d sessions, %d messages, %d metadata events (%s)", + res.ExportedSessions, res.ImportedSessions, res.ImportedMessages, + res.ImportedMetadata, reason, + ) + return nil +} + +// runSyncWatch runs continuous artifact folder sync: an initial local sync and +// artifact exchange, then debounced file-change exchanges and a periodic floor. +func runSyncWatch(cfg SyncConfig) { + appCfg, err := config.LoadMinimal() + if err != nil { + log.Fatalf("loading config: %v", err) + } + if err := os.MkdirAll(appCfg.DataDir, 0o755); err != nil { + log.Fatalf("creating data dir: %v", err) + } + setupLogFileNamed(appCfg.DataDir, "artifact-watch.log") + + if cfg.ArtifactFolder == "" { + fatal("artifact watch: folder target is required") + } + + debounce := cfg.Debounce + if debounce <= 0 { + debounce = defaultWatchDebounce + } + interval := cfg.Interval + if interval <= 0 { + interval = defaultWatchInterval + } + + lockPath, err := (daemon.RuntimeStore{ + Dir: appCfg.DataDir, + Prefix: "artifact-watch", + }).LockPath() + if err != nil { + fatal("artifact watch: %v", err) + } + lock := flock.New(lockPath) + locked, err := lock.TryLock() + if err != nil { + fatal("artifact watch: locking %s: %v", lockPath, err) + } + if !locked { + fatal("artifact watch: already locked (%s)", lockPath) + } + defer func() { + if rerr := lock.Unlock(); rerr != nil { + log.Printf("artifact watch: releasing lock: %v", rerr) + } + }() + + applyClassifierConfig(appCfg) + database, writeLock := mustOpenWriteDB(context.Background(), appCfg) + defer closeWriteDB(database, writeLock) + + for _, def := range parser.Registry { + if !appCfg.IsUserConfigured(def.Type) { + continue + } + warnMissingDirs(appCfg.ResolveDirs(def.Type), string(def.Type)) + } + cleanResyncTemp(appCfg.DBPath) + + origin, err := resolveArtifactOrigin(appCfg, database) + if err != nil { + fatal("artifact watch origin: %v", err) + } + + ctx, stop := signal.NotifyContext( + context.Background(), os.Interrupt, syscall.SIGTERM, + ) + defer stop() + + engine := newArtifactWatchEngine(database, appCfg) + defer engine.Close() + + didResync := cfg.Full || database.NeedsResync() + if didResync { + engine.ResyncAll(ctx, nil) + } else { + engine.SyncAll(ctx, nil) + } + if ctx.Err() != nil { + return + } + + pusher := &artifactFolderPusher{ + appCfg: appCfg, + database: database, + engine: engine, + target: cfg.ArtifactFolder, + origin: origin, + token: artifactPeerToken(cfg), + allowInsecure: cfg.AllowInsecure, + baseline: cfg.Init, + } + + log.Printf( + "artifact watch: starting (origin=%q target=%q debounce=%s interval=%s)", + origin, cfg.ArtifactFolder, debounce, interval, + ) + fmt.Printf( + "agentsview sync --watch: syncing artifacts as %q to %s "+ + "(debounce %s, floor %s)\n", + origin, cfg.ArtifactFolder, debounce, interval, + ) + + if err := pusher.push(ctx, reasonStartup); err != nil { + log.Printf("artifact watch: initial sync failed: %v", err) + } + + runWatchedSink(ctx, watchedSinkConfig{ + AppConfig: appCfg, + Engine: engine, + Debounce: debounce, + Interval: interval, + LogPrefix: "artifact watch", + Push: func(c context.Context, r pushReason) error { + return pusher.push(c, r) + }, + }) +} diff --git a/cmd/agentsview/sync_watch_test.go b/cmd/agentsview/sync_watch_test.go new file mode 100644 index 000000000..b073c98cb --- /dev/null +++ b/cmd/agentsview/sync_watch_test.go @@ -0,0 +1,300 @@ +package main + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/parser" + syncpkg "go.kenn.io/agentsview/internal/sync" + "go.kenn.io/agentsview/internal/testjsonl" +) + +type countingArtifactWatchSyncer struct { + syncAllCalls int + flushCalls int +} + +func (s *countingArtifactWatchSyncer) SyncAll( + _ context.Context, _ syncpkg.ProgressFunc, +) syncpkg.SyncStats { + s.syncAllCalls++ + return syncpkg.SyncStats{} +} + +func (s *countingArtifactWatchSyncer) FlushSignals() { + s.flushCalls++ +} + +func openWatchTestDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(filepath.Join(t.TempDir(), "test.db")) + require.NoError(t, err) + t.Cleanup(func() { database.Close() }) + return database +} + +func seedStarredSession(t *testing.T, database *db.DB, id string) { + t.Helper() + require.NoError(t, database.UpsertSession(db.Session{ + ID: id, + Project: "alpha", + Machine: "local", + Agent: "claude", + MessageCount: 1, + UserMessageCount: 1, + CreatedAt: "2026-06-14T01:02:03Z", + })) + require.NoError(t, database.ReplaceSessionMessages(id, []db.Message{ + {SessionID: id, Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + })) + starred, err := database.StarSession(id) + require.NoError(t, err, "StarSession") + require.True(t, starred, "session should be newly starred") +} + +func TestArtifactFolderPusherPublishesBaselineOnFirstSuccessfulPush(t *testing.T) { + dataDir := t.TempDir() + target := t.TempDir() + database := openWatchTestDB(t) + seedStarredSession(t, database, "sess-1") + + pusher := &artifactFolderPusher{ + appCfg: config.Config{DataDir: dataDir}, + database: database, + target: target, + origin: "laptop-a1b2c3", + baseline: true, + } + require.NoError(t, pusher.push(context.Background(), reasonStartup)) + assert.False(t, pusher.baseline, + "baseline must be published at most once after a successful push") + + events, err := filepath.Glob( + filepath.Join(target, "laptop-a1b2c3", "meta", "*"), + ) + require.NoError(t, err) + assert.NotEmpty(t, events, + "--init --watch must publish baseline metadata events to the target") +} + +func TestArtifactFolderPusherRetainsBaselineAfterFailedPush(t *testing.T) { + dataDir := t.TempDir() + // A regular file as the folder target makes the exchange fail. + target := filepath.Join(t.TempDir(), "not-a-dir") + require.NoError(t, os.WriteFile(target, []byte("x"), 0o600)) + database := openWatchTestDB(t) + seedStarredSession(t, database, "sess-1") + + pusher := &artifactFolderPusher{ + appCfg: config.Config{DataDir: dataDir}, + database: database, + target: target, + origin: "laptop-a1b2c3", + baseline: true, + } + require.Error(t, pusher.push(context.Background(), reasonStartup)) + assert.True(t, pusher.baseline, + "a failed push must keep the baseline pending for the next retry") +} + +func TestArtifactFolderPusherPlumbsInsecurePeerOptIn(t *testing.T) { + target, requests := insecureArtifactPeerTarget(t) + dataDir := t.TempDir() + database := openWatchTestDB(t) + pusher := &artifactFolderPusher{ + appCfg: config.Config{DataDir: dataDir}, + database: database, + target: target, + origin: "desk-a1b2c3", + allowInsecure: true, + } + + require.NoError(t, pusher.push(context.Background(), reasonStartup)) + assert.Positive(t, requests.Load(), + "watch sync must reach an explicitly allowed plaintext peer") +} + +func TestArtifactFolderPusherRunsFullDiscoveryForIntervalAndShutdown(t *testing.T) { + dataDir := t.TempDir() + target := t.TempDir() + database := openWatchTestDB(t) + syncer := &countingArtifactWatchSyncer{} + pusher := &artifactFolderPusher{ + appCfg: config.Config{DataDir: dataDir}, + database: database, + engine: syncer, + target: target, + origin: "desk-a1b2c3", + } + + for _, reason := range []pushReason{reasonStartup, reasonChange} { + require.NoError(t, pusher.push(context.Background(), reason)) + } + assert.Zero(t, syncer.syncAllCalls, + "startup and watcher-driven pushes already synchronized local files") + assert.Equal(t, 2, syncer.flushCalls, + "every export must flush pending signal recomputes") + + require.NoError(t, pusher.push(context.Background(), reasonShutdown)) + assert.Equal(t, 1, syncer.syncAllCalls, + "shutdown must discover changes still pending in the watcher batch") + + require.NoError(t, pusher.push(context.Background(), reasonInterval)) + assert.Equal(t, 2, syncer.syncAllCalls, + "the periodic floor must discover changes from unwatched roots") + assert.Equal(t, 4, syncer.flushCalls) +} + +func TestArtifactFolderPusherShutdownDiscoversPendingWatcherChange(t *testing.T) { + claudeDir := t.TempDir() + projectDir := filepath.Join(claudeDir, "-Users-alice-work") + require.NoError(t, os.MkdirAll(projectDir, 0o755)) + + database := openWatchTestDB(t) + appCfg := config.Config{ + DataDir: t.TempDir(), + AgentDirs: map[parser.AgentType][]string{ + parser.AgentClaude: {claudeDir}, + }, + } + engine := newArtifactWatchEngine(database, appCfg) + t.Cleanup(engine.Close) + require.Zero(t, engine.SyncAll(context.Background(), nil).Synced) + + content := testjsonl.NewSessionBuilder(). + AddClaudeUser("2026-01-01T00:00:00Z", "pending", "/Users/alice/work/project"). + AddClaudeAssistant("2026-01-01T00:00:01Z", "ok"). + String() + require.NoError(t, os.WriteFile( + filepath.Join(projectDir, "pending-session.jsonl"), []byte(content), 0o644, + )) + + target := t.TempDir() + pusher := &artifactFolderPusher{ + appCfg: appCfg, + database: database, + engine: engine, + target: target, + origin: "desk-a1b2c3", + } + require.NoError(t, pusher.push(context.Background(), reasonShutdown)) + + session, err := database.GetSession(context.Background(), "pending-session") + require.NoError(t, err) + require.NotNil(t, session, + "shutdown flush must ingest a file whose watcher event is still pending") + assert.Equal(t, "pending", *session.FirstMessage) + pending, err := database.PendingArtifactExports(context.Background(), 1) + require.NoError(t, err) + assert.Empty(t, pending, + "shutdown flush must publish every newly discovered session") + + checkpoints, err := filepath.Glob(filepath.Join( + target, "desk-a1b2c3", string(artifact.KindCheckpoints), "cp-*.json", + )) + require.NoError(t, err) + require.Len(t, checkpoints, 1) + latest, err := os.ReadFile(checkpoints[0]) + require.NoError(t, err) + var checkpoint struct { + Sessions map[string]string `json:"sessions"` + } + require.NoError(t, json.Unmarshal(latest, &checkpoint)) + assert.Contains(t, checkpoint.Sessions, "desk-a1b2c3~pending-session", + "shutdown flush must make the discovered session reachable from the final checkpoint") +} + +func TestArtifactWatchEngineHonorsConfiguredCwdPrefixes(t *testing.T) { + claudeDir := t.TempDir() + projectDir := filepath.Join(claudeDir, "-Users-alice-work") + require.NoError(t, os.MkdirAll(projectDir, 0o755)) + writeSession := func(name, cwd, prompt string) { + content := testjsonl.NewSessionBuilder(). + AddClaudeUser("2026-01-01T00:00:00Z", prompt, cwd). + AddClaudeAssistant("2026-01-01T00:00:01Z", "ok"). + String() + require.NoError(t, os.WriteFile( + filepath.Join(projectDir, name+".jsonl"), []byte(content), 0o644, + )) + } + writeSession("allowed-session", "/Users/alice/work/project", "allowed") + writeSession("blocked-session", "/Users/alice/personal/project", "blocked") + + database := openWatchTestDB(t) + engine := newArtifactWatchEngine(database, config.Config{ + AgentDirs: map[parser.AgentType][]string{ + parser.AgentClaude: {claudeDir}, + }, + SyncIncludeCwdPrefixes: []string{"/Users/alice/work"}, + }) + t.Cleanup(engine.Close) + stats := engine.SyncAll(context.Background(), nil) + require.Equal(t, 1, stats.Synced) + + allowed, err := database.GetSession(context.Background(), "allowed-session") + require.NoError(t, err) + require.NotNil(t, allowed) + assert.Equal(t, "allowed", *allowed.FirstMessage) + blocked, err := database.GetSession(context.Background(), "blocked-session") + require.NoError(t, err) + assert.Nil(t, blocked, + "watch mode must not ingest sessions outside sync_include_cwd_prefixes") +} + +func TestResolveArtifactOriginPromotesDBOriginIntoConfig(t *testing.T) { + dir := t.TempDir() + appCfg := config.Config{DataDir: dir} + database := openWatchTestDB(t) + require.NoError(t, artifact.AdoptOrigin(database, "laptop-a1b2c3"), + "seed DB-only origin") + + origin, err := resolveArtifactOrigin(appCfg, database) + require.NoError(t, err) + assert.Equal(t, "laptop-a1b2c3", origin, + "a DB-only origin must be reused, not replaced by a generated one") + + data, err := os.ReadFile(filepath.Join(dir, "config.toml")) + require.NoError(t, err) + assert.Contains(t, string(data), `artifact_origin_id = "laptop-a1b2c3"`, + "the DB origin must be promoted into config so serve adopts the same origin") +} + +func TestResolveArtifactOriginConfigWinsOverDB(t *testing.T) { + dir := t.TempDir() + appCfg := config.Config{DataDir: dir, ArtifactOriginID: "desktop-d4e5f6"} + database := openWatchTestDB(t) + require.NoError(t, artifact.AdoptOrigin(database, "laptop-a1b2c3"), + "seed DB origin") + + origin, err := resolveArtifactOrigin(appCfg, database) + require.NoError(t, err) + assert.Equal(t, "desktop-d4e5f6", origin) + + stored, err := artifact.StoredOrigin(database) + require.NoError(t, err) + assert.Equal(t, "desktop-d4e5f6", stored, + "the authoritative config origin must replace a divergent DB origin") +} + +func TestResolveArtifactOriginGeneratesWhenAbsentEverywhere(t *testing.T) { + dir := t.TempDir() + appCfg := config.Config{DataDir: dir} + database := openWatchTestDB(t) + + origin, err := resolveArtifactOrigin(appCfg, database) + require.NoError(t, err) + assert.Regexp(t, `^[a-z0-9]+(?:-[a-z0-9]+)*-[0-9a-f]{6}$`, origin) + + stored, err := artifact.StoredOrigin(database) + require.NoError(t, err) + assert.Equal(t, origin, stored, + "a directly generated CLI origin must be available to DB-only consumers") +} diff --git a/docs/artifact-sync.md b/docs/artifact-sync.md new file mode 100644 index 000000000..4e1ff44b3 --- /dev/null +++ b/docs/artifact-sync.md @@ -0,0 +1,299 @@ +--- +title: Trusted-Fleet Artifact Sync +description: Sync AgentsView archives between trusted personal machines without copying the live SQLite database +--- + +# Trusted-Fleet Artifact Sync + +Artifact sync exchanges immutable AgentsView artifacts between machines and +imports them into each machine's local SQLite archive. It is local-first: every +machine keeps its own complete database, and transports only move +content-addressed session artifacts plus metadata events. + +Use it for a fully trusted personal fleet: your laptop, desktop, home server, +NAS, or object-store bucket. Do not treat artifact sync as a team-sharing +security boundary. A peer that can write to the shared artifact target can +publish sessions and metadata for the fleet. + +## When To Use It + +Artifact sync makes sense when you want multiple machines to converge on the +same session archive without running PostgreSQL as the coordination point. + +It is a good fit for: + +- laptop plus desktop archives +- NAS, Syncthing, Dropbox, or rclone-backed rendezvous folders +- S3-compatible buckets such as MinIO, Backblaze B2, or AWS S3 +- trusted always-on AgentsView peers over HTTP + +Use [PostgreSQL Sync](/pg-sync/) or [DuckDB Mirror](/duckdb/) when you want a +read-only aggregation or analytics mirror. Those backends are still mirrors: +SQLite remains the local write/archive database, and artifact sync projects +foreign artifacts into ordinary SQLite rows before they can be pushed onward. + +## Quick Start + +Use a dedicated artifact share folder: + +```bash +agentsview sync --init /path/to/agentsview-artifacts +agentsview sync /path/to/agentsview-artifacts +``` + +Run `--init` once on each machine. It creates or adopts that machine's artifact +origin, backfills local sessions into the local artifact store, exchanges +artifacts with the target, and imports peer artifacts already present there. + +CLI artifact sync needs exclusive write access to the local archive. If a local +`agentsview` daemon is running, stop it first with `agentsview serve stop` and +retry. A running server still participates in a fleet as an +[HTTP peer](#http-peer) target for other machines. + +To keep a machine exchanging artifacts while it is online, run watch mode: + +```bash +agentsview sync --watch /path/to/agentsview-artifacts +``` + +Watch mode runs an initial local sync and artifact exchange, debounces local +session-file changes, retries failed exchanges on later changes or interval +ticks, and performs a final best-effort exchange on shutdown. + +## Targets + +### Folder + +```bash +agentsview sync /path/to/agentsview-artifacts +``` + +The folder may live on a local disk, NAS mount, Syncthing folder, Dropbox +folder, NFS share, or rclone-mounted bucket. The folder must be dedicated to +artifact sync. Do not point artifact sync at: + +- `AGENTSVIEW_DATA_DIR` +- the live SQLite database file or its WAL/SHM files +- a whole AgentsView data directory +- raw agent directories that contain live database files + +Copying the live SQLite database or the whole data directory with a general +file-sync tool is unsafe. Artifact sync exists specifically to avoid that. + +### HTTP Peer + +An AgentsView server can expose artifact exchange routes behind the existing +Bearer-token auth middleware: + +```bash +agentsview sync https://desktop:8080 --token +``` + +HTTP peer sync only sends an `Authorization` header when `--token` is provided. +It does not reuse the local server's `auth_token` for explicit peer URLs. HTTP +peer sync rejects redirects, so configure the final artifact API URL directly. +Credentials and artifact bodies are not forwarded to a redirect destination. + +Non-loopback peers require HTTPS by default. Plain `http://` remains available +for `localhost`, `127.0.0.0/8`, and `::1`. To connect to a remote plaintext peer +on a trusted test network, LAN, or VPN, opt in explicitly: + +```bash +agentsview sync http://desktop:8080 --token --allow-insecure +``` + +The override works for one-shot and `--watch` sync and logs a warning. It sends +the bearer token and full archive content without transport encryption, so use +HTTPS through a reverse proxy or VPN termination whenever possible. If you +expose a server beyond loopback, enable authentication and protect the token +like write access to the full archive. + +The HTTP client pulls every missing artifact from the peer and posts every local +artifact the peer is missing. Garbage collection of superseded artifacts on the +remote peer is the peer's own responsibility. + +### S3-Compatible Object Storage + +```bash +export AWS_ACCESS_KEY_ID=... +export AWS_SECRET_ACCESS_KEY=... +agentsview sync s3://my-bucket/agentsview +``` + +Credentials come from `AWS_ACCESS_KEY_ID`, `AWS_SECRET_ACCESS_KEY`, and optional +`AWS_SESSION_TOKEN`. Region resolves from `AGENTSVIEW_S3_REGION`, then +`AWS_REGION`, defaulting to `us-east-1`. + +For MinIO, Backblaze B2, or another S3-compatible service, set +`AGENTSVIEW_S3_ENDPOINT`. A custom endpoint automatically uses path-style +addressing; `AGENTSVIEW_S3_PATH_STYLE=true` forces it otherwise. + +Custom endpoints default to HTTPS when no scheme is given. Plain `http://` is +accepted only for loopback hosts or when +`AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT=true` is explicitly set. Use that +override only on a trusted test or private network because requests and artifact +content travel without TLS. Artifact sync rejects redirects. If a corrupt remote +object must be deleted so a later upload can heal the bucket, that deletion is +HTTPS-only, even when the insecure override permits other requests over HTTP. + +## How It Works + +Each install has a stable artifact origin ID. Locally owned sessions keep their +ordinary SQLite IDs and `machine='local'`. Foreign sessions are imported as +`~` with `machine=`. Source, parent, and +subagent relationship IDs are rewritten the same way before import, so SQLite, +PostgreSQL, and DuckDB see the same ordinary session graph after projection. + +The local artifact repository lives at `$AGENTSVIEW_DATA_DIR/artifacts`. It is a +private Docbank vault whose catalog, blobs, and packs are owned directly by the +running AgentsView process. Do not inspect, edit, copy, or synchronize files +inside this directory, and do not use it as a folder-sync target. A single +writable AgentsView daemon owns the vault; maintenance commands use that daemon +when it is available or acquire exclusive local ownership before opening the +vault directly. + +Folder targets keep the external artifact protocol layout: + +```text +// + checkpoints/cp-.json + manifests/.json.zst + segments/.ndjson.zst + meta/-.json + raw/ +``` + +Artifact kinds: + +- **checkpoints** list the current manifest hash for each session published by + an origin +- **manifests** hold the canonical session header, usage events, and segment + references +- **segments** hold canonical message NDJSON +- **metadata events** record user edits such as rename, trash/restore, star, + pin/unpin, and delete-everywhere +- **raw artifacts** are optional source snapshots when a parser can provide a + safe regular-file snapshot + +The canonical identity is always the SHA-256 and size of the uncompressed +logical bytes. Manifests and segments use zstd on the external wire, without +changing that identity. The local vault may also zstd-compress eligible loose +objects and later pack them; those physical changes are private storage details, +and every read still verifies and returns the exact canonical bytes. + +Logical artifacts are immutable. Folder writes use no-replace semantics, and S3 +writes use conditional create, so repeated syncs are idempotent set-union +operations. External names and extensions are independent of the Docbank layout +and remain stable when local objects move between loose and packed storage. + +Normal sync publishes only sessions in SQLite's durable changed-session queue. +Transports report the exact artifacts they create or repair, and import works +from those references instead of rescanning checkpoint or metadata history. +Incomplete checkpoints and metadata whose session has not landed remain in a +bounded durable retry queue across restarts. Peer status includes that pending +work; a later dependency transfer or startup retry resumes it. + +## Metadata And Deletes + +A machine records metadata events only after it has an artifact origin, which it +gets by running `agentsview sync --init`, running any artifact sync, or +receiving a peer exchange. Until then curation stays local and publishes no +artifact metadata; the `--init` baseline snapshot publishes the accumulated +curation state when the machine later joins a fleet. + +User curation converges through metadata events. Rename, trash/restore, star, +and pin changes are replayed deterministically with hybrid logical clocks. If +two peers edit the same metadata field close together, AgentsView records the +losing value in the local conflict log while still deriving one deterministic +current value. Conflicts and per-origin sync status are visible on the Peers +page in the UI, and a conflicted session shows a fork badge in its header. + +Emptying local trash is local-only. Fleet-wide permanent delete is explicit: +delete-everywhere writes a purge event and an exclusion tombstone so peers do +not resurrect the session from older artifacts. + +Checkpoint absence is never a deletion signal. A missing artifact, truncated +checkpoint, offline peer, or old target cannot remove local data. + +## Version And Failure Handling + +Artifact readers ignore unknown JSON fields. Unknown future metadata operations +are marked applied and skipped. Artifacts with a future format version are +deferred, not treated as successful imports, so older AgentsView versions keep +syncing the artifact kinds and versions they understand. + +Manifests that reference missing segments are also deferred. Import watermarks +advance only after all referenced content is hash-verified and applied. + +A corrupt artifact — a hash-mismatched segment or manifest, an undecodable +compressed object, or an unparseable checkpoint or metadata event — is skipped +so one damaged object never aborts a sync or spreads to other machines. Local +corruption is quarantined inside the private Docbank vault; it is not a loose +file beside the live artifact. A folder target may rename its own damaged wire +file with a `.corrupt` suffix so a valid holder can replace it. Invalid changed +artifacts do not advance landed provenance; a later valid checkpoint or repaired +dependency retries the durable import. Metadata events whose timestamps do not +parse are rejected at write time and quarantined during import. + +An object corrupt in place inside an S3 bucket is eligible for remote deletion +when a peer fetches it because it is missing locally and validation fails. +AgentsView attempts that deletion only over HTTPS; if it succeeds, a later push +from a valid holder can re-upload the object. A corrupt remote checkpoint found +while comparing a checkpoint already present locally is retained for its owner +to repair. Over permitted HTTP, corrupt remote objects are retained and skipped; +automatic deletion does not occur. + +## Garbage Collection + +`agentsview sync gc` performs conservative logical retention and one bounded +physical-maintenance pass in the local Docbank vault. It does not accept or +modify a folder, HTTP, or S3 target. Use `--dry-run` to preview logical +retention without trashing artifacts or running physical reclamation: + +```bash +agentsview sync gc --dry-run +agentsview sync gc --grace 168h --quarantine-grace 168h +``` + +GC keeps the latest checkpoint for each origin and every manifest, segment, and +raw artifact reachable from it. Origins without checkpoints are skipped rather +than interpreted as deleted. Unreachable logical artifacts first enter Docbank +trash. Physical maintenance then empties eligible trash, removes unreferenced +blobs, and repacks live content as separate bounded stages. By default, each +stage processes at most 1,024 objects. Blob garbage collection and repacking +also use a 256 MiB soft byte budget; trash emptying is bounded by object count. +When work remains, the command prints an exact resume command with opaque stage +cursors; do not edit the cursors. + +The writable daemon owns maintenance when it is running. If no daemon owns the +archive, the command opens the vault only after acquiring direct ownership. It +fails closed when a known owner is unreachable instead of risking concurrent +catalog or pack writes. + +## Resetting A Failed Local Vault + +`agentsview sync artifact-reset` is a fail-closed recovery operation for a local +artifact vault that Docbank cannot open or repair normally. It accepts no target +and never deletes or recreates the SQLite archive. The command verifies the +AgentsView ownership marker, moves the entire vault to a timestamped +`artifacts.reset-*` diagnostic path, creates a fresh Docbank vault, and +republishes this machine's artifact origin from authoritative SQLite rows. + +The moved-aside vault is never deleted automatically. Preserve it for diagnosis +and remove it manually only when it is no longer needed; repeated resets can +consume substantial disk. Foreign relay artifacts cannot be reconstructed from +local SQLite publication state, so they remain unavailable in the fresh vault +until a trusted peer or target sends them again. Sessions already imported into +SQLite are unchanged. If a reset is interrupted after the move, AgentsView +retains durable republish state and retries local publication when the fresh +repository is recovered. + +## Availability + +Two intermittent machines only sync directly while both are online and can reach +the same target. A NAS folder, always-on home server, S3-compatible bucket, +cloud-synced folder, or always-running AgentsView peer can act as a rendezvous. + +That rendezvous is a deployment convention, not a privileged architecture: +AgentsView still treats every participant as a peer and keeps the complete local +archive on each machine. diff --git a/docs/commands.md b/docs/commands.md index 72b5fdcf3..eb90d6203 100644 --- a/docs/commands.md +++ b/docs/commands.md @@ -182,15 +182,22 @@ starts a detached daemon so SQLite writes stay owned by one process. Set write-owner lock and exits when done. ```bash -agentsview sync [flags] +agentsview sync [artifact-target] [flags] ``` -| Flag | Default | Description | -| -------- | ------- | ---------------------------------------------- | -| `--full` | `false` | Force a full resync regardless of data version | -| `--host` | | SSH hostname for deprecated remote sync | -| `--user` | | SSH username for deprecated remote sync | -| `--port` | `22` | SSH port for deprecated remote sync | +| Flag | Default | Description | +| ------------------- | ------- | ---------------------------------------------------------- | +| `--full` | `false` | Force a full resync regardless of data version | +| `--host` | | SSH hostname for deprecated remote sync | +| `--user` | | SSH username for deprecated remote sync | +| `--port` | `22` | SSH port for deprecated remote sync | +| `--artifact-folder` | | Artifact folder, HTTP(S) peer, or `s3://` target | +| `--init` | `false` | Adopt an origin and publish the first artifact baseline | +| `--token` | | Bearer token for an HTTP(S) artifact peer | +| `--allow-insecure` | `false` | Permit plaintext HTTP to a non-loopback artifact peer | +| `--watch` | `false` | Continue artifact exchange after the initial sync | +| `--debounce` | `30s` | Coalesce changes before artifact exchange (`--watch` only) | +| `--interval` | `15m` | Periodic artifact-exchange floor (`--watch` only) | **Examples:** @@ -199,10 +206,25 @@ agentsview sync # incremental sync and exit agentsview sync --full # full resync and exit agentsview sync --host buildbox.local agentsview sync --host buildbox.local --user wes --port 2222 +agentsview sync --init /path/to/dedicated-artifact-share +agentsview sync --watch https://peer.example.test:8080 --token "$TOKEN" +agentsview sync s3://my-bucket/agentsview ``` After syncing, a summary of session and message counts is printed to stdout. +The optional positional artifact target and `--artifact-folder` select the same +folder, HTTP(S), or S3-compatible exchange target. Despite the legacy flag name, +the target does not have to be a folder. Do not point either form at +`AGENTSVIEW_DATA_DIR`, its private `artifacts` Docbank vault, a live SQLite +database, or a raw agent-session directory. `--debounce` and `--interval` have +no effect without `--watch`. + +Recurring artifact exchange publishes only SQLite sessions marked as changed. +Received artifacts are imported by exact changed reference; incomplete +checkpoint or metadata work is retained durably and retried in bounded batches +after dependencies arrive and on startup. + When `--host` is set, AgentsView syncs only that remote host and fails fast on error. If the local daemon has a matching configured `[[remote_hosts]]` entry, the daemon uses that stored entry and its configured transport. Otherwise, @@ -219,6 +241,41 @@ session source, using object size and `LastModified` metadata to skip unchanged sessions and downloading only objects that need parsing. See [Configuration — S3-Compatible Session Sources](/configuration/#s3-compatible-session-sources). +#### Artifact Vault Maintenance + +`agentsview sync gc` retains reachable logical artifacts and runs one bounded +physical-maintenance pass in the local Docbank vault. It accepts no positional +target. If a writable daemon is available, the command delegates maintenance to +that owner; otherwise it requires exclusive direct ownership. + +```bash +agentsview sync gc [flags] +``` + +| Flag | Default | Description | +| -------------------- | --------- | ---------------------------------------------------- | +| `--dry-run` | `false` | Preview logical retention; skip physical reclamation | +| `--grace` | `168h` | Age before unreachable logical artifacts enter trash | +| `--quarantine-grace` | `168h` | Diagnostic retention for quarantined artifacts | +| `--max-objects` | `1024` | Object budget for each physical stage | +| `--max-bytes` | `256 MiB` | Soft byte budget for blob GC and repacking | +| `--trash-cursor` | | Resume physical trash emptying | +| `--gc-cursor` | | Resume physical blob garbage collection | +| `--repack-cursor` | | Resume physical repacking | + +When work remains, output includes a complete resume command with the required +opaque cursors. + +#### Artifact Vault Reset + +`agentsview sync artifact-reset` accepts no arguments or operation-specific +flags. It is a last-resort, fail-closed recovery command: AgentsView verifies +ownership, moves the failed vault to a timestamped diagnostic path, creates a +fresh vault, and republishes local-origin artifacts from SQLite. The diagnostic +vault must be removed manually after investigation. Foreign relay artifacts +return only when a trusted peer or target sends them again; imported SQLite +sessions are not removed. + #### Configured Remote Hosts As of 0.33.0, remote hosts can also be declared in `~/.agentsview/config.toml` diff --git a/docs/index.md b/docs/index.md index a3bdde4db..a6ad8ab89 100644 --- a/docs/index.md +++ b/docs/index.md @@ -180,8 +180,10 @@ See [Activity](/activity/) for the full reference. AgentsView reads the session files that your [AI coding agents](/configuration/#session-discovery) leave on your machine and gives you a local-first desktop and web app to work with them. By default -everything stays on your machine. Optionally, [PostgreSQL sync](/pg-sync/) can -push session data to a shared database for team or multi-machine setups. +everything stays on your machine. Optionally, [artifact sync](/artifact-sync/) +can converge a trusted personal fleet without copying the live SQLite database, +and [PostgreSQL sync](/pg-sync/) can push session data to a shared database for +team or multi-machine dashboards.
diff --git a/docs/zensical.toml b/docs/zensical.toml index f0894bcad..3e5b5b7d0 100644 --- a/docs/zensical.toml +++ b/docs/zensical.toml @@ -30,6 +30,7 @@ nav = [ {"Semantic Search Internals" = "semantic-search-internals.md"}, {"Recall (Experimental)" = "recall.md"}, {"Remote Access" = "remote-access.md"}, + {"Artifact Sync" = "artifact-sync.md"}, {"PostgreSQL Sync" = "pg-sync.md"}, {"DuckDB Mirror" = "duckdb.md"}, {"Changelog" = "changelog.md"}, diff --git a/frontend/messages/en.json b/frontend/messages/en.json index 3d267e711..c0cd1c5f4 100644 --- a/frontend/messages/en.json +++ b/frontend/messages/en.json @@ -22,6 +22,7 @@ "nav_pinned": "Pinned", "nav_insights": "Insights", "nav_trash": "Trash", + "nav_peers": "Peers", "nav_search_sessions": "Search sessions...", "nav_search_sessions_shortcut": "Search sessions ({shortcut})", "header_transcript_normal_title": "Normal transcript - show all messages", @@ -219,6 +220,23 @@ "session_breadcrumb_find_in_session_shortcut": "Find in session (/)", "session_breadcrumb_rename": "Rename", "session_breadcrumb_delete": "Delete", + "session_breadcrumb_metadata_conflicts": "Metadata conflicts", + "session_breadcrumb_conflict_field_name": "Name", + "session_breadcrumb_conflict_field_trash": "Trash", + "session_breadcrumb_conflict_field_star": "Star", + "session_breadcrumb_conflict_field_delete_everywhere": "Delete everywhere", + "session_breadcrumb_conflict_field_pin": "Pin", + "session_breadcrumb_conflict_default_title": "Default title", + "session_breadcrumb_conflict_restored": "Restored", + "session_breadcrumb_conflict_in_trash": "In trash", + "session_breadcrumb_conflict_starred": "Starred", + "session_breadcrumb_conflict_unstarred": "Unstarred", + "session_breadcrumb_conflict_unpinned_target": "Unpinned {target}", + "session_breadcrumb_conflict_pinned_target": "Pinned {target}", + "session_breadcrumb_conflict_pinned_target_note": "Pinned {target}: {note}", + "session_breadcrumb_conflict_unknown_origin": "unknown origin", + "session_breadcrumb_conflict_current": "Current", + "session_breadcrumb_conflict_other": "Other", "session_breadcrumb_resumed_in": "Resumed in {target}", "session_breadcrumb_command_copied": "Command copied!", "session_breadcrumb_failed": "Failed", @@ -488,6 +506,61 @@ "sidebar_row_rename": "Rename", "sidebar_row_open_in_new_tab": "Open in new tab", "sidebar_row_delete": "Delete", + "peers_title": "Peers", + "peers_refresh_status": "Refresh peer status", + "peers_conflict_count_singular": "1 metadata conflict", + "peers_conflict_count_plural": "{count} metadata conflicts", + "peers_loading": "Loading peers...", + "peers_load_more": "Load more", + "peers_loading_more": "Loading more...", + "peers_sync_error": "Sync error", + "peers_pagination_error": "More peers could not be loaded.", + "peers_empty": "No peers yet", + "peers_unavailable": "Artifact sync is not available on this server.", + "peers_this_machine": "This machine", + "peers_in_sync": "In sync", + "peers_pending": "{count} pending", + "peers_local_sessions_title": "Sessions present in this database from this peer", + "peers_local_sessions": "{count} local", + "peers_published_sessions_title": "Sessions this peer has published", + "peers_published_sessions": "{count} published", + "peers_checkpoint_title": "Latest checkpoint sequence", + "peers_checkpoint": "checkpoint #{seq}", + "peers_last_published_title": "Last checkpoint published", + "peers_updated": "updated {time}", + "trash_title": "Trash", + "trash_empty_local_title": "Empty local trash", + "trash_emptying": "Emptying...", + "trash_empty_local": "Empty Local Trash", + "trash_loading": "Loading trash...", + "trash_empty": "Trash is empty", + "trash_empty_desc": "Deleted sessions will appear here.", + "trash_messages": [ + { + "declarations": [ + "input count", + "input countLabel", + "local countPlural = count: plural" + ], + "selectors": [ + "countPlural" + ], + "match": { + "countPlural=one": "{countLabel} msg", + "countPlural=other": "{countLabel} msgs" + } + } + ], + "trash_deleted": "deleted {time}", + "trash_restore_session": "Restore session", + "trash_restore": "Restore", + "trash_delete_everywhere_prompt": "Delete everywhere?", + "trash_confirm_delete_everywhere": "Confirm delete everywhere", + "trash_deleting": "Deleting...", + "trash_confirm": "Confirm", + "trash_cancel": "Cancel", + "trash_delete_everywhere": "Delete Everywhere", + "trash_delete_everywhere_title": "Delete everywhere", "shared_all_projects": "All Projects", "shared_project_filter_placeholder": "Filter projects...", "shared_select_project": "Select project", @@ -1731,33 +1804,6 @@ "system_boundary_stop_hook": "Stop hook feedback", "system_boundary_title": "System boundary: {subtype}", "system_boundary_show_content": "Show content", - "trash_loading": "Loading trash...", - "trash_empty": "Trash is empty", - "trash_empty_desc": "Deleted sessions will appear here.", - "trash_title": "Trash", - "trash_emptying": "Emptying...", - "trash_empty_trash": "Empty Trash", - "trash_msgs": [ - { - "declarations": [ - "input count", - "input countLabel", - "local countPlural = count: plural" - ], - "selectors": [ - "countPlural" - ], - "match": { - "countPlural=one": "{countLabel} msg", - "countPlural=other": "{countLabel} msgs" - } - } - ], - "trash_deleted_ago": "deleted {time}", - "trash_restore_session": "Restore session", - "trash_restore": "Restore", - "trash_permanently_delete": "Permanently delete", - "trash_delete_forever": "Delete Forever", "trends_term": "Term", "trends_per1k_messages": "Per 1k messages", "trends_count": "Count", diff --git a/frontend/messages/fr.json b/frontend/messages/fr.json index f733c90bf..225a31034 100644 --- a/frontend/messages/fr.json +++ b/frontend/messages/fr.json @@ -22,6 +22,7 @@ "nav_pinned": "Épinglés", "nav_insights": "Analyses", "nav_trash": "Corbeille", + "nav_peers": "Pairs", "nav_search_sessions": "Rechercher des sessions...", "nav_search_sessions_shortcut": "Rechercher des sessions ({shortcut})", "header_transcript_normal_title": "Transcription normale — afficher tous les messages", @@ -219,6 +220,23 @@ "session_breadcrumb_find_in_session_shortcut": "Rechercher dans la session (/)", "session_breadcrumb_rename": "Renommer", "session_breadcrumb_delete": "Supprimer", + "session_breadcrumb_metadata_conflicts": "Conflits de métadonnées", + "session_breadcrumb_conflict_field_name": "Nom", + "session_breadcrumb_conflict_field_trash": "Corbeille", + "session_breadcrumb_conflict_field_star": "Favori", + "session_breadcrumb_conflict_field_delete_everywhere": "Supprimer partout", + "session_breadcrumb_conflict_field_pin": "Épinglage", + "session_breadcrumb_conflict_default_title": "Titre par défaut", + "session_breadcrumb_conflict_restored": "Restaurée", + "session_breadcrumb_conflict_in_trash": "Dans la corbeille", + "session_breadcrumb_conflict_starred": "Ajoutée aux favoris", + "session_breadcrumb_conflict_unstarred": "Retirée des favoris", + "session_breadcrumb_conflict_unpinned_target": "{target} désépinglé", + "session_breadcrumb_conflict_pinned_target": "{target} épinglé", + "session_breadcrumb_conflict_pinned_target_note": "{target} épinglé : {note}", + "session_breadcrumb_conflict_unknown_origin": "origine inconnue", + "session_breadcrumb_conflict_current": "Actuel", + "session_breadcrumb_conflict_other": "Autre", "session_breadcrumb_resumed_in": "Reprise dans {target}", "session_breadcrumb_command_copied": "Commande copiée !", "session_breadcrumb_failed": "Échec", @@ -488,6 +506,61 @@ "sidebar_row_rename": "Renommer", "sidebar_row_open_in_new_tab": "Ouvrir dans un nouvel onglet", "sidebar_row_delete": "Supprimer", + "peers_title": "Pairs", + "peers_refresh_status": "Actualiser l'état des pairs", + "peers_conflict_count_singular": "1 conflit de métadonnées", + "peers_conflict_count_plural": "{count} conflits de métadonnées", + "peers_loading": "Chargement des pairs...", + "peers_load_more": "Charger plus", + "peers_loading_more": "Chargement...", + "peers_sync_error": "Erreur de synchronisation", + "peers_pagination_error": "Impossible de charger plus de pairs.", + "peers_empty": "Aucun pair pour le moment", + "peers_unavailable": "La synchronisation des artefacts n'est pas disponible sur ce serveur.", + "peers_this_machine": "Cette machine", + "peers_in_sync": "Synchronisé", + "peers_pending": "{count} en attente", + "peers_local_sessions_title": "Sessions de ce pair présentes dans cette base de données", + "peers_local_sessions": "{count} locales", + "peers_published_sessions_title": "Sessions publiées par ce pair", + "peers_published_sessions": "{count} publiées", + "peers_checkpoint_title": "Numéro du dernier point de contrôle", + "peers_checkpoint": "point de contrôle n°{seq}", + "peers_last_published_title": "Dernier point de contrôle publié", + "peers_updated": "mis à jour {time}", + "trash_title": "Corbeille", + "trash_empty_local_title": "Vider la corbeille locale", + "trash_emptying": "Vidage...", + "trash_empty_local": "Vider la corbeille locale", + "trash_loading": "Chargement de la corbeille...", + "trash_empty": "La corbeille est vide", + "trash_empty_desc": "Les sessions supprimées apparaîtront ici.", + "trash_messages": [ + { + "declarations": [ + "input count", + "input countLabel", + "local countPlural = count: plural" + ], + "selectors": [ + "countPlural" + ], + "match": { + "countPlural=one": "{countLabel} message", + "countPlural=other": "{countLabel} messages" + } + } + ], + "trash_deleted": "supprimée {time}", + "trash_restore_session": "Restaurer la session", + "trash_restore": "Restaurer", + "trash_delete_everywhere_prompt": "Supprimer partout ?", + "trash_confirm_delete_everywhere": "Confirmer la suppression partout", + "trash_deleting": "Suppression...", + "trash_confirm": "Confirmer", + "trash_cancel": "Annuler", + "trash_delete_everywhere": "Supprimer partout", + "trash_delete_everywhere_title": "Supprimer partout", "shared_all_projects": "Tous les projets", "shared_project_filter_placeholder": "Filtrer les projets...", "shared_select_project": "Sélectionner un projet", @@ -1730,33 +1803,6 @@ "system_boundary_stop_hook": "Retour du hook d'arrêt", "system_boundary_title": "Limite système : {subtype}", "system_boundary_show_content": "Afficher le contenu", - "trash_loading": "Chargement de la corbeille...", - "trash_empty": "La corbeille est vide", - "trash_empty_desc": "Les sessions supprimées apparaîtront ici.", - "trash_title": "Corbeille", - "trash_emptying": "Vidage...", - "trash_empty_trash": "Vider la corbeille", - "trash_msgs": [ - { - "declarations": [ - "input count", - "input countLabel", - "local countPlural = count: plural" - ], - "selectors": [ - "countPlural" - ], - "match": { - "countPlural=one": "{countLabel} msg", - "countPlural=other": "{countLabel} msgs" - } - } - ], - "trash_deleted_ago": "supprimée {time}", - "trash_restore_session": "Restaurer la session", - "trash_restore": "Restaurer", - "trash_permanently_delete": "Supprimer définitivement", - "trash_delete_forever": "Supprimer définitivement", "trends_term": "Terme", "trends_per1k_messages": "Pour 1k messages", "trends_count": "Nombre", diff --git a/frontend/messages/ko.json b/frontend/messages/ko.json index 001930f7f..4a37d6750 100644 --- a/frontend/messages/ko.json +++ b/frontend/messages/ko.json @@ -22,6 +22,7 @@ "nav_pinned": "고정됨", "nav_insights": "인사이트", "nav_trash": "휴지통", + "nav_peers": "피어", "nav_search_sessions": "세션 검색...", "nav_search_sessions_shortcut": "세션 검색 ({shortcut})", "header_transcript_normal_title": "일반 트랜스크립트 - 모든 메시지 표시", @@ -202,7 +203,9 @@ "input countLabel", "local countPlural = count: plural" ], - "selectors": ["countPlural"], + "selectors": [ + "countPlural" + ], "match": { "countPlural=other": "{countLabel}개 단계" } @@ -215,6 +218,23 @@ "session_breadcrumb_find_in_session_shortcut": "세션에서 찾기 (/)", "session_breadcrumb_rename": "이름 변경", "session_breadcrumb_delete": "삭제", + "session_breadcrumb_metadata_conflicts": "메타데이터 충돌", + "session_breadcrumb_conflict_field_name": "이름", + "session_breadcrumb_conflict_field_trash": "휴지통", + "session_breadcrumb_conflict_field_star": "고정", + "session_breadcrumb_conflict_field_delete_everywhere": "모든 기기에서 삭제", + "session_breadcrumb_conflict_field_pin": "메시지 고정", + "session_breadcrumb_conflict_default_title": "기본 제목", + "session_breadcrumb_conflict_restored": "복원됨", + "session_breadcrumb_conflict_in_trash": "휴지통에 있음", + "session_breadcrumb_conflict_starred": "고정됨", + "session_breadcrumb_conflict_unstarred": "고정 해제됨", + "session_breadcrumb_conflict_unpinned_target": "{target} 고정 해제됨", + "session_breadcrumb_conflict_pinned_target": "{target} 고정됨", + "session_breadcrumb_conflict_pinned_target_note": "{target} 고정됨: {note}", + "session_breadcrumb_conflict_unknown_origin": "알 수 없는 출처", + "session_breadcrumb_conflict_current": "현재", + "session_breadcrumb_conflict_other": "기타", "session_breadcrumb_resumed_in": "{target}에서 재개됨", "session_breadcrumb_command_copied": "명령어가 복사되었습니다!", "session_breadcrumb_failed": "실패", @@ -475,6 +495,60 @@ "sidebar_row_rename": "이름 변경", "sidebar_row_open_in_new_tab": "새 탭에서 열기", "sidebar_row_delete": "삭제", + "peers_title": "피어", + "peers_refresh_status": "피어 상태 새로 고침", + "peers_conflict_count_singular": "메타데이터 충돌 1건", + "peers_conflict_count_plural": "메타데이터 충돌 {count}건", + "peers_loading": "피어를 불러오는 중...", + "peers_load_more": "더 보기", + "peers_loading_more": "더 불러오는 중...", + "peers_sync_error": "동기화 오류", + "peers_pagination_error": "피어를 더 불러올 수 없습니다.", + "peers_empty": "피어가 없습니다", + "peers_unavailable": "이 서버에서는 아티팩트 동기화를 사용할 수 없습니다.", + "peers_this_machine": "이 컴퓨터", + "peers_in_sync": "동기화됨", + "peers_pending": "{count}개 대기 중", + "peers_local_sessions_title": "이 피어의 세션 중 이 데이터베이스에 있는 세션", + "peers_local_sessions": "로컬 {count}개", + "peers_published_sessions_title": "이 피어가 게시한 세션", + "peers_published_sessions": "게시됨 {count}개", + "peers_checkpoint_title": "최신 체크포인트 순번", + "peers_checkpoint": "체크포인트 #{seq}", + "peers_last_published_title": "마지막 체크포인트 게시 시각", + "peers_updated": "{time} 업데이트됨", + "trash_title": "휴지통", + "trash_empty_local_title": "로컬 휴지통 비우기", + "trash_emptying": "비우는 중...", + "trash_empty_local": "로컬 휴지통 비우기", + "trash_loading": "휴지통을 불러오는 중...", + "trash_empty": "휴지통이 비어 있습니다", + "trash_empty_desc": "삭제된 세션이 여기에 표시됩니다.", + "trash_messages": [ + { + "declarations": [ + "input count", + "input countLabel", + "local countPlural = count: plural" + ], + "selectors": [ + "countPlural" + ], + "match": { + "countPlural=other": "메시지 {countLabel}개" + } + } + ], + "trash_deleted": "{time} 삭제됨", + "trash_restore_session": "세션 복원", + "trash_restore": "복원", + "trash_delete_everywhere_prompt": "모든 기기에서 삭제할까요?", + "trash_confirm_delete_everywhere": "모든 기기에서 삭제 확인", + "trash_deleting": "삭제 중...", + "trash_confirm": "확인", + "trash_cancel": "취소", + "trash_delete_everywhere": "모든 기기에서 삭제", + "trash_delete_everywhere_title": "모든 기기에서 삭제", "shared_all_projects": "모든 프로젝트", "shared_project_filter_placeholder": "프로젝트 필터링...", "shared_select_project": "프로젝트 선택", @@ -614,7 +688,9 @@ "input countLabel", "local countPlural = count: plural" ], - "selectors": ["countPlural"], + "selectors": [ + "countPlural" + ], "match": { "countPlural=other": "세션 {countLabel}개" } @@ -1694,32 +1770,6 @@ "system_boundary_stop_hook": "중지 훅 피드백", "system_boundary_title": "시스템 경계: {subtype}", "system_boundary_show_content": "내용 표시", - "trash_loading": "휴지통을 불러오는 중...", - "trash_empty": "휴지통이 비어 있습니다", - "trash_empty_desc": "삭제된 세션이 여기에 표시됩니다.", - "trash_title": "휴지통", - "trash_emptying": "비우는 중...", - "trash_empty_trash": "휴지통 비우기", - "trash_msgs": [ - { - "declarations": [ - "input count", - "input countLabel", - "local countPlural = count: plural" - ], - "selectors": [ - "countPlural" - ], - "match": { - "countPlural=other": "메시지 {countLabel}개" - } - } - ], - "trash_deleted_ago": "{time} 삭제됨", - "trash_restore_session": "세션 복원", - "trash_restore": "복원", - "trash_permanently_delete": "영구 삭제", - "trash_delete_forever": "영구 삭제", "trends_term": "용어", "trends_per1k_messages": "메시지 1,000건당", "trends_count": "개수", diff --git a/frontend/messages/zh-CN.json b/frontend/messages/zh-CN.json index 4bce0a529..049152509 100644 --- a/frontend/messages/zh-CN.json +++ b/frontend/messages/zh-CN.json @@ -22,6 +22,7 @@ "nav_pinned": "已固定", "nav_insights": "洞察", "nav_trash": "回收站", + "nav_peers": "Peers", "nav_search_sessions": "搜索会话...", "nav_search_sessions_shortcut": "搜索会话 ({shortcut})", "header_transcript_normal_title": "普通 transcript - 显示所有消息", @@ -215,6 +216,23 @@ "session_breadcrumb_find_in_session_shortcut": "在会话中查找 (/)", "session_breadcrumb_rename": "重命名", "session_breadcrumb_delete": "删除", + "session_breadcrumb_metadata_conflicts": "元数据冲突", + "session_breadcrumb_conflict_field_name": "名称", + "session_breadcrumb_conflict_field_trash": "回收站", + "session_breadcrumb_conflict_field_star": "固定", + "session_breadcrumb_conflict_field_delete_everywhere": "全局删除", + "session_breadcrumb_conflict_field_pin": "固定", + "session_breadcrumb_conflict_default_title": "默认标题", + "session_breadcrumb_conflict_restored": "已恢复", + "session_breadcrumb_conflict_in_trash": "在回收站中", + "session_breadcrumb_conflict_starred": "已固定", + "session_breadcrumb_conflict_unstarred": "已取消固定", + "session_breadcrumb_conflict_unpinned_target": "已取消固定 {target}", + "session_breadcrumb_conflict_pinned_target": "已固定 {target}", + "session_breadcrumb_conflict_pinned_target_note": "已固定 {target}: {note}", + "session_breadcrumb_conflict_unknown_origin": "未知来源", + "session_breadcrumb_conflict_current": "当前", + "session_breadcrumb_conflict_other": "其他", "session_breadcrumb_resumed_in": "已在 {target} 中继续", "session_breadcrumb_command_copied": "命令已复制!", "session_breadcrumb_failed": "失败", @@ -475,6 +493,60 @@ "sidebar_row_rename": "重命名", "sidebar_row_open_in_new_tab": "在新标签页打开", "sidebar_row_delete": "删除", + "peers_title": "Peers", + "peers_refresh_status": "刷新 peer 状态", + "peers_conflict_count_singular": "1 个元数据冲突", + "peers_conflict_count_plural": "{count} 个元数据冲突", + "peers_loading": "正在加载 peers...", + "peers_load_more": "加载更多", + "peers_loading_more": "正在加载更多...", + "peers_sync_error": "同步错误", + "peers_pagination_error": "无法加载更多 peers。", + "peers_empty": "暂无 peers", + "peers_unavailable": "此服务器不可用 artifact sync。", + "peers_this_machine": "此机器", + "peers_in_sync": "已同步", + "peers_pending": "{count} 个待处理", + "peers_local_sessions_title": "此数据库中来自该 peer 的会话", + "peers_local_sessions": "{count} 个本地", + "peers_published_sessions_title": "此 peer 已发布的会话", + "peers_published_sessions": "{count} 个已发布", + "peers_checkpoint_title": "最新 checkpoint 序号", + "peers_checkpoint": "checkpoint #{seq}", + "peers_last_published_title": "上次发布 checkpoint", + "peers_updated": "更新于 {time}", + "trash_title": "回收站", + "trash_empty_local_title": "清空本地回收站", + "trash_emptying": "正在清空...", + "trash_empty_local": "清空本地回收站", + "trash_loading": "正在加载回收站...", + "trash_empty": "回收站为空", + "trash_empty_desc": "已删除的会话会显示在这里。", + "trash_messages": [ + { + "declarations": [ + "input count", + "input countLabel", + "local countPlural = count: plural" + ], + "selectors": [ + "countPlural" + ], + "match": { + "countPlural=other": "{countLabel} 条消息" + } + } + ], + "trash_deleted": "删除于 {time}", + "trash_restore_session": "恢复会话", + "trash_restore": "恢复", + "trash_delete_everywhere_prompt": "全局删除?", + "trash_confirm_delete_everywhere": "确认全局删除", + "trash_deleting": "正在删除...", + "trash_confirm": "确认", + "trash_cancel": "取消", + "trash_delete_everywhere": "全局删除", + "trash_delete_everywhere_title": "全局删除", "shared_all_projects": "所有项目", "shared_project_filter_placeholder": "筛选项目...", "shared_select_project": "选择项目", @@ -1692,32 +1764,6 @@ "system_boundary_stop_hook": "停止钩子反馈", "system_boundary_title": "系统边界:{subtype}", "system_boundary_show_content": "显示内容", - "trash_loading": "正在加载回收站...", - "trash_empty": "回收站为空", - "trash_empty_desc": "已删除的会话会显示在这里。", - "trash_title": "回收站", - "trash_emptying": "正在清空...", - "trash_empty_trash": "清空回收站", - "trash_msgs": [ - { - "declarations": [ - "input count", - "input countLabel", - "local countPlural = count: plural" - ], - "selectors": [ - "countPlural" - ], - "match": { - "countPlural=other": "{countLabel} 条消息" - } - } - ], - "trash_deleted_ago": "{time} 删除", - "trash_restore_session": "恢复会话", - "trash_restore": "恢复", - "trash_permanently_delete": "永久删除", - "trash_delete_forever": "永久删除", "trends_term": "词项", "trends_per1k_messages": "每千条消息", "trends_count": "数量", diff --git a/frontend/messages/zh-TW.json b/frontend/messages/zh-TW.json index 5b5b879d2..75890c0dc 100644 --- a/frontend/messages/zh-TW.json +++ b/frontend/messages/zh-TW.json @@ -22,6 +22,7 @@ "nav_pinned": "已固定", "nav_insights": "洞察", "nav_trash": "回收站", + "nav_peers": "Peers", "nav_search_sessions": "搜索會話...", "nav_search_sessions_shortcut": "搜索會話 ({shortcut})", "header_transcript_normal_title": "普通 transcript - 顯示所有消息", @@ -215,6 +216,23 @@ "session_breadcrumb_find_in_session_shortcut": "在會話中查找 (/)", "session_breadcrumb_rename": "重命名", "session_breadcrumb_delete": "刪除", + "session_breadcrumb_metadata_conflicts": "元資料衝突", + "session_breadcrumb_conflict_field_name": "名稱", + "session_breadcrumb_conflict_field_trash": "回收站", + "session_breadcrumb_conflict_field_star": "固定", + "session_breadcrumb_conflict_field_delete_everywhere": "全域刪除", + "session_breadcrumb_conflict_field_pin": "固定", + "session_breadcrumb_conflict_default_title": "預設標題", + "session_breadcrumb_conflict_restored": "已恢復", + "session_breadcrumb_conflict_in_trash": "在回收站中", + "session_breadcrumb_conflict_starred": "已固定", + "session_breadcrumb_conflict_unstarred": "已取消固定", + "session_breadcrumb_conflict_unpinned_target": "已取消固定 {target}", + "session_breadcrumb_conflict_pinned_target": "已固定 {target}", + "session_breadcrumb_conflict_pinned_target_note": "已固定 {target}: {note}", + "session_breadcrumb_conflict_unknown_origin": "未知來源", + "session_breadcrumb_conflict_current": "目前", + "session_breadcrumb_conflict_other": "其他", "session_breadcrumb_resumed_in": "已在 {target} 中繼續", "session_breadcrumb_command_copied": "命令已複製!", "session_breadcrumb_failed": "失敗", @@ -475,6 +493,60 @@ "sidebar_row_rename": "重命名", "sidebar_row_open_in_new_tab": "在新標籤頁打開", "sidebar_row_delete": "刪除", + "peers_title": "Peers", + "peers_refresh_status": "重新整理 peer 狀態", + "peers_conflict_count_singular": "1 個元資料衝突", + "peers_conflict_count_plural": "{count} 個元資料衝突", + "peers_loading": "正在加載 peers...", + "peers_load_more": "載入更多", + "peers_loading_more": "正在載入更多...", + "peers_sync_error": "同步錯誤", + "peers_pagination_error": "無法載入更多 peers。", + "peers_empty": "暫無 peers", + "peers_unavailable": "此伺服器無法使用 artifact sync。", + "peers_this_machine": "此機器", + "peers_in_sync": "已同步", + "peers_pending": "{count} 個待處理", + "peers_local_sessions_title": "此資料庫中來自該 peer 的會話", + "peers_local_sessions": "{count} 個本地", + "peers_published_sessions_title": "此 peer 已發佈的會話", + "peers_published_sessions": "{count} 個已發佈", + "peers_checkpoint_title": "最新 checkpoint 序號", + "peers_checkpoint": "checkpoint #{seq}", + "peers_last_published_title": "上次發佈 checkpoint", + "peers_updated": "更新於 {time}", + "trash_title": "回收站", + "trash_empty_local_title": "清空本地回收站", + "trash_emptying": "正在清空...", + "trash_empty_local": "清空本地回收站", + "trash_loading": "正在加載回收站...", + "trash_empty": "回收站為空", + "trash_empty_desc": "已刪除的會話會顯示在這裡。", + "trash_messages": [ + { + "declarations": [ + "input count", + "input countLabel", + "local countPlural = count: plural" + ], + "selectors": [ + "countPlural" + ], + "match": { + "countPlural=other": "{countLabel} 條消息" + } + } + ], + "trash_deleted": "{time} 刪除", + "trash_restore_session": "恢復會話", + "trash_restore": "恢復", + "trash_delete_everywhere_prompt": "全域刪除?", + "trash_confirm_delete_everywhere": "確認全域刪除", + "trash_deleting": "正在刪除...", + "trash_confirm": "確認", + "trash_cancel": "取消", + "trash_delete_everywhere": "全域刪除", + "trash_delete_everywhere_title": "全域刪除", "shared_all_projects": "所有項目", "shared_project_filter_placeholder": "篩選項目...", "shared_select_project": "選擇項目", @@ -1692,32 +1764,6 @@ "system_boundary_stop_hook": "停止鉤子反饋", "system_boundary_title": "系統邊界:{subtype}", "system_boundary_show_content": "顯示內容", - "trash_loading": "正在加載回收站...", - "trash_empty": "回收站為空", - "trash_empty_desc": "已刪除的會話會顯示在這裡。", - "trash_title": "回收站", - "trash_emptying": "正在清空...", - "trash_empty_trash": "清空回收站", - "trash_msgs": [ - { - "declarations": [ - "input count", - "input countLabel", - "local countPlural = count: plural" - ], - "selectors": [ - "countPlural" - ], - "match": { - "countPlural=other": "{countLabel} 條消息" - } - } - ], - "trash_deleted_ago": "{time} 刪除", - "trash_restore_session": "恢復會話", - "trash_restore": "恢復", - "trash_permanently_delete": "永久刪除", - "trash_delete_forever": "永久刪除", "trends_term": "詞項", "trends_per1k_messages": "每千條消息", "trends_count": "數量", diff --git a/frontend/package.json b/frontend/package.json index ad39b5a7d..3e183e0d7 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -11,6 +11,7 @@ "build": "vp build", "i18n:compile": "paraglide-js compile --project ./project.inlang --outdir ./src/lib/paraglide --strategy localStorage preferredLanguage baseLocale --emit-ts-declarations --silent", "generate:api": "node scripts/generate-api-client.mjs", + "test:generate-api": "node --test scripts/generated-file-normalizer.test.mjs", "preview": "vp preview", "precheck": "npm run i18n:compile", "check": "svelte-check --tsconfig ./tsconfig.json", diff --git a/frontend/scripts/generate-api-client.mjs b/frontend/scripts/generate-api-client.mjs index f0c871285..0b138d530 100644 --- a/frontend/scripts/generate-api-client.mjs +++ b/frontend/scripts/generate-api-client.mjs @@ -13,6 +13,8 @@ import { } from "node:path"; import { fileURLToPath } from "node:url"; +import { normalizeChangedGeneratedFiles } from "./generated-file-normalizer.mjs"; + const frontendDir = resolve( dirname(fileURLToPath(import.meta.url)), "..", @@ -46,6 +48,44 @@ function suppressExpectedAbortLogging() { ); } +function preserveBinaryResponseBodies() { + const requestPath = join( + frontendDir, + "src/lib/api/generated/core/request.ts", + ); + const source = readFileSync(requestPath, "utf8"); + const generatedDecoder = ` const jsonTypes = ['application/json', 'application/problem+json'] + const isJSON = jsonTypes.some(type => contentType.toLowerCase().startsWith(type)); + if (isJSON) { + return await response.json(); + } else { + return await response.text(); + } +`; + const binaryAwareDecoder = ` const normalizedContentType = contentType.toLowerCase(); + const jsonTypes = ['application/json', 'application/problem+json']; + const binaryTypes = ['application/octet-stream', 'application/zstd']; + const isJSON = jsonTypes.some(type => normalizedContentType.startsWith(type)); + const isBinary = binaryTypes.some(type => normalizedContentType.startsWith(type)); + if (isJSON) { + return await response.json(); + } else if (isBinary) { + return await response.blob(); + } else { + return await response.text(); + } +`; + if (!source.includes(generatedDecoder)) { + throw new Error( + "generated request body handler no longer matches the binary-body patch", + ); + } + writeFileSync( + requestPath, + source.replace(generatedDecoder, binaryAwareDecoder), + ); +} + function run(cmd, args, options = {}) { const result = spawnSync(cmd, args, { cwd: options.cwd, @@ -89,6 +129,8 @@ try { run("npx", openapiArgs, { cwd: frontendDir }); } suppressExpectedAbortLogging(); + preserveBinaryResponseBodies(); + normalizeChangedGeneratedFiles(repoRoot); } finally { rmSync(tempDir, { recursive: true, force: true }); } diff --git a/frontend/scripts/generated-file-normalizer.mjs b/frontend/scripts/generated-file-normalizer.mjs new file mode 100644 index 000000000..91d557f65 --- /dev/null +++ b/frontend/scripts/generated-file-normalizer.mjs @@ -0,0 +1,89 @@ +import { spawnSync } from "node:child_process"; +import { lstatSync, readFileSync, realpathSync, writeFileSync } from "node:fs"; +import { isAbsolute, relative, resolve, sep } from "node:path"; + +function isolatedGitEnvironment() { + return Object.fromEntries( + Object.entries(process.env).filter(([name]) => !name.startsWith("GIT_")), + ); +} + +function gitPathList(repoRoot, args) { + const result = spawnSync("git", args, { + cwd: repoRoot, + encoding: null, + env: isolatedGitEnvironment(), + stdio: ["ignore", "pipe", "pipe"], + }); + if (result.error) throw result.error; + if (result.status !== 0) { + throw new Error( + `git ${args.join(" ")} exited ${result.status}: ${result.stderr.toString("utf8").trim()}`, + ); + } + return result.stdout + .toString("utf8") + .split("\0") + .filter((path) => path !== ""); +} + +function isWithin(root, path) { + const child = relative(root, path); + return child === "" || (!isAbsolute(child) && child !== ".." && !child.startsWith(`..${sep}`)); +} + +export function normalizeChangedGeneratedFiles( + repoRoot, + generatedDir = "frontend/src/lib/api/generated", +) { + const absoluteRoot = resolve(repoRoot); + const generatedRoot = resolve(absoluteRoot, generatedDir); + if (!isWithin(absoluteRoot, generatedRoot)) { + throw new Error(`generated directory is outside repository: ${generatedDir}`); + } + const realGeneratedRoot = realpathSync(generatedRoot); + const changed = gitPathList(absoluteRoot, [ + "diff", + "--name-only", + "-z", + "--diff-filter=ACMRT", + "HEAD", + "--", + generatedDir, + ]); + const untracked = gitPathList(absoluteRoot, [ + "ls-files", + "-z", + "--others", + "--exclude-standard", + "--", + generatedDir, + ]); + + for (const path of new Set([...changed, ...untracked])) { + if (!path.endsWith(".ts")) continue; + const absolutePath = resolve(absoluteRoot, path); + if (!isWithin(generatedRoot, absolutePath)) { + throw new Error(`generated path is outside generated directory: ${path}`); + } + let info; + try { + info = lstatSync(absolutePath); + } catch (error) { + if (error?.code === "ENOENT") continue; + throw error; + } + if (info.isSymbolicLink()) { + throw new Error(`generated path is a symbolic link: ${path}`); + } + if (!info.isFile()) { + throw new Error(`generated path is not a regular file: ${path}`); + } + const realPath = realpathSync(absolutePath); + if (!isWithin(realGeneratedRoot, realPath)) { + throw new Error(`generated path is outside generated directory: ${path}`); + } + const source = readFileSync(absolutePath, "utf8"); + writeFileSync(absolutePath, `${source.trimEnd()}\n`); + } +} diff --git a/frontend/scripts/generated-file-normalizer.test.mjs b/frontend/scripts/generated-file-normalizer.test.mjs new file mode 100644 index 000000000..4fb169fb4 --- /dev/null +++ b/frontend/scripts/generated-file-normalizer.test.mjs @@ -0,0 +1,86 @@ +import { spawnSync } from "node:child_process"; +import { + existsSync, + mkdtempSync, + mkdirSync, + readFileSync, + rmSync, + symlinkSync, + writeFileSync, +} from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import test from "node:test"; +import assert from "node:assert/strict"; + +import { normalizeChangedGeneratedFiles } from "./generated-file-normalizer.mjs"; + +const generatedDir = "frontend/src/lib/api/generated"; + +function git(root, ...args) { + const env = Object.fromEntries( + Object.entries(process.env).filter(([name]) => !name.startsWith("GIT_")), + ); + const result = spawnSync("git", args, { cwd: root, encoding: "utf8", env }); + assert.equal(result.status, 0, result.stderr); +} + +function writeGenerated(root, name, contents) { + const path = join(root, generatedDir, name); + mkdirSync(join(path, ".."), { recursive: true }); + writeFileSync(path, contents); + return path; +} + +function fixture(t) { + const root = mkdtempSync(join(tmpdir(), "generated-normalizer-test-")); + t.after(() => rmSync(root, { recursive: true, force: true })); + git(root, "init", "--quiet"); + git(root, "config", "user.name", "Normalizer Test"); + git(root, "config", "user.email", "normalizer@example.invalid"); + return root; +} + +test("normalizes staged, unstaged, and untracked files while skipping deletions", (t) => { + const root = fixture(t); + const staged = writeGenerated(root, "staged.ts", "export const staged = 1;\n"); + const unstaged = writeGenerated(root, "unstaged.ts", "export const unstaged = 1;\n"); + const deleted = writeGenerated(root, "deleted.ts", "export const deleted = 1;\n"); + const untouched = writeGenerated(root, "untouched.ts", "export const untouched = 1; \n"); + git(root, "add", "."); + git(root, "commit", "--quiet", "-m", "fixture"); + + writeFileSync(staged, "export const staged = 2; \n\n"); + git(root, "add", staged); + writeFileSync(unstaged, "export const unstaged = 2; \n\n"); + rmSync(deleted); + const untracked = writeGenerated(root, "untracked.ts", "export const untracked = 1; \n\n"); + const newlineName = writeGenerated(root, "odd\nname.ts", "export const odd = 1; \n\n"); + + normalizeChangedGeneratedFiles(root, generatedDir); + + assert.equal(readFileSync(staged, "utf8"), "export const staged = 2;\n"); + assert.equal(readFileSync(unstaged, "utf8"), "export const unstaged = 2;\n"); + assert.equal(readFileSync(untracked, "utf8"), "export const untracked = 1;\n"); + assert.equal(readFileSync(newlineName, "utf8"), "export const odd = 1;\n"); + assert.equal(readFileSync(untouched, "utf8"), "export const untouched = 1; \n"); + assert.equal(existsSync(deleted), false); +}); + +test("rejects generated paths that escape through a symlink", (t) => { + const root = fixture(t); + writeGenerated(root, "tracked.ts", "export const tracked = true;\n"); + git(root, "add", "."); + git(root, "commit", "--quiet", "-m", "fixture"); + + const outside = join(root, "outside.ts"); + writeFileSync(outside, "private contents \n"); + const link = join(root, generatedDir, "escaped.ts"); + symlinkSync(outside, link); + + assert.throws( + () => normalizeChangedGeneratedFiles(root, generatedDir), + /symbolic link|outside generated directory/, + ); + assert.equal(readFileSync(outside, "utf8"), "private contents \n"); +}); diff --git a/frontend/src/App.svelte b/frontend/src/App.svelte index 1c735681f..b4bf37fe5 100644 --- a/frontend/src/App.svelte +++ b/frontend/src/App.svelte @@ -64,6 +64,7 @@ import PinnedPage from "./lib/components/pinned/PinnedPage.svelte"; import TrashPage from "./lib/components/trash/TrashPage.svelte"; import RecentEditsPage from "./lib/components/recentedits/RecentEditsPage.svelte"; + import PeersPage from "./lib/components/peers/PeersPage.svelte"; import SettingsPage from "./lib/components/settings/SettingsPage.svelte"; import { sessions, filtersToParams } from "./lib/stores/sessions.svelte.js"; import { messages } from "./lib/stores/messages.svelte.js"; @@ -682,6 +683,10 @@
+{:else if router.route === "peers"} +
+ +
{:else if router.route === "settings"}
diff --git a/frontend/src/lib/api/generated/core/request.ts b/frontend/src/lib/api/generated/core/request.ts index 320db4f1e..c788c6ce8 100644 --- a/frontend/src/lib/api/generated/core/request.ts +++ b/frontend/src/lib/api/generated/core/request.ts @@ -234,10 +234,15 @@ export const getResponseBody = async (response: Response): Promise => { try { const contentType = response.headers.get('Content-Type'); if (contentType) { - const jsonTypes = ['application/json', 'application/problem+json'] - const isJSON = jsonTypes.some(type => contentType.toLowerCase().startsWith(type)); + const normalizedContentType = contentType.toLowerCase(); + const jsonTypes = ['application/json', 'application/problem+json']; + const binaryTypes = ['application/octet-stream', 'application/zstd']; + const isJSON = jsonTypes.some(type => normalizedContentType.startsWith(type)); + const isBinary = binaryTypes.some(type => normalizedContentType.startsWith(type)); if (isJSON) { return await response.json(); + } else if (isBinary) { + return await response.blob(); } else { return await response.text(); } diff --git a/frontend/src/lib/api/generated/index.ts b/frontend/src/lib/api/generated/index.ts index 6685840b6..19f3e3d60 100644 --- a/frontend/src/lib/api/generated/index.ts +++ b/frontend/src/lib/api/generated/index.ts @@ -18,6 +18,12 @@ export type { AgentsResponse } from './models/AgentsResponse'; export type { AgentTotal } from './models/AgentTotal'; export type { ApiErrorResponse } from './models/ApiErrorResponse'; export type { ApplyWorktreeMappingsResponse } from './models/ApplyWorktreeMappingsResponse'; +export type { ArtifactFinalizeResponse } from './models/ArtifactFinalizeResponse'; +export type { ArtifactIndexResponse } from './models/ArtifactIndexResponse'; +export type { ArtifactOriginsResponse } from './models/ArtifactOriginsResponse'; +export type { ArtifactPeer } from './models/ArtifactPeer'; +export type { ArtifactPeersResponse } from './models/ArtifactPeersResponse'; +export type { ArtifactPostResponse } from './models/ArtifactPostResponse'; export type { BatchDeleteInputBody } from './models/BatchDeleteInputBody'; export type { BranchesResponse } from './models/BranchesResponse'; export type { BulkStarInputBody } from './models/BulkStarInputBody'; @@ -53,6 +59,7 @@ export type { DbHourOfWeekResponse } from './models/DbHourOfWeekResponse'; export type { DbInsight } from './models/DbInsight'; export type { DbMachineBreakdown } from './models/DbMachineBreakdown'; export type { DbMessage } from './models/DbMessage'; +export type { DbMetadataConflict } from './models/DbMetadataConflict'; export type { DbModelBreakdown } from './models/DbModelBreakdown'; export type { DbPeakContextDistribution } from './models/DbPeakContextDistribution'; export type { DbPercentiles } from './models/DbPercentiles'; @@ -148,6 +155,7 @@ export type { GithubConfigResponse } from './models/GithubConfigResponse'; export type { InsightCannedSessionFilters } from './models/InsightCannedSessionFilters'; export type { InsightsResponse } from './models/InsightsResponse'; export type { MachinesResponse } from './models/MachinesResponse'; +export type { MetadataConflictsResponse } from './models/MetadataConflictsResponse'; export type { ModelTotal } from './models/ModelTotal'; export type { Opener } from './models/Opener'; export type { OpenersResponse } from './models/OpenersResponse'; @@ -211,6 +219,7 @@ export type { WorktreeMappingsResponse } from './models/WorktreeMappingsResponse export { ActivityService } from './services/ActivityService'; export { AnalyticsService } from './services/AnalyticsService'; +export { ArtifactsService } from './services/ArtifactsService'; export { AssetsService } from './services/AssetsService'; export { ConfigService } from './services/ConfigService'; export { EmbeddingsService } from './services/EmbeddingsService'; diff --git a/frontend/src/lib/api/generated/models/ArtifactFinalizeResponse.ts b/frontend/src/lib/api/generated/models/ArtifactFinalizeResponse.ts new file mode 100644 index 000000000..f28bae902 --- /dev/null +++ b/frontend/src/lib/api/generated/models/ArtifactFinalizeResponse.ts @@ -0,0 +1,10 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type ArtifactFinalizeResponse = { + deferred: number; + imported_messages: number; + imported_metadata: number; + imported_sessions: number; +}; diff --git a/frontend/src/lib/api/generated/models/ArtifactIndexResponse.ts b/frontend/src/lib/api/generated/models/ArtifactIndexResponse.ts new file mode 100644 index 000000000..81a419f6b --- /dev/null +++ b/frontend/src/lib/api/generated/models/ArtifactIndexResponse.ts @@ -0,0 +1,13 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type ArtifactIndexResponse = { + checkpoints: any[] | null; + manifests: any[] | null; + meta: any[] | null; + next_cursor?: string; + origin: string; + raw: any[] | null; + segments: any[] | null; +}; diff --git a/frontend/src/lib/api/generated/models/ArtifactOriginsResponse.ts b/frontend/src/lib/api/generated/models/ArtifactOriginsResponse.ts new file mode 100644 index 000000000..8a110c915 --- /dev/null +++ b/frontend/src/lib/api/generated/models/ArtifactOriginsResponse.ts @@ -0,0 +1,8 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type ArtifactOriginsResponse = { + next_cursor?: string; + origins: any[] | null; +}; diff --git a/frontend/src/lib/api/generated/models/ArtifactPeer.ts b/frontend/src/lib/api/generated/models/ArtifactPeer.ts new file mode 100644 index 000000000..2bbd15c4d --- /dev/null +++ b/frontend/src/lib/api/generated/models/ArtifactPeer.ts @@ -0,0 +1,13 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type ArtifactPeer = { + checkpoint_seq: number; + is_local: boolean; + last_published?: string; + local_sessions: number; + origin: string; + published_sessions: number; + status: string; +}; diff --git a/frontend/src/lib/api/generated/models/ArtifactPeersResponse.ts b/frontend/src/lib/api/generated/models/ArtifactPeersResponse.ts new file mode 100644 index 000000000..3db49c058 --- /dev/null +++ b/frontend/src/lib/api/generated/models/ArtifactPeersResponse.ts @@ -0,0 +1,12 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type ArtifactPeersResponse = { + conflict_count: number; + local_origin: string; + next_cursor?: string; + oldest_pending_at?: string; + peers: any[] | null; + pending_imports: number; +}; diff --git a/frontend/src/lib/api/generated/models/ArtifactPostResponse.ts b/frontend/src/lib/api/generated/models/ArtifactPostResponse.ts new file mode 100644 index 000000000..3caad6ae8 --- /dev/null +++ b/frontend/src/lib/api/generated/models/ArtifactPostResponse.ts @@ -0,0 +1,12 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type ArtifactPostResponse = { + duplicate: boolean; + hash?: string; + kind: string; + name: string; + origin: string; + size: number; +}; diff --git a/frontend/src/lib/api/generated/models/DaemonPushRequest.ts b/frontend/src/lib/api/generated/models/DaemonPushRequest.ts index 979142cf5..1959ad55d 100644 --- a/frontend/src/lib/api/generated/models/DaemonPushRequest.ts +++ b/frontend/src/lib/api/generated/models/DaemonPushRequest.ts @@ -5,6 +5,7 @@ import type { ConfigDuckDBConfig } from './ConfigDuckDBConfig'; import type { ConfigPGConfig } from './ConfigPGConfig'; export type DaemonPushRequest = { + automatic?: boolean; duckdb?: ConfigDuckDBConfig; exclude_projects?: any[] | null; full: boolean; @@ -14,4 +15,3 @@ export type DaemonPushRequest = { projects?: any[] | null; sync_state_target?: string; }; - diff --git a/frontend/src/lib/api/generated/models/DbMetadataConflict.ts b/frontend/src/lib/api/generated/models/DbMetadataConflict.ts new file mode 100644 index 000000000..66afad6d0 --- /dev/null +++ b/frontend/src/lib/api/generated/models/DbMetadataConflict.ts @@ -0,0 +1,18 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type DbMetadataConflict = { + created_at: string; + field: string; + id: number; + losing_op: string; + losing_order_key: string; + losing_origin: string; + losing_value: string; + session_gid: string; + winning_op: string; + winning_order_key: string; + winning_origin: string; + winning_value: string; +}; diff --git a/frontend/src/lib/api/generated/models/DbTopSession.ts b/frontend/src/lib/api/generated/models/DbTopSession.ts index 11808f97d..ba2c1497a 100644 --- a/frontend/src/lib/api/generated/models/DbTopSession.ts +++ b/frontend/src/lib/api/generated/models/DbTopSession.ts @@ -15,4 +15,3 @@ export type DbTopSession = { started_at?: string; termination_status?: string; }; - diff --git a/frontend/src/lib/api/generated/models/MetadataConflictsResponse.ts b/frontend/src/lib/api/generated/models/MetadataConflictsResponse.ts new file mode 100644 index 000000000..1dc167d78 --- /dev/null +++ b/frontend/src/lib/api/generated/models/MetadataConflictsResponse.ts @@ -0,0 +1,7 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +export type MetadataConflictsResponse = { + conflicts: any[] | null; +}; diff --git a/frontend/src/lib/api/generated/services/ArtifactsService.ts b/frontend/src/lib/api/generated/services/ArtifactsService.ts new file mode 100644 index 000000000..5930ad4ba --- /dev/null +++ b/frontend/src/lib/api/generated/services/ArtifactsService.ts @@ -0,0 +1,292 @@ +/* generated using openapi-typescript-codegen -- do not edit */ +/* istanbul ignore file */ +/* tslint:disable */ +/* eslint-disable */ +import type { ArtifactFinalizeResponse } from '../models/ArtifactFinalizeResponse'; +import type { ArtifactIndexResponse } from '../models/ArtifactIndexResponse'; +import type { ArtifactOriginsResponse } from '../models/ArtifactOriginsResponse'; +import type { ArtifactPeersResponse } from '../models/ArtifactPeersResponse'; +import type { ArtifactPostResponse } from '../models/ArtifactPostResponse'; +import type { CancelablePromise } from '../core/CancelablePromise'; +import { OpenAPI } from '../core/OpenAPI'; +import { request as __request } from '../core/request'; +export class ArtifactsService { + /** + * Finalize artifact uploads + * @returns ArtifactFinalizeResponse OK + * @throws ApiError + */ + public static postApiV1ArtifactsFinalize(): CancelablePromise { + return __request(OpenAPI, { + method: 'POST', + url: '/api/v1/artifacts/finalize', + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } + /** + * List artifact origins + * @returns ArtifactOriginsResponse OK + * @throws ApiError + */ + public static getApiV1ArtifactsOrigins({ + cursor, + limit = 512, + }: { + /** + * Opaque artifact origin cursor + */ + cursor?: string, + /** + * Maximum origins to return + */ + limit?: number, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'GET', + url: '/api/v1/artifacts/origins', + query: { + 'cursor': cursor, + 'limit': limit, + }, + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 422: `Unprocessable Entity`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } + /** + * List artifact peers + * @returns ArtifactPeersResponse OK + * @throws ApiError + */ + public static getApiV1ArtifactsPeers({ + cursor, + limit = 512, + }: { + /** + * Opaque peer origin cursor + */ + cursor?: string, + /** + * Maximum peer origins to return + */ + limit?: number, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'GET', + url: '/api/v1/artifacts/peers', + query: { + 'cursor': cursor, + 'limit': limit, + }, + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 422: `Unprocessable Entity`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } + /** + * Get latest artifact checkpoint + * @returns binary OK + * @throws ApiError + */ + public static getApiV1ArtifactsOriginCheckpoint({ + origin, + }: { + /** + * Artifact origin ID + */ + origin: string, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'GET', + url: '/api/v1/artifacts/{origin}/checkpoint', + path: { + 'origin': origin, + }, + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } + /** + * List artifact index for an origin + * @returns ArtifactIndexResponse OK + * @throws ApiError + */ + public static getApiV1ArtifactsOriginIndex({ + origin, + cursor, + limit = 512, + }: { + /** + * Artifact origin ID + */ + origin: string, + /** + * Opaque artifact index cursor + */ + cursor?: string, + /** + * Maximum artifact names to return + */ + limit?: number, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'GET', + url: '/api/v1/artifacts/{origin}/index', + path: { + 'origin': origin, + }, + query: { + 'cursor': cursor, + 'limit': limit, + }, + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 422: `Unprocessable Entity`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } + /** + * Get artifact + * @returns binary OK + * @throws ApiError + */ + public static getApiV1ArtifactsOriginKindName({ + origin, + kind, + name, + }: { + /** + * Artifact origin ID + */ + origin: string, + /** + * Artifact kind + */ + kind: string, + /** + * Artifact filename or hash + */ + name: string, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'GET', + url: '/api/v1/artifacts/{origin}/{kind}/{name}', + path: { + 'origin': origin, + 'kind': kind, + 'name': name, + }, + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } + /** + * Post artifact + * @returns ArtifactPostResponse OK + * @throws ApiError + */ + public static postApiV1ArtifactsOriginKindName({ + origin, + kind, + name, + requestBody, + }: { + /** + * Artifact origin ID + */ + origin: string, + /** + * Artifact kind + */ + kind: string, + /** + * Artifact filename or hash + */ + name: string, + requestBody: Blob, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'POST', + url: '/api/v1/artifacts/{origin}/{kind}/{name}', + path: { + 'origin': origin, + 'kind': kind, + 'name': name, + }, + body: requestBody, + mediaType: 'application/octet-stream', + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } +} diff --git a/frontend/src/lib/api/generated/services/SessionsService.ts b/frontend/src/lib/api/generated/services/SessionsService.ts index e28701793..b707bb821 100644 --- a/frontend/src/lib/api/generated/services/SessionsService.ts +++ b/frontend/src/lib/api/generated/services/SessionsService.ts @@ -8,6 +8,7 @@ import type { DbSessionActivityResponse } from '../models/DbSessionActivityRespo import type { DbSessionTiming } from '../models/DbSessionTiming'; import type { DbSidebarSessionIndex } from '../models/DbSidebarSessionIndex'; import type { EmptyTrashResponse } from '../models/EmptyTrashResponse'; +import type { MetadataConflictsResponse } from '../models/MetadataConflictsResponse'; import type { OpenRequest } from '../models/OpenRequest'; import type { OpenSessionResponse } from '../models/OpenSessionResponse'; import type { OrdinalsResponse } from '../models/OrdinalsResponse'; @@ -847,6 +848,40 @@ export class SessionsService { }, }); } + /** + * List session metadata conflicts + * @returns MetadataConflictsResponse OK + * @throws ApiError + */ + public static getApiV1SessionsIdMetadataConflicts({ + id, + }: { + /** + * Session ID + */ + id: string, + }): CancelablePromise { + return __request(OpenAPI, { + method: 'GET', + url: '/api/v1/sessions/{id}/metadata-conflicts', + path: { + 'id': id, + }, + errors: { + 400: `Bad Request`, + 401: `Unauthorized`, + 403: `Forbidden`, + 404: `Not Found`, + 409: `Conflict`, + 422: `Unprocessable Entity`, + 500: `Internal Server Error`, + 501: `Not Implemented`, + 502: `Bad Gateway`, + 503: `Service Unavailable`, + 504: `Gateway Timeout`, + }, + }); + } /** * Open session directory * @returns OpenSessionResponse OK diff --git a/frontend/src/lib/api/request.test.ts b/frontend/src/lib/api/request.test.ts new file mode 100644 index 000000000..6b2eb14eb --- /dev/null +++ b/frontend/src/lib/api/request.test.ts @@ -0,0 +1,37 @@ +import { describe, expect, it } from "vitest"; +import { getResponseBody } from "./generated/core/request"; + +async function expectExactBlob(contentType: string, expected: Uint8Array): Promise { + const response = new Response(Uint8Array.from(expected).buffer, { + headers: { "Content-Type": contentType }, + }); + + const body = await getResponseBody(response); + + expect(body).toBeInstanceOf(Blob); + expect(new Uint8Array(await (body as Blob).arrayBuffer())).toEqual(expected); +} + +describe("generated response decoding", () => { + it("preserves non-UTF-8 octet-stream response bytes", async () => { + await expectExactBlob("application/octet-stream", new Uint8Array([0x00, 0xff, 0x80, 0x7f])); + }); + + it("preserves zstd response bytes", async () => { + await expectExactBlob( + "application/zstd; charset=binary", + new Uint8Array([0x28, 0xb5, 0x2f, 0xfd, 0xff, 0x00]), + ); + }); + + it("continues parsing JSON error bodies", async () => { + const response = new Response('{"detail":"invalid artifact"}', { + status: 400, + headers: { "Content-Type": "application/problem+json; charset=utf-8" }, + }); + + await expect(getResponseBody(response)).resolves.toEqual({ + detail: "invalid artifact", + }); + }); +}); diff --git a/frontend/src/lib/components/layout/AppHeader.svelte b/frontend/src/lib/components/layout/AppHeader.svelte index 73a466a8a..e0568de7e 100644 --- a/frontend/src/lib/components/layout/AppHeader.svelte +++ b/frontend/src/lib/components/layout/AppHeader.svelte @@ -90,6 +90,7 @@ "insights", "trash", "recent-edits", + "peers", ] as const; const tabs: TopBarTab[] = $derived([ @@ -101,6 +102,7 @@ { id: "insights", label: m.nav_insights() }, { id: "trash", label: m.nav_trash() }, { id: "recent-edits", label: m.nav_recent_edits() }, + { id: "peers", label: m.nav_peers() }, ]); const activeTab = $derived( diff --git a/frontend/src/lib/components/layout/SessionBreadcrumb.svelte b/frontend/src/lib/components/layout/SessionBreadcrumb.svelte index 86e5932f8..08dfeaaf7 100644 --- a/frontend/src/lib/components/layout/SessionBreadcrumb.svelte +++ b/frontend/src/lib/components/layout/SessionBreadcrumb.svelte @@ -13,12 +13,14 @@ LinkIcon, SearchIcon, SquareTerminalIcon, + TriangleAlertIcon, } from "../../icons.js"; import { onDestroy, onMount } from "svelte"; import type { Session } from "../../api/types.js"; import { OpenersService, SessionsService, + type DbMetadataConflict, type ResumeRequest, type ResumeResponse, } from "../../api/generated/index"; @@ -77,6 +79,9 @@ const directoryRead = new LatestRead(); const costRead = new LatestRead(); const breakdownRead = new LatestRead(); + let metadataConflicts = $state([]); + let conflictsOpen = $state(false); + let conflictRequestSeq = 0; interface Opener { id: string; @@ -164,6 +169,31 @@ }); }); + $effect(() => { + if (!session) { + metadataConflicts = []; + conflictsOpen = false; + conflictRequestSeq++; + return; + } + const id = session.id; + metadataConflicts = []; + conflictsOpen = false; + const seq = ++conflictRequestSeq; + configureGeneratedClient(); + SessionsService.getApiV1SessionsIdMetadataConflicts({ id }) + .then((res) => { + if (seq !== conflictRequestSeq || session?.id !== id) return; + metadataConflicts = Array.isArray(res.conflicts) + ? (res.conflicts as DbMetadataConflict[]) + : []; + }) + .catch(() => { + if (seq !== conflictRequestSeq) return; + metadataConflicts = []; + }); + }); + let sessionCost = $state(null); let sessionCostIsRollup = $state(false); let sessionRollupSubagentCount = $state(0); @@ -582,6 +612,90 @@ : m.session_breadcrumb_failed()); } + type ConflictSide = "winning" | "losing"; + + function conflictFieldLabel(field: string): string { + if (field === "display_name") return m.session_breadcrumb_conflict_field_name(); + if (field === "deleted_at") return m.session_breadcrumb_conflict_field_trash(); + if (field === "starred") return m.session_breadcrumb_conflict_field_star(); + if (field === "purge") return m.session_breadcrumb_conflict_field_delete_everywhere(); + if (field.startsWith("pin:")) return m.session_breadcrumb_conflict_field_pin(); + return field.replaceAll("_", " "); + } + + function parseJSONValue(value: string): Record | null { + if (!value) return null; + try { + const parsed = JSON.parse(value) as unknown; + return parsed !== null && typeof parsed === "object" + ? parsed as Record + : null; + } catch { + return null; + } + } + + function conflictSideValue( + conflict: DbMetadataConflict, + side: ConflictSide, + ): string { + const op = side === "winning" + ? conflict.winning_op + : conflict.losing_op; + const value = side === "winning" + ? conflict.winning_value + : conflict.losing_value; + if (conflict.field === "display_name") { + const parsed = parseJSONValue(value); + const displayName = parsed?.display_name; + if (typeof displayName === "string" && displayName.trim()) { + return displayName; + } + return m.session_breadcrumb_conflict_default_title(); + } + if (conflict.field === "deleted_at") { + if (op === "restore") return m.session_breadcrumb_conflict_restored(); + if (op === "soft_delete") return m.session_breadcrumb_conflict_in_trash(); + } + if (conflict.field === "starred") { + if (op === "star") return m.session_breadcrumb_conflict_starred(); + if (op === "unstar") return m.session_breadcrumb_conflict_unstarred(); + } + if (conflict.field === "purge") { + return op === "purge" + ? m.session_breadcrumb_conflict_field_delete_everywhere() + : op; + } + if (conflict.field.startsWith("pin:")) { + const parsed = parseJSONValue(value); + const ordinal = parsed?.ordinal; + const sourceUUID = parsed?.source_uuid; + const note = parsed?.note; + const target = typeof sourceUUID === "string" && sourceUUID + ? sourceUUID.slice(0, 8) + : typeof ordinal === "number" + ? `#${ordinal}` + : "pin"; + if (op === "unpin") { + return m.session_breadcrumb_conflict_unpinned_target({ target }); + } + if (typeof note === "string" && note.trim()) { + return m.session_breadcrumb_conflict_pinned_target_note({ target, note }); + } + return m.session_breadcrumb_conflict_pinned_target({ target }); + } + return value || op; + } + + function conflictOrigin( + conflict: DbMetadataConflict, + side: ConflictSide, + ): string { + return side === "winning" + ? conflict.winning_origin || m.session_breadcrumb_conflict_unknown_origin() + : conflict.losing_origin || m.session_breadcrumb_conflict_unknown_origin(); + } + async function handleOpenIn(opener: Opener) { if (!session) return; showOpenMenu = false; @@ -699,6 +813,9 @@ } else if (showOpenMenu) { showOpenMenu = false; e.preventDefault(); + } else if (conflictsOpen) { + conflictsOpen = false; + e.preventDefault(); } return; } @@ -730,6 +847,9 @@ if (!(target as HTMLElement).closest?.(".open-group")) { showOpenMenu = false; } + if (!(target as HTMLElement).closest?.(".conflict-group")) { + conflictsOpen = false; + } } @@ -822,6 +942,63 @@ > {getGradeLabel(session.health_grade)} + {#if metadataConflicts.length > 0} + + + {#if conflictsOpen} +
+
{m.session_breadcrumb_metadata_conflicts()}
+ {#each metadataConflicts as conflict (conflict.id)} +
+
+ {conflictFieldLabel(conflict.field)} +
+
+ {m.session_breadcrumb_conflict_current()} + + {conflictSideValue(conflict, "winning")} + + + {conflictOrigin(conflict, "winning")} + +
+
+ {m.session_breadcrumb_conflict_other()} + + {conflictSideValue(conflict, "losing")} + + + {conflictOrigin(conflict, "losing")} + +
+
+ {/each} +
+ {/if} +
+ {/if} {#if showDropdown} +
+ + {#if conflictCount > 0} +
+
+ {/if} + + {#if loading} +
+ + {m.peers_loading()} +
+ {:else if peers.length === 0} + + {#snippet icon()} + + {:else} +
+ {#each peers as peer (peer.origin)} + {@const state = syncState(peer)} +
+
+ {#if peer.is_local} +
+
+
+ {peer.origin} + {#if peer.is_local} + {m.peers_this_machine()} + {/if} + {#if state === "synced"} + {m.peers_in_sync()} + {:else if state === "behind"} + + {m.peers_pending({ + count: formatNumber(peer.published_sessions - peer.local_sessions), + })} + + {:else if state === "error"} + {m.peers_sync_error()} + {/if} +
+
+ + {m.peers_local_sessions({ count: formatNumber(peer.local_sessions) })} + + / + + {m.peers_published_sessions({ count: formatNumber(peer.published_sessions) })} + + {#if peer.checkpoint_seq > 0} + · + + {m.peers_checkpoint({ seq: String(peer.checkpoint_seq) })} + + {/if} + {#if peer.last_published} + · + + {m.peers_updated({ time: formatRelativeTime(peer.last_published) })} + + {/if} +
+
+
+ {/each} +
+ {#if paginationError} +

{m.peers_pagination_error()}

+ {/if} + {#if nextCursor} +
+
+ {/if} + {/if} +
+ + diff --git a/frontend/src/lib/components/peers/PeersPage.test.ts b/frontend/src/lib/components/peers/PeersPage.test.ts new file mode 100644 index 000000000..b5517bece --- /dev/null +++ b/frontend/src/lib/components/peers/PeersPage.test.ts @@ -0,0 +1,166 @@ +// @vitest-environment jsdom +import { + afterEach, + beforeEach, + describe, + expect, + it, + vi, +} from "vite-plus/test"; +import { + cleanup, + fireEvent, + render, + screen, + waitFor, +} from "@testing-library/svelte"; +import { setLocale } from "../../i18n/index.js"; + +const mocks = vi.hoisted(() => ({ + getPeers: vi.fn(), +})); + +vi.mock("../../api/generated/index", () => ({ + ArtifactsService: { + getApiV1ArtifactsPeers: mocks.getPeers, + }, +})); + +vi.mock("../../api/runtime.js", () => ({ + configureGeneratedClient: vi.fn(), +})); + +// @ts-ignore +import PeersPage from "./PeersPage.svelte"; + +function peer( + origin: string, + status = "pending", + publishedSessions = 1, + localSessions = 0, +) { + return { + origin, + status, + is_local: false, + checkpoint_seq: 1, + published_sessions: publishedSessions, + local_sessions: localSessions, + }; +} + +function page( + peers: ReturnType[], + nextCursor?: string, + conflictCount = 0, +) { + return { + peers, + local_origin: "local-a1b2c3", + conflict_count: conflictCount, + next_cursor: nextCursor, + }; +} + +describe("PeersPage", () => { + beforeEach(() => { + mocks.getPeers.mockReset(); + setLocale("en"); + }); + + afterEach(() => cleanup()); + + it("renders only the first bounded page on mount", async () => { + mocks.getPeers + .mockResolvedValueOnce(page([peer("peer-one")], "cursor-1")) + .mockResolvedValueOnce(page([peer("peer-two")])); + + render(PeersPage); + + await screen.findByText("peer-one"); + expect(mocks.getPeers).toHaveBeenCalledTimes(1); + expect(screen.queryByText("peer-two")).toBeNull(); + }); + + it("loads the next page only after the user requests it", async () => { + mocks.getPeers + .mockResolvedValueOnce(page([peer("peer-one")], "cursor-1")) + .mockResolvedValueOnce(page([peer("peer-two")])); + render(PeersPage); + await screen.findByText("peer-one"); + + await fireEvent.click(screen.getByRole("button", { name: "Load more" })); + + await screen.findByText("peer-two"); + expect(mocks.getPeers).toHaveBeenNthCalledWith(2, { cursor: "cursor-1" }); + }); + + it("renders the backend error status instead of inferring sync from counts", async () => { + mocks.getPeers.mockResolvedValueOnce( + page([peer("corrupt-peer", "error", 3, 3)]), + ); + render(PeersPage); + + await screen.findByText("corrupt-peer"); + expect(screen.getByText("Sync error")).toBeTruthy(); + expect(screen.queryByText("In sync")).toBeNull(); + }); + + it("terminates a repeated cursor and renders a pagination error", async () => { + mocks.getPeers + .mockResolvedValueOnce(page([peer("peer-one")], "repeat")) + .mockResolvedValueOnce(page([peer("peer-two")], "repeat")); + render(PeersPage); + await screen.findByText("peer-one"); + + await fireEvent.click(screen.getByRole("button", { name: "Load more" })); + + await screen.findByText("More peers could not be loaded."); + expect(screen.queryByText("peer-two")).toBeNull(); + expect(screen.queryByRole("button", { name: "Load more" })).toBeNull(); + expect(mocks.getPeers).toHaveBeenCalledTimes(2); + }); + + it("refresh replaces the first page and resets pagination state", async () => { + mocks.getPeers + .mockResolvedValueOnce(page([peer("old-one")], "old-cursor")) + .mockResolvedValueOnce(page([peer("old-two")])) + .mockResolvedValueOnce(page([peer("fresh-one")], "fresh-cursor")) + .mockResolvedValueOnce(page([peer("fresh-two")])); + render(PeersPage); + await screen.findByText("old-one"); + await fireEvent.click(screen.getByRole("button", { name: "Load more" })); + await screen.findByText("old-two"); + + await fireEvent.click( + screen.getByRole("button", { name: "Refresh peer status" }), + ); + + await screen.findByText("fresh-one"); + expect(screen.queryByText("old-one")).toBeNull(); + expect(screen.queryByText("old-two")).toBeNull(); + await fireEvent.click(screen.getByRole("button", { name: "Load more" })); + await screen.findByText("fresh-two"); + await waitFor(() => { + expect(mocks.getPeers).toHaveBeenNthCalledWith(3, { cursor: undefined }); + expect(mocks.getPeers).toHaveBeenNthCalledWith(4, { + cursor: "fresh-cursor", + }); + }); + }); + + it("clears a stale conflict banner when refresh fails", async () => { + mocks.getPeers + .mockResolvedValueOnce(page([peer("peer-one")], undefined, 1)) + .mockRejectedValueOnce(new Error("artifact store unavailable")); + render(PeersPage); + + await screen.findByText("1 metadata conflict"); + await fireEvent.click( + screen.getByRole("button", { name: "Refresh peer status" }), + ); + + await screen.findByText("Artifact sync is not available on this server."); + expect(screen.queryByText("1 metadata conflict")).toBeNull(); + }); +}); diff --git a/frontend/src/lib/components/trash/TrashPage.svelte b/frontend/src/lib/components/trash/TrashPage.svelte index 1e125b0ef..2ec7e9a31 100644 --- a/frontend/src/lib/components/trash/TrashPage.svelte +++ b/frontend/src/lib/components/trash/TrashPage.svelte @@ -18,6 +18,8 @@ let loading = $state(true); let emptying = $state(false); const trashRead = new LatestRead(); + let confirmingDeleteId = $state(null); + let deletingId = $state(null); interface TrashResponse { sessions: Session[]; @@ -53,6 +55,7 @@ configureGeneratedClient(); await SessionsService.postApiV1SessionsIdRestore({ id }); trashedSessions = trashedSessions.filter((s) => s.id !== id); + if (confirmingDeleteId === id) confirmingDeleteId = null; sessions.clearRecentlyDeleted(id); sessions.invalidateFilterCaches(); sessions.load(); @@ -61,15 +64,27 @@ } } + function requestPermanentDelete(id: string) { + confirmingDeleteId = id; + } + + function cancelPermanentDelete(id: string) { + if (confirmingDeleteId === id) confirmingDeleteId = null; + } + async function permanentDelete(id: string) { + deletingId = id; try { configureGeneratedClient(); await SessionsService.deleteApiV1SessionsIdPermanent({ id }); trashedSessions = trashedSessions.filter((s) => s.id !== id); + if (confirmingDeleteId === id) confirmingDeleteId = null; sessions.clearRecentlyDeleted(id); sessions.invalidateFilterCaches(); } catch { // silently fail + } finally { + if (deletingId === id) deletingId = null; } } @@ -95,6 +110,22 @@
+
+
+ {#if loading}
{m.trash_loading()}
{:else if trashedSessions.length === 0} @@ -104,19 +135,6 @@ {/snippet} {:else} -
-
-
{#each trashedSessions as session (session.id)}
@@ -125,12 +143,12 @@
{session.agent} {session.project} - {m.trash_msgs({ + {m.trash_messages({ count: session.user_message_count, countLabel: session.user_message_count.toLocaleString(), })} {#if session.deleted_at} - {m.trash_deleted_ago({ time: formatRelativeTime(session.deleted_at) })} + {m.trash_deleted({ time: formatRelativeTime(session.deleted_at) })} {/if}
@@ -142,13 +160,32 @@ > {m.trash_restore()} - + {#if confirmingDeleteId === session.id} + {m.trash_delete_everywhere_prompt()} + + + {:else} + + {/if}
{/each} @@ -282,6 +319,7 @@ .trash-card-actions { display: flex; + align-items: center; gap: 6px; flex-shrink: 0; } @@ -317,4 +355,38 @@ .perm-delete-btn:hover { background: color-mix(in srgb, var(--accent-red, #e55) 8%, transparent); } + + .perm-delete-btn--confirm { + border-color: var(--accent-red, #e55); + } + + .delete-confirm-label { + color: var(--text-muted); + font-size: 11px; + font-weight: 500; + white-space: nowrap; + } + + .cancel-delete-btn { + font-size: 11px; + font-weight: 500; + color: var(--text-muted); + background: none; + border: 1px solid var(--border-muted); + border-radius: var(--radius-sm); + padding: 4px 10px; + cursor: pointer; + transition: background 0.12s, color 0.12s; + } + + .cancel-delete-btn:hover:not(:disabled) { + color: var(--text-secondary); + background: var(--bg-surface-hover); + } + + .perm-delete-btn:disabled, + .cancel-delete-btn:disabled { + cursor: default; + opacity: 0.6; + } diff --git a/frontend/src/lib/components/trash/TrashPage.test.ts b/frontend/src/lib/components/trash/TrashPage.test.ts new file mode 100644 index 000000000..80843a33b --- /dev/null +++ b/frontend/src/lib/components/trash/TrashPage.test.ts @@ -0,0 +1,152 @@ +// @vitest-environment jsdom +import { + afterEach, + beforeEach, + describe, + expect, + it, + vi, +} from "vite-plus/test"; +import { mount, tick, unmount } from "svelte"; +// @ts-ignore +import TrashPage from "./TrashPage.svelte"; +import type { Session } from "../../api/types.js"; +import { SessionsService } from "../../api/generated/index"; + +vi.mock("../../api/runtime.js", async (importOriginal) => { + const orig = + await importOriginal(); + return { + ...orig, + configureGeneratedClient: vi.fn(), + }; +}); + +vi.mock("../../api/generated/index", async (importOriginal) => { + const orig = + await importOriginal(); + return { + ...orig, + SessionsService: { + deleteApiV1SessionsIdPermanent: vi.fn(), + deleteApiV1Trash: vi.fn(), + getApiV1Trash: vi.fn(), + postApiV1SessionsIdRestore: vi.fn(), + }, + }; +}); + +vi.mock("../../stores/sessions.svelte.js", () => ({ + sessions: { + clearRecentlyDeleted: vi.fn(), + invalidateFilterCaches: vi.fn(), + load: vi.fn(), + }, +})); + +const sessionsService = SessionsService as unknown as { + deleteApiV1SessionsIdPermanent: ReturnType; + deleteApiV1Trash: ReturnType; + getApiV1Trash: ReturnType; + postApiV1SessionsIdRestore: ReturnType; +}; + +function makeTrashedSession(overrides: Partial = {}): Session { + return { + id: "s1", + project: "alpha", + machine: "local", + agent: "claude", + first_message: "build this", + display_name: null, + started_at: "2026-06-14T01:02:03Z", + ended_at: "2026-06-14T01:03:03Z", + created_at: "2026-06-14T01:02:03Z", + deleted_at: "2026-06-14T01:04:03Z", + message_count: 2, + user_message_count: 1, + total_output_tokens: 0, + peak_context_tokens: 0, + is_automated: false, + ...overrides, + } as Session; +} + +function buttonByText(label: string): HTMLButtonElement { + const button = Array.from(document.querySelectorAll("button")) + .find((el) => el.textContent?.trim() === label); + expect(button).toBeTruthy(); + return button as HTMLButtonElement; +} + +beforeEach(() => { + sessionsService.getApiV1Trash + .mockReset() + .mockResolvedValue({ sessions: [makeTrashedSession()] }); + sessionsService.deleteApiV1SessionsIdPermanent + .mockReset() + .mockResolvedValue(undefined); + sessionsService.deleteApiV1Trash + .mockReset() + .mockResolvedValue({ deleted: 1 }); + sessionsService.postApiV1SessionsIdRestore + .mockReset() + .mockResolvedValue(undefined); +}); + +afterEach(() => { + document.body.innerHTML = ""; +}); + +describe("TrashPage", () => { + it("requires confirmation before deleting a session everywhere", async () => { + const component = mount(TrashPage, { target: document.body }); + + await vi.waitFor(() => { + expect(buttonByText("Delete Everywhere")).toBeTruthy(); + }); + + buttonByText("Delete Everywhere").dispatchEvent( + new MouseEvent("click", { bubbles: true }), + ); + await tick(); + + expect( + sessionsService.deleteApiV1SessionsIdPermanent, + ).not.toHaveBeenCalled(); + expect(document.body.textContent).toContain("Delete everywhere?"); + + buttonByText("Confirm").dispatchEvent( + new MouseEvent("click", { bubbles: true }), + ); + + await vi.waitFor(() => { + expect( + sessionsService.deleteApiV1SessionsIdPermanent, + ).toHaveBeenCalledWith({ id: "s1" }); + }); + + unmount(component); + }); + + it("keeps empty trash on the local endpoint", async () => { + const component = mount(TrashPage, { target: document.body }); + + await vi.waitFor(() => { + expect(buttonByText("Empty Local Trash")).toBeTruthy(); + }); + + buttonByText("Empty Local Trash").dispatchEvent( + new MouseEvent("click", { bubbles: true }), + ); + + await vi.waitFor(() => { + expect(sessionsService.deleteApiV1Trash).toHaveBeenCalled(); + }); + expect( + sessionsService.deleteApiV1SessionsIdPermanent, + ).not.toHaveBeenCalled(); + + unmount(component); + }); +}); diff --git a/frontend/src/lib/i18n/i18n.test.ts b/frontend/src/lib/i18n/i18n.test.ts index dde91a59f..4e443188d 100644 --- a/frontend/src/lib/i18n/i18n.test.ts +++ b/frontend/src/lib/i18n/i18n.test.ts @@ -154,7 +154,7 @@ describe("i18n locale selection", () => { count: 2, countLabel: "2", })).toBe("2 sessions"); - expect(m.trash_msgs({ + expect(m.trash_messages({ count: 1, countLabel: "1", })).toBe("1 msg"); diff --git a/frontend/src/lib/stores/router.svelte.ts b/frontend/src/lib/stores/router.svelte.ts index e6bc5cd24..3f9cafe14 100644 --- a/frontend/src/lib/stores/router.svelte.ts +++ b/frontend/src/lib/stores/router.svelte.ts @@ -7,6 +7,7 @@ export type Route = | "pinned" | "trash" | "recent-edits" + | "peers" | "settings"; const VALID_ROUTES: ReadonlySet = new Set([ @@ -18,6 +19,7 @@ const VALID_ROUTES: ReadonlySet = new Set([ "pinned", "trash", "recent-edits", + "peers", "settings", ]); diff --git a/frontend/src/lib/stores/router.test.ts b/frontend/src/lib/stores/router.test.ts index f70881788..813c0c024 100644 --- a/frontend/src/lib/stores/router.test.ts +++ b/frontend/src/lib/stores/router.test.ts @@ -68,6 +68,7 @@ describe("parsePath", () => { "insights", "pinned", "trash", + "peers", "settings", ]) { setURL(`/${route}`); diff --git a/frontend/vite.config.ts b/frontend/vite.config.ts index c6ded9a04..07050236d 100644 --- a/frontend/vite.config.ts +++ b/frontend/vite.config.ts @@ -136,7 +136,11 @@ export default defineConfig({ }, test: { environment: "jsdom", - exclude: ["e2e/**", "node_modules/**"], + exclude: [ + "e2e/**", + "node_modules/**", + "scripts/generated-file-normalizer.test.mjs", + ], setupFiles: ["./src/vitest-setup.ts"], server: { deps: { diff --git a/go.mod b/go.mod index 89c99f595..a5f23f811 100644 --- a/go.mod +++ b/go.mod @@ -25,7 +25,8 @@ require ( github.com/testcontainers/testcontainers-go v0.43.0 github.com/testcontainers/testcontainers-go/modules/postgres v0.43.0 github.com/tidwall/gjson v1.19.0 - go.kenn.io/kit v0.7.1 + go.kenn.io/docbank v0.10.2-0.20260722131024-46be436fbd38 + go.kenn.io/kit v0.11.0 golang.org/x/mod v0.37.0 golang.org/x/perf v0.0.0-20260615155930-9e4b9ddef5b6 golang.org/x/sync v0.21.0 @@ -96,6 +97,7 @@ require ( github.com/moby/term v0.5.2 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/ncruces/go-strftime v1.0.0 // indirect + github.com/oklog/ulid/v2 v2.1.1 // indirect github.com/opencontainers/go-digest v1.0.0 // indirect github.com/opencontainers/image-spec v1.1.1 // indirect github.com/pascaldekloe/name v1.0.0 // indirect diff --git a/go.sum b/go.sum index 6750deca2..b26999901 100644 --- a/go.sum +++ b/go.sum @@ -140,6 +140,8 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= +github.com/leanovate/gopter v0.2.11 h1:vRjThO1EKPb/1NsDXuDrzldR28RLkBflWYcU9CvzWu4= +github.com/leanovate/gopter v0.2.11/go.mod h1:aK3tzZP/C+p1m3SPRE4SYZFGP7jjkuSI4f7Xvpt0S9c= github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw= github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 h1:6E+4a0GO5zZEnZ81pIr0yLvtUWk2if982qA3F3QD6H4= @@ -188,12 +190,15 @@ github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/oklog/ulid/v2 v2.1.1 h1:suPZ4ARWLOJLegGFiZZ1dFAkqzhMjL3J1TzI+5wHz8s= +github.com/oklog/ulid/v2 v2.1.1/go.mod h1:rcEKHmBBKfef9DhnvX7y1HZBYxjXb0cP5ExxNsTT1QQ= github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM= github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040= github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M= github.com/pascaldekloe/name v1.0.0 h1:n7LKFgHixETzxpRv2R77YgPUFo85QHGZKrdaYm7eY5U= github.com/pascaldekloe/name v1.0.0/go.mod h1:Z//MfYJnH4jVpQ9wkclwu2I2MkHmXTlT9wR5UZScttM= +github.com/pborman/getopt v0.0.0-20170112200414-7148bc3a4c30/go.mod h1:85jBQOZwpVEaDAr341tbn15RS4fCAsIst0qp7i8ex1o= github.com/philhofer/fwd v1.2.0 h1:e6DnBTl7vGY+Gz322/ASL4Gyp1FspeMvx1RNDoToZuM= github.com/philhofer/fwd v1.2.0/go.mod h1:RqIHx9QI14HlwKwm98g9Re5prTQ6LdeRQn+gXJFxsJM= github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0= @@ -217,8 +222,8 @@ github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEy github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= -github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= -github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= +github.com/rogpeppe/go-internal v1.15.0 h1:D0RCU5rMAp+SpgkiNdrjfJ+LX4J1M32V2NeCY7EJ6hc= +github.com/rogpeppe/go-internal v1.15.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs= github.com/rs/xid v1.6.0 h1:fV591PaemRlL6JfRxGDEPl69wICngIQ3shQtzfy2gxU= github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= @@ -274,8 +279,14 @@ github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ= github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0= github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs= github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s= -go.kenn.io/kit v0.7.1 h1:4jUIPmBagCtPqhQe1wcZtgUo2HqAeWM3qTqoJe1SG0U= -go.kenn.io/kit v0.7.1/go.mod h1:qcJxICjZlquj1gKxtsCJeuLmlkAP6+Hui/hXqtJNZtc= +go.kenn.io/docbank v0.10.2-0.20260722122306-e50791ba45cb h1:J2H0CT/kT4vTjVwpH1cvZ5o84gf7eChyWE6+vbS2HCM= +go.kenn.io/docbank v0.10.2-0.20260722122306-e50791ba45cb/go.mod h1:fPbZ60TNZwMqXBQydiXiROHN8asWY5Y2nvposu6dQlg= +go.kenn.io/docbank v0.10.2-0.20260722130401-ea3191fb4fa6 h1:7s9HArIYlkjY4T38CJe6VeSCAZdBh5KzG1NlO8xpqm8= +go.kenn.io/docbank v0.10.2-0.20260722130401-ea3191fb4fa6/go.mod h1:fPbZ60TNZwMqXBQydiXiROHN8asWY5Y2nvposu6dQlg= +go.kenn.io/docbank v0.10.2-0.20260722131024-46be436fbd38 h1:mLlGZN8z5NDv//4E8tjDqRiahqJyxrP25afWTZOWWlw= +go.kenn.io/docbank v0.10.2-0.20260722131024-46be436fbd38/go.mod h1:fPbZ60TNZwMqXBQydiXiROHN8asWY5Y2nvposu6dQlg= +go.kenn.io/kit v0.11.0 h1:OdEaI8i3R7M0OTptrP2Osu+WJSg+lTClmHalRgn/T/U= +go.kenn.io/kit v0.11.0/go.mod h1:dComZhFNb4LR+Tj4ZD0slEDMkPJk0gGd8q/HEHXTGSM= go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= go.opentelemetry.io/contrib/bridges/prometheus v0.69.0 h1:saQoWg5845Q8TojpqeVStS7zGwVZ6bc5W2PJavTPiBM= diff --git a/internal/artifact/compression_test.go b/internal/artifact/compression_test.go new file mode 100644 index 000000000..e226776cc --- /dev/null +++ b/internal/artifact/compression_test.go @@ -0,0 +1,516 @@ +package artifact + +import ( + "bytes" + "context" + "fmt" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +func createCompressedTestArtifact( + t *testing.T, store ArtifactStore, origin string, kind Kind, hash string, encoded []byte, +) (Ref, error) { + t.Helper() + name := hash + switch kind { + case KindManifests: + name += ".json" + case KindSegments: + name += ".ndjson" + case KindMeta: + if !strings.HasSuffix(name, metadataEventExtension) { + name += metadataEventExtension + } + } + ref, err := NewRef(origin, kind, name) + require.NoError(t, err) + wire, err := ToWireRef(ref) + require.NoError(t, err) + _, err = CreateFromWire( + t.Context(), store, wire, bytes.NewReader(encoded), transportWireLimits(kind), + ) + return ref, err +} + +func assertArtifactNotStored(t *testing.T, store ArtifactStore, ref Ref) { + t.Helper() + _, err := store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) +} + +func latestTestStoreManifest(t *testing.T, store ArtifactStore, origin string) manifest { + t.Helper() + _, cp, err := latestStoreCheckpointSummary(t.Context(), store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + hash := cp.Sessions[origin+"~sess-1"] + require.NotEmpty(t, hash) + ref, err := NewRef(origin, KindManifests, hash+".json") + require.NoError(t, err) + m, err := decodeManifestWithLimits(readContractArtifact(t, store, ref), productionArtifactLimits()) + require.NoError(t, err) + return m +} + +func testStoreManifestMessages(t *testing.T, store ArtifactStore, origin string, m manifest) []db.Message { + t.Helper() + var messages []db.Message + for _, hash := range m.Segments { + ref, err := NewRef(origin, KindSegments, hash+".ndjson") + require.NoError(t, err) + segment, err := decodeSegment(readContractArtifact(t, store, ref)) + require.NoError(t, err) + messages = append(messages, segment...) + } + return messages +} + +func assertNoPublishedArtifacts(t *testing.T, store ArtifactStore, origin string) { + t.Helper() + for _, kind := range transportKinds { + page, err := firstStoreEntryPage(t.Context(), store, origin, kind, maxArtifactListPageSize) + require.NoError(t, err) + assert.Empty(t, page.Items, "unexpected published %s artifact", kind) + require.Empty(t, page.Next) + } +} + +func assertNoPublishedAuthority(t *testing.T, store ArtifactStore, origin string) { + t.Helper() + for _, kind := range []Kind{KindManifests, KindCheckpoints} { + page, err := firstStoreEntryPage(t.Context(), store, origin, kind, maxArtifactListPageSize) + require.NoError(t, err) + assert.Empty(t, page.Items, "unexpected published %s artifact", kind) + require.Empty(t, page.Next) + } +} + +func TestWriteArtifactRejectsManifestStructuralAmplification(t *testing.T) { + origin := "peer-a1b2c3" + validHash := strings64("a") + tests := []struct { + name string + manifest manifest + wantError string + }{ + { + name: "duplicate segment references", + manifest: manifest{ + Segments: []string{validHash, validHash}, + }, + wantError: "duplicate segment reference", + }, + { + name: "too many segment references", + manifest: manifest{ + Segments: syntheticSegmentHashes(17), + }, + wantError: "segment reference limit", + }, + { + name: "too many usage events", + manifest: manifest{ + Segments: []string{validHash}, + UsageEvents: make([]artifactUsageEvent, 32_769), + }, + wantError: "usage event limit", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + store := newTestArtifactStore(t) + m := tt.manifest + m.Version = formatVersion + m.Origin = origin + m.NativeSessionID = "sess-1" + m.Session = manifestSession{ID: "sess-1", Machine: origin} + data, err := canonicalJSON(m) + require.NoError(t, err) + hash := hashHex(data) + + ref, err := createCompressedTestArtifact( + t, store, origin, KindManifests, hash, compressPeerTestData(t, data), + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), tt.wantError) + assertArtifactNotStored(t, store, ref) + }) + } +} + +func TestWriteArtifactRejectsSegmentRecordAmplification(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + data := syntheticSegmentRecords(t, 4_097) + hash := hashHex(data) + + ref, err := createCompressedTestArtifact( + t, store, origin, KindSegments, hash, compressPeerTestData(t, data), + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), "message record limit") + assertArtifactNotStored(t, store, ref) +} + +func TestWriteArtifactRejectsSegmentNestedAmplificationWithoutWriting(t *testing.T) { + origin := "peer-a1b2c3" + tests := []struct { + name string + record segmentMessage + wantError string + }{ + { + name: "too many tool calls in one message", + record: segmentMessage{ + ToolCalls: make([]segmentToolCall, 257), + }, + wantError: "tool call limit", + }, + { + name: "too many result events in one tool call", + record: segmentMessage{ + ToolCalls: []segmentToolCall{{ + ResultEvents: make([]segmentResultEvent, 1_025), + }}, + }, + wantError: "result event limit", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + store := newTestArtifactStore(t) + record := tt.record + record.Version = formatVersion + record.Ordinal = 7 + record.Role = "assistant" + data, err := canonicalJSON(record) + require.NoError(t, err) + hash := hashHex(data) + + ref, err := createCompressedTestArtifact( + t, store, origin, KindSegments, hash, compressPeerTestData(t, data), + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), tt.wantError) + assertArtifactNotStored(t, store, ref) + }) + } +} + +func TestWriteArtifactRejectsBlankSegmentRecordsWithoutWriting(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + data := bytes.Repeat([]byte("\n"), 4_097) + hash := hashHex(data) + + ref, err := createCompressedTestArtifact( + t, store, origin, KindSegments, hash, compressPeerTestData(t, data), + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), "blank message record") + assertArtifactNotStored(t, store, ref) +} + +func TestExportChunksOnMessageRecordLimit(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + msgs := make([]db.Message, 4_097) + for i := range msgs { + msgs[i] = db.Message{SessionID: "sess-1", Ordinal: i, Role: "user"} + } + require.NoError(t, database.ReplaceSessionMessages("sess-1", msgs)) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + m := latestTestStoreManifest(t, store, origin) + require.Len(t, m.Segments, 2) + got := testStoreManifestMessages(t, store, origin, m) + assert.Len(t, got, 4_097) +} + +func TestExportRejectsOversizedGeneratedManifestBeforePublication(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha", func(sess *db.Session) { + first := strings.Repeat("x", int(manifestDecodedLimit)) + sess.FirstMessage = &first + }) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), "generated manifest exceeds") + assertNoPublishedAuthority(t, store, origin) +} + +func TestExportRejectsSessionMessageAmplificationBeforePublication(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + msgs := make([]db.Message, 32_769) + for i := range msgs { + msgs[i] = db.Message{SessionID: "sess-1", Ordinal: i, Role: "user"} + } + require.NoError(t, database.ReplaceSessionMessages("sess-1", msgs)) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), "session message limit") + assertNoPublishedArtifacts(t, store, origin) +} + +func TestExportRejectsUsageEventAmplificationBeforePublication(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + events := make([]db.UsageEvent, 32_769) + for i := range events { + events[i] = db.UsageEvent{ + SessionID: "sess-1", + Source: "fixture", + DedupKey: fmt.Sprintf("usage-%d", i), + } + } + require.NoError(t, database.ReplaceSessionUsageEvents("sess-1", events)) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), "usage event limit") + assertNoPublishedArtifacts(t, store, origin) +} + +func TestExportSessionRejectsAggregateLimitsBeforeWritingWithSmallLimits(t *testing.T) { + tests := []struct { + name string + configure func(*artifactLimits) + wantError string + }{ + { + name: "decoded bytes", + configure: func(limits *artifactLimits) { + limits.sessionDecodedBytes = 1 + }, + wantError: "session decoded byte limit", + }, + { + name: "segment references", + configure: func(limits *artifactLimits) { + limits.segmentMessages = 1 + limits.manifestSegments = 1 + }, + wantError: "segment reference limit", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + limits := productionArtifactLimits() + tt.configure(&limits) + + sess, err := database.GetSessionFull(ctx, "sess-1") + require.NoError(t, err) + require.NotNil(t, sess) + messages, err := database.GetAllMessages(ctx, "sess-1") + require.NoError(t, err) + _, _, err = exportLoadedSessionToStore( + ctx, store, origin, sess, messages, nil, limits, + ) + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantError) + assertNoPublishedAuthority(t, store, origin) + }) + } +} + +func TestImportQuarantinesOversizedSegmentWithoutAdvancingState(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + localOrigin := "desktop-d4e5f6" + importDB := testDB(t) + + prefix, err := canonicalJSON(segmentMessage{ + Version: formatVersion, + Ordinal: 0, + Role: "user", + Content: "hello", + ContentLength: 5, + }) + require.NoError(t, err) + segmentData := append(prefix, bytes.Repeat([]byte("\n"), 64<<20+1-len(prefix))...) + segmentHash := hashHex(segmentData) + segmentRef, err := NewRef(origin, KindSegments, segmentHash+".ndjson") + require.NoError(t, err) + createContractArtifact(t, store, segmentRef, segmentData) + + gid := origin + "~sess-1" + m := manifest{ + Version: formatVersion, + Origin: origin, + NativeSessionID: "sess-1", + Session: manifestSession{ + ID: "sess-1", + Machine: origin, + }, + Segments: []string{segmentHash}, + } + manifestData, err := canonicalJSON(m) + require.NoError(t, err) + manifestHash := hashHex(manifestData) + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + createContractArtifact(t, store, manifestRef, manifestData) + cpData, err := canonicalJSON(checkpoint{ + Version: formatVersion, + Origin: origin, + Sequence: 1, + Sessions: map[string]string{gid: manifestHash}, + }) + require.NoError(t, err) + cpRef, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + createContractArtifact(t, store, cpRef, cpData) + + res, err := importResultFromTestStore(ctx, importDB, store, localOrigin) + require.NoError(t, err) + assert.False(t, res.Changed()) + state, err := importDB.GetSyncState(importStateKey(origin, gid)) + require.NoError(t, err) + assert.Empty(t, state) + got, err := importDB.GetSessionFull(ctx, gid) + require.NoError(t, err) + assert.Nil(t, got) + _, err = store.Stat(ctx, segmentRef) + assert.ErrorIs(t, err, ErrArtifactNotFound) +} + +func TestExportChunksLargeMultiMessageSessionInOrder(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + content := strings.Repeat("x", 9<<20) + msgs := make([]db.Message, 4) + for i := range msgs { + msgs[i] = db.Message{ + SessionID: "sess-1", + Ordinal: i, + Role: "user", + Content: content, + ContentLength: len(content), + } + } + require.NoError(t, database.ReplaceSessionMessages("sess-1", msgs)) + + exportResult, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + assert.Equal(t, 1, exportResult.ExportedSessions) + m := latestTestStoreManifest(t, store, origin) + require.Len(t, m.Segments, 2) + + got := testStoreManifestMessages(t, store, origin, m) + require.Len(t, got, 4) + for i := range got { + assert.Equal(t, i, got[i].Ordinal) + assert.Equal(t, content, got[i].Content) + } +} + +func TestExportRejectsSingleEncodedRecordAboveReadableLimit(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + content := strings.Repeat("x", 64<<20) + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{{ + SessionID: "sess-1", + Ordinal: 0, + Role: "user", + Content: content, + ContentLength: len(content), + }})) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), "encoded message record") + assert.Contains(t, err.Error(), "67108864-byte readable limit") + assertNoPublishedArtifacts(t, store, origin) +} + +func TestExportPreservesSmallSingleSegmentHash(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + msgs, err := database.GetAllMessages(ctx, "sess-1") + require.NoError(t, err) + segmentData, err := encodeSegment(canonicalMessages(msgs)) + require.NoError(t, err) + wantHash := hashHex(segmentData) + + _, err = ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + m := latestTestStoreManifest(t, store, origin) + assert.Equal(t, []string{wantHash}, m.Segments) +} + +type repeatedByteReader byte + +func (r repeatedByteReader) Read(p []byte) (int, error) { + for i := range p { + p[i] = byte(r) + } + return len(p), nil +} + +func syntheticSegmentHashes(count int) []string { + hashes := make([]string, count) + for i := range hashes { + hashes[i] = fmt.Sprintf("%064x", i+1) + } + return hashes +} + +func syntheticSegmentRecords(t *testing.T, count int) []byte { + return syntheticSegmentRecordsFrom(t, 0, count) +} + +func syntheticSegmentRecordsFrom(t *testing.T, start, count int) []byte { + t.Helper() + var data bytes.Buffer + for i := range count { + record, err := canonicalJSON(segmentMessage{ + Version: formatVersion, + Ordinal: start + i, + Role: "user", + }) + require.NoError(t, err) + _, err = data.Write(record) + require.NoError(t, err) + } + return data.Bytes() +} diff --git a/internal/artifact/export_usage_test.go b/internal/artifact/export_usage_test.go new file mode 100644 index 000000000..820b0b808 --- /dev/null +++ b/internal/artifact/export_usage_test.go @@ -0,0 +1,91 @@ +package artifact + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +func TestExportIncludesUsageOnlySession(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + database := testDB(t) + + // A usage-only session: owned, zero messages, with a usage event. The + // sidebar list filter (message_count > 0) hides it, so export must not rely + // on that path. + require.NoError(t, database.UpsertSession(db.Session{ + ID: "usage-1", + Project: "alpha", + Machine: "local", + Agent: "claude", + CreatedAt: "2026-06-14T01:02:03Z", + })) + require.NoError(t, database.ReplaceSessionUsageEvents("usage-1", []db.UsageEvent{{ + SessionID: "usage-1", + Source: "assistant", + Model: "claude-opus-4-8", + InputTokens: 10, + OutputTokens: 20, + OccurredAt: "2026-06-14T01:02:30Z", + DedupKey: "usage-1:0", + }})) + + exportResult, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + assert.Equal(t, 1, exportResult.ExportedSessions) + + gid := origin + "~usage-1" + _, cp, err := latestStoreCheckpointSummary(ctx, store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + assert.Contains(t, cp.Sessions, gid, "usage-only session must be in the checkpoint") + + // It imports into a peer as a real session with its usage events intact. + importDB := testDB(t) + res, err := importResultFromTestStore(ctx, importDB, store, "desktop-d4e5f6") + require.NoError(t, err) + assert.Equal(t, 1, res.Sessions) + + got, err := importDB.GetSessionFull(ctx, gid) + require.NoError(t, err) + require.NotNil(t, got) + events, err := importDB.GetUsageEvents(ctx, gid) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, 10, events[0].InputTokens) + assert.Equal(t, 20, events[0].OutputTokens) +} + +func TestExportSkipsDeletedAndForeignSessions(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + database := testDB(t) + + seedSession(t, database, "owned", "alpha") + // A soft-deleted owned session and a foreign-owned session are both excluded. + seedSession(t, database, "trashed", "alpha") + require.NoError(t, database.SoftDeleteSession("trashed")) + require.NoError(t, database.UpsertSession(db.Session{ + ID: "foreign", + Project: "alpha", + Machine: "desktop-d4e5f6", + Agent: "claude", + CreatedAt: "2026-06-14T01:02:03Z", + })) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + + _, cp, err := latestStoreCheckpointSummary(ctx, store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + assert.Contains(t, cp.Sessions, origin+"~owned") + assert.NotContains(t, cp.Sessions, origin+"~trashed") + assert.NotContains(t, cp.Sessions, origin+"~foreign") +} diff --git a/internal/artifact/format_test.go b/internal/artifact/format_test.go new file mode 100644 index 000000000..d740028d8 --- /dev/null +++ b/internal/artifact/format_test.go @@ -0,0 +1,195 @@ +package artifact + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +func TestCanonicalCheckpointGolden(t *testing.T) { + cp := checkpoint{ + Version: formatVersion, + Origin: "laptop-a1b2c3", + Sequence: 7, + Sessions: map[string]string{ + "laptop-a1b2c3~sess-b": "b222", + "laptop-a1b2c3~sess-a": "a111", + }, + } + + data, err := canonicalJSON(cp) + require.NoError(t, err) + + assert.Equal(t, + "{\"origin\":\"laptop-a1b2c3\",\"seq\":7,\"sessions\":{\"laptop-a1b2c3~sess-a\":\"a111\",\"laptop-a1b2c3~sess-b\":\"b222\"},\"v\":1}\n", + string(data), + ) + assert.Equal(t, "56fd64d35ebd700bfa2a50d41a97857b871e033e5c1e1d02dea55c25c7df7655", hashHex(data)) +} + +func TestCanonicalManifestGolden(t *testing.T) { + cost := 0.03125 + ordinal := 2 + parent := "parent-1" + name := "Fixture" + raw := rawSourceRef{ + Hash: "raw123", + Size: 4096, + MediaType: "application/jsonl", + Path: "claude/session.jsonl", + } + m := manifest{ + Version: formatVersion, + Origin: "laptop-a1b2c3", + NativeSessionID: "sess-1", + Session: manifestSession{ + ID: "sess-1", + Project: "alpha", + Machine: "laptop-a1b2c3", + Agent: "claude", + FirstMessage: new("hello"), + StartedAt: new("2026-06-14T01:02:03Z"), + EndedAt: new("2026-06-14T01:03:03Z"), + MessageCount: 2, + UserMessageCount: 1, + ParentSessionID: &parent, + RelationshipType: "subagent", + TotalOutputTokens: 42, + CreatedAt: "2026-06-14T01:02:03Z", + }, + SessionName: &name, + Segments: []string{"seg222", "seg111"}, + UsageEvents: []artifactUsageEvent{ + { + MessageOrdinal: &ordinal, + Source: "fixture", + Model: "claude-test", + InputTokens: 11, + OutputTokens: 7, + CostUSD: &cost, + CostStatus: "known", + CostSource: "fixture", + OccurredAt: "2026-06-14T01:02:04Z", + DedupKey: "usage-1", + }, + }, + RawSource: &raw, + DataVersion: 99, + Generation: 3, + SessionHasToolCalls: true, + SessionHasContextData: true, + SessionQualitySignals: &manifestQualitySignals{ + Version: 3, + ShortPromptCount: 2, + UnstructuredStart: true, + MissingSuccessCriteriaCount: 4, + MissingVerificationCount: 5, + DuplicatePromptCount: 6, + NoCodeContextCount: 7, + RunawayToolLoopCount: 1, + }, + } + + data, err := canonicalJSON(m) + require.NoError(t, err) + + assert.Equal(t, + "{\"data_version\":99,\"generation\":3,\"native_session_id\":\"sess-1\",\"origin\":\"laptop-a1b2c3\",\"raw_source\":{\"hash\":\"raw123\",\"media_type\":\"application/jsonl\",\"path\":\"claude/session.jsonl\",\"size\":4096},\"segments\":[\"seg222\",\"seg111\"],\"session\":{\"agent\":\"claude\",\"compaction_count\":0,\"consecutive_failure_max\":0,\"created_at\":\"2026-06-14T01:02:03Z\",\"edit_churn_count\":0,\"ended_at\":\"2026-06-14T01:03:03Z\",\"ended_with_role\":\"\",\"final_failure_streak\":0,\"first_message\":\"hello\",\"has_peak_context_tokens\":false,\"has_total_output_tokens\":false,\"id\":\"sess-1\",\"is_automated\":false,\"machine\":\"laptop-a1b2c3\",\"message_count\":2,\"mid_task_compaction_count\":0,\"outcome\":\"\",\"outcome_confidence\":\"\",\"parent_session_id\":\"parent-1\",\"peak_context_tokens\":0,\"project\":\"alpha\",\"relationship_type\":\"subagent\",\"secret_leak_count\":0,\"started_at\":\"2026-06-14T01:02:03Z\",\"tool_failure_signal_count\":0,\"tool_retry_count\":0,\"total_output_tokens\":42,\"user_message_count\":1},\"session_has_context_data\":true,\"session_has_tool_calls\":true,\"session_name\":\"Fixture\",\"session_quality_signals\":{\"duplicate_prompt_count\":6,\"missing_success_criteria_count\":4,\"missing_verification_count\":5,\"no_code_context_count\":7,\"runaway_tool_loop_count\":1,\"short_prompt_count\":2,\"unstructured_start\":true,\"version\":3},\"usage_events\":[{\"cost_source\":\"fixture\",\"cost_status\":\"known\",\"cost_usd\":0.03125,\"dedup_key\":\"usage-1\",\"input_tokens\":11,\"message_ordinal\":2,\"model\":\"claude-test\",\"occurred_at\":\"2026-06-14T01:02:04Z\",\"output_tokens\":7,\"source\":\"fixture\"}],\"v\":1}\n", + string(data), + ) + assert.Equal(t, "1a563d1b1642cf850bb2253643d8cb628a91499cbf48607a6c40b15af01b4a6f", hashHex(data)) +} + +func TestCanonicalMessageSegmentGolden(t *testing.T) { + msgs := []db.Message{ + { + ID: 99, + SessionID: "sess-1", + Ordinal: 2, + Role: "assistant", + Content: "world", + ContentLength: 5, + Timestamp: "2026-06-14T01:02:05Z", + HasToolUse: true, + Model: "claude-test", + TokenUsage: json.RawMessage(`{"output":2,"input":1}`), + OutputTokens: 2, + HasOutputTokens: true, + ClaudeMessageID: "msg-1", + ClaudeRequestID: "req-1", + SourceType: "jsonl", + SourceSubtype: "assistant", + SourceUUID: "uuid-msg-1", + SourceParentUUID: "uuid-parent", + ToolCalls: []db.ToolCall{ + { + MessageID: 99, + SessionID: "sess-1", + ToolName: "Read", + Category: "file", + ToolUseID: "tool-1", + InputJSON: "{\"file_path\":\"README.md\"}", + FilePath: "README.md", + ResultContentLength: 12, + ResultContent: "file content", + SubagentSessionID: "child-1", + ResultEvents: []db.ToolResultEvent{ + { + ToolUseID: "tool-1", + AgentID: "agent-1", + SubagentSessionID: "child-1", + Source: "tool_result", + Status: "success", + Content: "done", + ContentLength: 4, + Timestamp: "2026-06-14T01:02:06Z", + EventIndex: 0, + }, + }, + }, + }, + }, + } + + data, err := encodeSegment(msgs) + require.NoError(t, err) + + assert.Equal(t, + "{\"claude_message_id\":\"msg-1\",\"claude_request_id\":\"req-1\",\"content\":\"world\",\"content_length\":5,\"has_output_tokens\":true,\"has_tool_use\":true,\"model\":\"claude-test\",\"ordinal\":2,\"output_tokens\":2,\"role\":\"assistant\",\"source_parent_uuid\":\"uuid-parent\",\"source_subtype\":\"assistant\",\"source_type\":\"jsonl\",\"source_uuid\":\"uuid-msg-1\",\"timestamp\":\"2026-06-14T01:02:05Z\",\"token_usage\":{\"input\":1,\"output\":2},\"tool_calls\":[{\"call_index\":0,\"category\":\"file\",\"file_path\":\"README.md\",\"input_json\":\"{\\\"file_path\\\":\\\"README.md\\\"}\",\"result_content\":\"file content\",\"result_content_length\":12,\"result_events\":[{\"agent_id\":\"agent-1\",\"content\":\"done\",\"content_length\":4,\"event_index\":0,\"source\":\"tool_result\",\"status\":\"success\",\"subagent_session_id\":\"child-1\",\"timestamp\":\"2026-06-14T01:02:06Z\",\"tool_use_id\":\"tool-1\"}],\"subagent_session_id\":\"child-1\",\"tool_name\":\"Read\",\"tool_use_id\":\"tool-1\"}],\"v\":1}\n", + string(data), + ) + assert.NotContains(t, string(data), `"id"`) + assert.NotContains(t, string(data), `"session_id"`) + assert.NotContains(t, string(data), `"message_id"`) + assert.Equal(t, "f46c1edbc77dab4eb15f43bcb3ce196243c784445b07e9135c223b5d58c6dea5", hashHex(data)) +} + +func TestCanonicalMetadataEventGolden(t *testing.T) { + value := json.RawMessage(`{"display_name":"Renamed session"}`) + note := "remember this" + event := metadataEvent{ + Version: formatVersion, + HLC: "2026-06-14T010203.000000001Z-laptop-a1b2c3", + Origin: "laptop-a1b2c3", + SessionGID: "desktop-d4e5f6~sess-1", + Op: "rename", + Value: value, + Pin: &MetadataPin{ + SourceUUID: "uuid-msg-1", + Ordinal: 2, + Note: ¬e, + }, + } + + data, err := canonicalJSON(event) + require.NoError(t, err) + + assert.Equal(t, + "{\"hlc\":\"2026-06-14T010203.000000001Z-laptop-a1b2c3\",\"op\":\"rename\",\"origin\":\"laptop-a1b2c3\",\"pin\":{\"note\":\"remember this\",\"ordinal\":2,\"source_uuid\":\"uuid-msg-1\"},\"session_gid\":\"desktop-d4e5f6~sess-1\",\"v\":1,\"value\":{\"display_name\":\"Renamed session\"}}\n", + string(data), + ) + assert.Equal(t, "fcb36d602e56fe1616ba6e2f86e973adde4ef547e0ecf280b37eb534b60e4b71", hashHex(data)) +} diff --git a/internal/artifact/gc.go b/internal/artifact/gc.go new file mode 100644 index 000000000..dfa1630e6 --- /dev/null +++ b/internal/artifact/gc.go @@ -0,0 +1,436 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "time" +) + +const retentionPageSize = 256 + +// GCOptions configures conservative logical artifact retention. Store exposes +// canonical references; retention never inspects or removes physical files. +type GCOptions struct { + Store ArtifactStore + Grace time.Duration + QuarantineGrace time.Duration + DryRun bool + Now time.Time + Logf func(string, ...any) +} + +// GCResult summarizes one logical retention scan. +type GCResult struct { + DryRun bool + Origins int + SkippedOrigins int + Scanned int + Candidates int + Eligible int + KeptByGrace int + Deleted int + BytesEligible int64 + BytesDeleted int64 + + QuarantinedScanned int + QuarantinedEligible int + QuarantinedDeleted int + QuarantineSkipped bool +} + +type gcLiveRef struct { + kind Kind + name string +} + +type gcOriginClosure map[gcLiveRef]struct{} + +// GarbageCollect trashes canonical nodes unreachable from each origin's newest +// safe checkpoint. Origins without such a checkpoint are left untouched. +func GarbageCollect(ctx context.Context, opts GCOptions) (_ GCResult, retErr error) { + if ctx == nil { + return GCResult{}, fmt.Errorf("artifact gc context is required") + } + if opts.Store == nil { + return GCResult{}, fmt.Errorf("artifact gc store is required") + } + if opts.Grace < 0 { + return GCResult{}, fmt.Errorf("artifact gc grace must be >= 0") + } + if opts.QuarantineGrace < 0 { + return GCResult{}, fmt.Errorf("artifact quarantine grace must be >= 0") + } + now := opts.Now + if now.IsZero() { + now = time.Now() + } + result := GCResult{DryRun: opts.DryRun} + + origins, err := openStoreOriginIterator(ctx, opts.Store) + if err != nil { + return result, fmt.Errorf("listing artifact origins: %w", err) + } + defer func() { retErr = errors.Join(retErr, origins.Close()) }() + for { + if err := ctx.Err(); err != nil { + return result, err + } + page, nextErr := origins.Next(ctx, retentionPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return result, fmt.Errorf("listing artifact origins: %w", nextErr) + } + for _, origin := range page { + result.Origins++ + closure, safe, err := collectGCOriginClosure(ctx, opts.Store, origin) + if err != nil { + return result, fmt.Errorf("scanning artifact origin %s: %w", origin, err) + } + if !safe { + result.SkippedOrigins++ + logGC(opts, "artifact gc: skipping %s without a safe current checkpoint", origin) + continue + } + if err := collectGCOrigin(ctx, opts, origin, closure, now, &result); err != nil { + return result, fmt.Errorf("collecting artifact origin %s: %w", origin, err) + } + } + if errors.Is(nextErr, io.EOF) { + break + } + } + if err := collectGCQuarantine(ctx, opts, now, &result); err != nil { + return result, err + } + return result, nil +} + +func collectGCOriginClosure( + ctx context.Context, store ArtifactStore, origin string, +) (gcOriginClosure, bool, error) { + latest, found, err := latestGCCheckpoint(ctx, store, origin) + if err != nil || !found { + return nil, false, err + } + closure := make(gcOriginClosure) + if err := closure.indexCheckpoint(ctx, store, origin, latest); err != nil { + if isUnsafeRetentionArtifact(err) { + return nil, false, nil + } + return nil, false, err + } + return closure, true, nil +} + +func (c gcOriginClosure) indexCheckpoint( + ctx context.Context, store ArtifactStore, origin string, latest Entry, +) error { + data, err := readGCStoreArtifact(ctx, store, latest, checkpointDecodedLimit) + if err != nil { + return err + } + var cp checkpoint + if err := json.Unmarshal(data, &cp); err != nil { + return fmt.Errorf("%w: invalid checkpoint JSON", ErrArtifactInvalid) + } + canonical, err := canonicalJSON(cp) + if err != nil || !bytes.Equal(canonical, data) { + return fmt.Errorf("%w: checkpoint JSON is not canonical", ErrArtifactInvalid) + } + if err := validateCheckpoint(&cp, origin); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + if err := validateCheckpointSequenceIdentity(cp, latest.Ref.Name); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + c[gcLiveRef{KindCheckpoints, latest.Ref.Name}] = struct{}{} + manifestOwners := make(map[string]string, len(cp.Sessions)) + for gid, manifestHash := range cp.Sessions { + if err := ctx.Err(); err != nil { + return err + } + if owner, found := manifestOwners[manifestHash]; found && owner != gid { + return fmt.Errorf( + "%w: manifest %s is referenced by multiple sessions", ErrArtifactInvalid, manifestHash) + } + manifestOwners[manifestHash] = gid + manifestRef := Ref{Origin: origin, Kind: KindManifests, Name: manifestHash + ".json"} + manifestEntry, err := store.Stat(ctx, manifestRef) + if err != nil { + return err + } + manifestData, err := readGCStoreArtifact(ctx, store, manifestEntry, manifestDecodedLimit) + if err != nil { + return err + } + if hashHex(manifestData) != manifestHash { + return fmt.Errorf("%w: manifest hash mismatch", ErrArtifactInvalid) + } + manifest, err := decodeManifestWithLimits(manifestData, productionArtifactLimits()) + if err != nil || validateManifest(manifest, origin, gid) != nil { + return fmt.Errorf("%w: invalid manifest", ErrArtifactInvalid) + } + if err := validateManifestReferencesWithLimits(manifest, productionArtifactLimits()); err != nil { + return fmt.Errorf("%w: invalid manifest references: %v", ErrArtifactInvalid, err) + } + c[gcLiveRef{KindManifests, manifestRef.Name}] = struct{}{} + for _, segmentHash := range manifest.Segments { + key := gcLiveRef{KindSegments, segmentHash + ".ndjson"} + if _, found := c[key]; found { + continue + } + segmentRef := Ref{Origin: origin, Kind: key.kind, Name: key.name} + entry, err := store.Stat(ctx, segmentRef) + if err != nil { + return err + } + segmentData, err := readGCStoreArtifact(ctx, store, entry, segmentDecodedLimit) + if err != nil { + return err + } + if hashHex(segmentData) != segmentHash { + return fmt.Errorf("%w: segment hash mismatch", ErrArtifactInvalid) + } + if _, err := preflightSegmentData(segmentData, productionArtifactLimits()); err != nil { + return fmt.Errorf("%w: invalid segment: %v", ErrArtifactInvalid, err) + } + c[key] = struct{}{} + } + if manifest.RawSource != nil && manifest.RawSource.Hash != "" { + key := gcLiveRef{KindRaw, manifest.RawSource.Hash} + if _, found := c[key]; found { + continue + } + rawRef := Ref{Origin: origin, Kind: key.kind, Name: key.name} + entry, err := store.Stat(ctx, rawRef) + if err != nil { + return err + } + if manifest.RawSource.Size != 0 && entry.Identity.Size != manifest.RawSource.Size { + return fmt.Errorf("%w: raw source size mismatch", ErrArtifactInvalid) + } + if err := verifyGCStoreArtifact(ctx, store, entry); err != nil { + return err + } + c[key] = struct{}{} + } + } + return nil +} + +func (c gcOriginClosure) contains(ctx context.Context, kind Kind, name string) (bool, error) { + if err := ctx.Err(); err != nil { + return false, err + } + _, found := c[gcLiveRef{kind, name}] + return found, nil +} + +func latestGCCheckpoint( + ctx context.Context, store ArtifactStore, origin string, +) (_ Entry, found bool, retErr error) { + iterator, err := openStoreEntryIterator(ctx, store, origin, KindCheckpoints) + if err != nil { + return Entry{}, false, err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + var latest Entry + for { + entries, nextErr := iterator.Next(ctx, retentionPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return Entry{}, false, nextErr + } + for _, entry := range entries { + if _, err := checkpointSequence(entry.Ref.Name); err == nil { + latest, found = entry, true + } + } + if errors.Is(nextErr, io.EOF) { + return latest, found, nil + } + } +} + +func collectGCOrigin( + ctx context.Context, + opts GCOptions, + origin string, + closure gcOriginClosure, + now time.Time, + result *GCResult, +) error { + for _, kind := range []Kind{KindCheckpoints, KindManifests, KindSegments, KindRaw} { + iterator, err := openStoreEntryIterator(ctx, opts.Store, origin, kind) + if err != nil { + return err + } + for { + entries, nextErr := iterator.Next(ctx, retentionPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return errors.Join(nextErr, iterator.Close()) + } + for _, entry := range entries { + if err := ctx.Err(); err != nil { + return errors.Join(err, iterator.Close()) + } + result.Scanned++ + live, err := closure.contains(ctx, kind, entry.Ref.Name) + if err != nil { + return errors.Join(err, iterator.Close()) + } + if live { + continue + } + result.Candidates++ + if entry.Modified.Add(opts.Grace).After(now) { + result.KeptByGrace++ + continue + } + result.Eligible++ + result.BytesEligible += entry.Identity.Size + if opts.DryRun { + logGC(opts, "artifact gc: would trash %s/%s/%s (%d bytes)", + entry.Ref.Origin, entry.Ref.Kind, entry.Ref.Name, entry.Identity.Size) + continue + } + if err := opts.Store.Trash(ctx, entry.Ref); err != nil { + if errors.Is(err, ErrArtifactNotFound) { + continue + } + return errors.Join(err, iterator.Close()) + } + result.Deleted++ + result.BytesDeleted += entry.Identity.Size + logGC(opts, "artifact gc: trashed %s/%s/%s (%d bytes)", + entry.Ref.Origin, entry.Ref.Kind, entry.Ref.Name, entry.Identity.Size) + } + if errors.Is(nextErr, io.EOF) { + break + } + } + if err := iterator.Close(); err != nil { + return err + } + } + return nil +} + +func readGCStoreArtifact( + ctx context.Context, store ArtifactStore, entry Entry, limit int64, +) (_ []byte, retErr error) { + actual, reader, err := store.Open(ctx, entry.Ref) + if err != nil { + return nil, err + } + defer func() { retErr = errors.Join(retErr, reader.Close()) }() + if actual.Identity != entry.Identity || actual.Ref != entry.Ref { + return nil, fmt.Errorf("%w: artifact changed during retention read", ErrArtifactCorrupt) + } + data, err := io.ReadAll(io.LimitReader(reader, limit+1)) + if err != nil { + return nil, err + } + if int64(len(data)) > limit { + return nil, fmt.Errorf("%w: artifact exceeds retention read limit", ErrArtifactCorrupt) + } + if err := reader.Verify(); err != nil { + return nil, err + } + return data, nil +} + +func verifyGCStoreArtifact( + ctx context.Context, store ArtifactStore, entry Entry, +) (retErr error) { + actual, reader, err := store.Open(ctx, entry.Ref) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, reader.Close()) }() + if actual.Identity != entry.Identity || actual.Ref != entry.Ref { + return fmt.Errorf("%w: artifact changed during retention read", ErrArtifactCorrupt) + } + if _, err := io.Copy(io.Discard, reader); err != nil { + return err + } + return reader.Verify() +} + +func isUnsafeRetentionArtifact(err error) bool { + return errors.Is(err, ErrArtifactNotFound) || + errors.Is(err, ErrArtifactCorrupt) || + errors.Is(err, ErrArtifactUnsupported) || + errors.Is(err, ErrArtifactInvalid) +} + +func collectGCQuarantine( + ctx context.Context, opts GCOptions, now time.Time, result *GCResult, +) (retErr error) { + quarantine, ok := opts.Store.(ArtifactQuarantineStore) + if !ok { + result.QuarantineSkipped = true + logGC(opts, "artifact gc: quarantine retention unsupported by store") + return nil + } + iterator, err := quarantine.Quarantined(ctx) + if err != nil { + return fmt.Errorf("listing artifact quarantine: %w", err) + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + for { + if err := ctx.Err(); err != nil { + return err + } + items, nextErr := iterator.Next(ctx, retentionPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return fmt.Errorf("listing artifact quarantine: %w", nextErr) + } + for _, entry := range items { + if err := ctx.Err(); err != nil { + return err + } + result.QuarantinedScanned++ + if entry.Modified.Add(opts.QuarantineGrace).After(now) { + continue + } + result.QuarantinedEligible++ + if opts.DryRun { + continue + } + if err := quarantine.TrashQuarantined(ctx, entry.Token); err != nil { + if errors.Is(err, ErrArtifactNotFound) { + continue + } + return fmt.Errorf("trashing artifact quarantine: %w", err) + } + result.QuarantinedDeleted++ + } + if errors.Is(nextErr, io.EOF) { + return nil + } + } +} + +// QuarantinedEntry names one hidden logical artifact retained for diagnosis. +type QuarantinedEntry struct { + Token string + Ref Ref + Identity Identity + Modified time.Time +} + +// ArtifactQuarantineStore is an optional stable-page view of hidden quarantine +// nodes. Tokens are opaque and may only be passed back to TrashQuarantined. +type ArtifactQuarantineStore interface { + Quarantined(context.Context) (QuarantineIterator, error) + TrashQuarantined(context.Context, string) error +} + +func logGC(opts GCOptions, format string, args ...any) { + if opts.Logf != nil { + opts.Logf(format, args...) + } +} diff --git a/internal/artifact/gc_test.go b/internal/artifact/gc_test.go new file mode 100644 index 000000000..f7047bc2d --- /dev/null +++ b/internal/artifact/gc_test.go @@ -0,0 +1,652 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "runtime" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/docbank" +) + +const retentionOrigin = "retention-a1b2c3" + +type retentionFixture struct { + store ArtifactStore + old map[Kind][]Ref + live map[Kind][]Ref + meta Ref +} + +func TestGarbageCollectTrashesOnlySupersededLogicalClosure(t *testing.T) { + fixture := newRetentionFixture(t) + now := time.Now().Add(48 * time.Hour) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: fixture.store, + Grace: time.Hour, + Now: now, + }) + require.NoError(t, err) + + assert.Equal(t, 1, result.Origins) + assert.Zero(t, result.SkippedOrigins) + assert.Equal(t, 3, result.Eligible) + assert.Equal(t, 3, result.Deleted) + for _, refs := range fixture.old { + for _, ref := range refs { + _, err := fixture.store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound, ref.Name) + } + } + for _, refs := range fixture.live { + for _, ref := range refs { + _, err := fixture.store.Stat(t.Context(), ref) + assert.NoError(t, err, ref.Name) + } + } + _, err = fixture.store.Stat(t.Context(), fixture.meta) + assert.NoError(t, err, "metadata and purge tombstones are retained by logical GC") +} + +func TestGarbageCollectDryRunAndGracePreserveCandidates(t *testing.T) { + for _, tc := range []struct { + name string + dryRun bool + now func(time.Time) time.Time + want func(t *testing.T, result GCResult) + }{ + { + name: "dry run", + dryRun: true, + now: func(created time.Time) time.Time { return created.Add(48 * time.Hour) }, + want: func(t *testing.T, result GCResult) { + assert.Equal(t, 3, result.Eligible) + assert.Zero(t, result.Deleted) + }, + }, + { + name: "inside grace", + now: func(created time.Time) time.Time { return created.Add(30 * time.Minute) }, + want: func(t *testing.T, result GCResult) { + assert.Zero(t, result.Eligible) + assert.Equal(t, 3, result.KeptByGrace) + assert.Zero(t, result.Deleted) + }, + }, + } { + t.Run(tc.name, func(t *testing.T) { + fixture := newRetentionFixture(t) + created, err := fixture.store.Stat(t.Context(), fixture.old[KindCheckpoints][0]) + require.NoError(t, err) + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: fixture.store, + Grace: time.Hour, + DryRun: tc.dryRun, + Now: tc.now(created.Modified), + }) + require.NoError(t, err) + tc.want(t, result) + for _, refs := range fixture.old { + for _, ref := range refs { + _, err := fixture.store.Stat(t.Context(), ref) + assert.NoError(t, err, ref.Name) + } + } + }) + } +} + +func TestGarbageCollectPreservesRawSourceReachableFromLatestCheckpoint(t *testing.T) { + store := newRetentionStore(t) + database := testDB(t) + seedSession(t, database, "raw-session", "alpha") + _, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: retentionOrigin, Full: true, + }) + require.NoError(t, err) + + latest, found, err := latestGCCheckpoint(t.Context(), store, retentionOrigin) + require.NoError(t, err) + require.True(t, found) + checkpointData, err := readGCStoreArtifact(t.Context(), store, latest, checkpointDecodedLimit) + require.NoError(t, err) + var cp checkpoint + require.NoError(t, json.Unmarshal(checkpointData, &cp)) + gid := retentionOrigin + "~raw-session" + manifestRef := Ref{ + Origin: retentionOrigin, Kind: KindManifests, Name: cp.Sessions[gid] + ".json", + } + manifestEntry, err := store.Stat(t.Context(), manifestRef) + require.NoError(t, err) + manifestData, err := readGCStoreArtifact(t.Context(), store, manifestEntry, manifestDecodedLimit) + require.NoError(t, err) + var m manifest + require.NoError(t, json.Unmarshal(manifestData, &m)) + + rawBody := []byte("canonical raw source") + rawRef := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex(rawBody), + }, rawBody) + m.RawSource = &rawSourceRef{Hash: rawRef.Name, Size: int64(len(rawBody))} + updatedManifest, err := canonicalJSON(m) + require.NoError(t, err) + updatedManifestHash := hashHex(updatedManifest) + createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindManifests, Name: updatedManifestHash + ".json", + }, updatedManifest) + cp.Sequence++ + cp.Sessions[gid] = updatedManifestHash + updatedCheckpoint, err := canonicalJSON(cp) + require.NoError(t, err) + createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindCheckpoints, + Name: fmt.Sprintf("cp-%010d.json", cp.Sequence), + }, updatedCheckpoint) + staleRaw := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("stale raw")), + }, []byte("stale raw")) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, Now: time.Now().Add(time.Hour), + }) + require.NoError(t, err) + _, err = store.Stat(t.Context(), rawRef) + assert.NoError(t, err, "latest checkpoint raw dependency must remain live") + _, err = store.Stat(t.Context(), staleRaw) + assert.ErrorIs(t, err, ErrArtifactNotFound) + assert.Positive(t, result.Deleted) +} + +func TestGarbageCollectLeavesOriginsWithoutSafeCheckpointUntouched(t *testing.T) { + for _, tc := range []struct { + name string + checkpoint []byte + }{ + {name: "no checkpoint"}, + {name: "corrupt checkpoint", checkpoint: []byte(`{"v":1`)}, + {name: "future checkpoint", checkpoint: []byte(`{"v":2,"origin":"retention-a1b2c3","seq":1,"sessions":{}}\n`)}, + {name: "missing sessions field", checkpoint: []byte(`{"v":1,"origin":"retention-a1b2c3","seq":1}`)}, + {name: "incomplete checkpoint", checkpoint: []byte(`{"v":1,"origin":"retention-a1b2c3","seq":1,"sessions":{"retention-a1b2c3~missing":"aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}}\n`)}, + } { + t.Run(tc.name, func(t *testing.T) { + store := newRetentionStore(t) + orphan := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, + Kind: KindRaw, + Name: hashHex([]byte("orphan")), + }, []byte("orphan")) + if tc.checkpoint != nil { + createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, + Kind: KindCheckpoints, + Name: "cp-0000000001.json", + }, tc.checkpoint) + } + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, + Grace: 0, + Now: time.Now().Add(time.Hour), + }) + require.NoError(t, err) + assert.Equal(t, 1, result.SkippedOrigins) + _, err = store.Stat(t.Context(), orphan) + assert.NoError(t, err, "unsafe origins must remain untouched") + }) + } +} + +func TestGarbageCollectRejectsDuplicateCheckpointKeysBeforeTrashing(t *testing.T) { + t.Run("top-level sessions field", func(t *testing.T) { + store := newRetentionStore(t) + body := []byte(`{"v":1,"origin":"retention-a1b2c3","seq":1,"sessions":{},"sessions":{}}`) + createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindCheckpoints, Name: "cp-0000000001.json", + }, body) + orphan := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("top-level orphan")), + }, []byte("top-level orphan")) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, Now: time.Now().Add(time.Hour), + }) + + require.NoError(t, err) + assert.Equal(t, 1, result.SkippedOrigins) + _, err = store.Stat(t.Context(), orphan) + assert.NoError(t, err) + }) + + t.Run("session id", func(t *testing.T) { + store := newRetentionStore(t) + database := testDB(t) + seedSession(t, database, "duplicate-session", "alpha") + _, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: retentionOrigin, Full: true, + }) + require.NoError(t, err) + latest, found, err := latestGCCheckpoint(t.Context(), store, retentionOrigin) + require.NoError(t, err) + require.True(t, found) + data, err := readGCStoreArtifact(t.Context(), store, latest, checkpointDecodedLimit) + require.NoError(t, err) + var cp checkpoint + require.NoError(t, json.Unmarshal(data, &cp)) + gid := retentionOrigin + "~duplicate-session" + hash := cp.Sessions[gid] + body := fmt.Appendf(nil, + `{"v":1,"origin":%q,"seq":2,"sessions":{%q:%q,%q:%q}}`, + retentionOrigin, gid, hash, gid, hash) + createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindCheckpoints, Name: "cp-0000000002.json", + }, body) + orphan := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("session orphan")), + }, []byte("session orphan")) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, Now: time.Now().Add(time.Hour), + }) + + require.NoError(t, err) + assert.Equal(t, 1, result.SkippedOrigins) + _, err = store.Stat(t.Context(), orphan) + assert.NoError(t, err) + }) +} + +func TestGarbageCollectCancellationStopsBeforeLaterOrigin(t *testing.T) { + base := newRetentionStore(t) + for _, origin := range []string{"cancel-a1b2c3", "cancel-d4e5f6"} { + createEmptyCheckpoint(t, base, origin, 1) + createRetentionRef(t, base, Ref{ + Origin: origin, Kind: KindRaw, Name: hashHex([]byte(origin)), + }, []byte(origin)) + } + ctx, cancel := context.WithCancel(t.Context()) + store := &cancelAfterTrashStore{ArtifactStore: base, cancel: cancel} + + result, err := GarbageCollect(ctx, GCOptions{ + Store: store, + Grace: 0, + Now: time.Now().Add(time.Hour), + }) + require.ErrorIs(t, err, context.Canceled) + assert.Equal(t, 1, result.Deleted) + assert.Equal(t, 1, store.trashCalls) +} + +func TestGarbageCollectExpiresQuarantineThroughLogicalCapability(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + ref := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("quarantine")), + }, []byte("quarantine")) + require.NoError(t, store.Quarantine(t.Context(), ref, "invalid protocol body")) + items, err := firstStoreQuarantinePage(t.Context(), store, 1) + require.NoError(t, err) + require.Len(t, items, 1) + assert.Equal(t, ref, items[0].Ref) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, + QuarantineGrace: time.Hour, + Now: items[0].Modified.Add(2 * time.Hour), + }) + require.NoError(t, err) + assert.False(t, result.QuarantineSkipped) + assert.Equal(t, 1, result.QuarantinedScanned) + assert.Equal(t, 1, result.QuarantinedEligible) + assert.Equal(t, 1, result.QuarantinedDeleted) + items, err = firstStoreQuarantinePage(t.Context(), store, 1) + require.NoError(t, err) + assert.Empty(t, items) +} + +func TestGarbageCollectQuarantineDryRunDoesNotTrash(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + ref := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("dry quarantine")), + }, []byte("dry quarantine")) + require.NoError(t, store.Quarantine(t.Context(), ref, "diagnostic")) + items, err := firstStoreQuarantinePage(t.Context(), store, 1) + require.NoError(t, err) + require.Len(t, items, 1) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, DryRun: true, + QuarantineGrace: time.Hour, Now: items[0].Modified.Add(2 * time.Hour), + }) + + require.NoError(t, err) + assert.Equal(t, 1, result.QuarantinedEligible) + assert.Zero(t, result.QuarantinedDeleted) + items, err = firstStoreQuarantinePage(t.Context(), store, 1) + require.NoError(t, err) + assert.Len(t, items, 1) +} + +func TestGarbageCollectCheckpointReachabilityHeapStaysBounded(t *testing.T) { + small := measureCheckpointReachabilityHeap(t, 10) + large := measureCheckpointReachabilityHeap(t, 10_000) + + assert.LessOrEqual(t, large, small+int64(32<<20), + "10,000 checkpoint sessions must stay within the explicit GC memory budget") +} + +func TestGarbageCollectKeepsQuarantineInsideGraceAndReportsUnsupportedStore(t *testing.T) { + t.Run("inside grace", func(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + ref := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("recent")), + }, []byte("recent")) + require.NoError(t, store.Quarantine(t.Context(), ref, "recent")) + items, err := firstStoreQuarantinePage(t.Context(), store, 1) + require.NoError(t, err) + require.Len(t, items, 1) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, + QuarantineGrace: time.Hour, + Now: items[0].Modified.Add(30 * time.Minute), + }) + require.NoError(t, err) + assert.Zero(t, result.QuarantinedEligible) + assert.Zero(t, result.QuarantinedDeleted) + items, err = firstStoreQuarantinePage(t.Context(), store, 1) + require.NoError(t, err) + assert.Len(t, items, 1) + }) + + t.Run("unsupported", func(t *testing.T) { + store := struct{ ArtifactStore }{ArtifactStore: newRetentionStore(t)} + result, err := GarbageCollect(t.Context(), GCOptions{Store: store}) + require.NoError(t, err) + assert.True(t, result.QuarantineSkipped) + }) +} + +type cancelAfterTrashStore struct { + ArtifactStore + cancel context.CancelFunc + trashCalls int +} + +func (s *cancelAfterTrashStore) Trash(ctx context.Context, ref Ref) error { + s.trashCalls++ + err := s.ArtifactStore.Trash(ctx, ref) + s.cancel() + return err +} + +func newRetentionFixture(t *testing.T) retentionFixture { + t.Helper() + store := newRetentionStore(t) + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + _, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: retentionOrigin, + Full: true, + }) + require.NoError(t, err) + firstCheckpoint, found, err := latestGCCheckpoint(t.Context(), store, retentionOrigin) + require.NoError(t, err) + require.True(t, found) + firstCheckpointData, err := readGCStoreArtifact( + t.Context(), store, firstCheckpoint, checkpointDecodedLimit) + require.NoError(t, err) + var first checkpoint + require.NoError(t, json.Unmarshal(firstCheckpointData, &first)) + firstManifestHash := first.Sessions[retentionOrigin+"~sess-1"] + require.NotEmpty(t, firstManifestHash) + firstManifestRef := Ref{ + Origin: retentionOrigin, Kind: KindManifests, Name: firstManifestHash + ".json", + } + firstManifestEntry, err := store.Stat(t.Context(), firstManifestRef) + require.NoError(t, err) + firstManifestData, err := readGCStoreArtifact( + t.Context(), store, firstManifestEntry, manifestDecodedLimit) + require.NoError(t, err) + var firstManifest manifest + require.NoError(t, json.Unmarshal(firstManifestData, &firstManifest)) + require.Len(t, firstManifest.Segments, 1) + firstSegmentRef := Ref{ + Origin: retentionOrigin, Kind: KindSegments, + Name: firstManifest.Segments[0] + ".ndjson", + } + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "changed", ContentLength: 7}, + })) + _, err = ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: retentionOrigin, + Full: true, + }) + require.NoError(t, err) + latestCheckpoint, found, err := latestGCCheckpoint(t.Context(), store, retentionOrigin) + require.NoError(t, err) + require.True(t, found) + latestCheckpointData, err := readGCStoreArtifact( + t.Context(), store, latestCheckpoint, checkpointDecodedLimit) + require.NoError(t, err) + var latest checkpoint + require.NoError(t, json.Unmarshal(latestCheckpointData, &latest)) + latestManifestHash := latest.Sessions[retentionOrigin+"~sess-1"] + require.NotEmpty(t, latestManifestHash) + require.NotEqual(t, firstManifestHash, latestManifestHash) + latestManifestRef := Ref{ + Origin: retentionOrigin, Kind: KindManifests, Name: latestManifestHash + ".json", + } + latestManifestEntry, err := store.Stat(t.Context(), latestManifestRef) + require.NoError(t, err) + latestManifestData, err := readGCStoreArtifact( + t.Context(), store, latestManifestEntry, manifestDecodedLimit) + require.NoError(t, err) + var latestManifest manifest + require.NoError(t, json.Unmarshal(latestManifestData, &latestManifest)) + require.Len(t, latestManifest.Segments, 1) + latestSegmentRef := Ref{ + Origin: retentionOrigin, Kind: KindSegments, + Name: latestManifest.Segments[0] + ".ndjson", + } + require.NotEqual(t, firstSegmentRef, latestSegmentRef) + old := map[Kind][]Ref{ + KindCheckpoints: {firstCheckpoint.Ref}, + KindManifests: {firstManifestRef}, + KindSegments: {firstSegmentRef}, + } + live := map[Kind][]Ref{ + KindCheckpoints: {latestCheckpoint.Ref}, + KindManifests: {latestManifestRef}, + KindSegments: {latestSegmentRef}, + } + + metadata := metadataEvent{ + Version: formatVersion, HLC: "20260721T120000.000000000Z-000000-retention-a1b2c3", + Origin: retentionOrigin, SessionGID: retentionOrigin + "~sess-1", Op: MetadataOpPurge, + } + metadataBody, err := canonicalJSON(metadata) + require.NoError(t, err) + metaHash := hashHex(metadataBody) + meta := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindMeta, + Name: metadata.HLC + "-" + metaHash + ".json", + }, metadataBody) + + return retentionFixture{ + store: store, + old: old, + live: live, + meta: meta, + } +} + +func newRetentionStore(t *testing.T) ArtifactStore { + t.Helper() + _, store := newTestDocbankStore(t, docbank.Config{}) + return store +} + +func createRetentionRef(t *testing.T, store ArtifactStore, ref Ref, body []byte) Ref { + t.Helper() + identity, err := NewIdentity(hashHex(body), int64(len(body))) + require.NoError(t, err) + _, err = store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + return ref +} + +func createEmptyCheckpoint(t *testing.T, store ArtifactStore, origin string, sequence int) Ref { + t.Helper() + body, err := canonicalJSON(checkpoint{ + Version: formatVersion, Origin: origin, Sequence: sequence, + Sessions: map[string]string{}, + }) + require.NoError(t, err) + return createRetentionRef(t, store, Ref{ + Origin: origin, Kind: KindCheckpoints, + Name: fmt.Sprintf("cp-%010d.json", sequence), + }, body) +} + +func TestGarbageCollectRejectsInvalidOptions(t *testing.T) { + _, err := GarbageCollect(t.Context(), GCOptions{}) + assert.Error(t, err) + _, err = GarbageCollect(t.Context(), GCOptions{ + Store: newRetentionStore(t), Grace: -time.Second, + }) + assert.Error(t, err) +} + +func TestGarbageCollectRejectsCheckpointSequenceMismatch(t *testing.T) { + store := newRetentionStore(t) + body, err := json.Marshal(checkpoint{ + Version: formatVersion, Origin: retentionOrigin, Sequence: 2, + Sessions: map[string]string{}, + }) + require.NoError(t, err) + createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindCheckpoints, Name: "cp-0000000001.json", + }, append(body, '\n')) + orphan := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("keep")), + }, []byte("keep")) + + result, err := GarbageCollect(t.Context(), GCOptions{ + Store: store, Now: time.Now().Add(time.Hour), + }) + require.NoError(t, err) + assert.Equal(t, 1, result.SkippedOrigins) + _, err = store.Stat(t.Context(), orphan) + assert.NoError(t, err) +} + +func TestGarbageCollectPropagatesStoreFailure(t *testing.T) { + want := errors.New("list failed") + store := &failingRetentionStore{err: want} + _, err := GarbageCollect(t.Context(), GCOptions{Store: store}) + assert.ErrorIs(t, err, want) +} + +type failingRetentionStore struct { + ArtifactStore + err error +} + +func (s *failingRetentionStore) Origins(context.Context) (OriginIterator, error) { + return nil, s.err +} + +type checkpointHeapStore struct { + ArtifactStore + entry Entry + body []byte + baseline uint64 + peak uint64 +} + +func newCheckpointHeapStore(t *testing.T, sessions int) *checkpointHeapStore { + t.Helper() + const manifestHash = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" + body := bytes.NewBuffer(make([]byte, 0, sessions*140)) + _, _ = fmt.Fprintf(body, `{"v":1,"origin":%q,"seq":1,"sessions":{`, retentionOrigin) + for index := range sessions { + if index > 0 { + _ = body.WriteByte(',') + } + _, _ = fmt.Fprintf(body, `%q:%q`, fmt.Sprintf("%s~session-%05d", retentionOrigin, index), manifestHash) + } + _, _ = body.WriteString("}}") + data := body.Bytes() + identity, err := NewIdentity(hashHex(data), int64(len(data))) + require.NoError(t, err) + return &checkpointHeapStore{ + entry: Entry{ + Ref: Ref{Origin: retentionOrigin, Kind: KindCheckpoints, Name: "cp-0000000001.json"}, + Identity: identity, Modified: time.Now(), + }, + body: data, + } +} + +func (s *checkpointHeapStore) Origins(context.Context) (OriginIterator, error) { + return &testOriginIterator{next: func(context.Context, int) ([]string, error) { + return []string{retentionOrigin}, io.EOF + }}, nil +} + +func (s *checkpointHeapStore) Entries( + _ context.Context, _ string, kind Kind, +) (EntryIterator, error) { + return &testEntryIterator{next: func(context.Context, int) ([]Entry, error) { + if kind == KindCheckpoints { + return []Entry{s.entry}, io.EOF + } + return nil, io.EOF + }}, nil +} + +func (s *checkpointHeapStore) Open( + context.Context, Ref, +) (Entry, VerifiedReader, error) { + return s.entry, &checkpointHeapReader{Reader: bytes.NewReader(s.body)}, nil +} + +func (s *checkpointHeapStore) Stat(context.Context, Ref) (Entry, error) { + runtime.GC() + var stats runtime.MemStats + runtime.ReadMemStats(&stats) + if stats.HeapAlloc > s.peak { + s.peak = stats.HeapAlloc + } + return Entry{}, ErrArtifactNotFound +} + +type checkpointHeapReader struct{ *bytes.Reader } + +func (r *checkpointHeapReader) Verify() error { return nil } +func (r *checkpointHeapReader) Close() error { return nil } + +func measureCheckpointReachabilityHeap(t *testing.T, sessions int) int64 { + t.Helper() + store := newCheckpointHeapStore(t, sessions) + runtime.GC() + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + store.baseline = baseline.HeapAlloc + result, err := GarbageCollect(t.Context(), GCOptions{Store: store}) + require.NoError(t, err) + assert.Equal(t, 1, result.SkippedOrigins) + if store.peak <= store.baseline { + return 0 + } + return int64(store.peak - store.baseline) +} diff --git a/internal/artifact/hlc.go b/internal/artifact/hlc.go new file mode 100644 index 000000000..32cf7f205 --- /dev/null +++ b/internal/artifact/hlc.go @@ -0,0 +1,267 @@ +package artifact + +import ( + "errors" + "fmt" + "strconv" + "strings" + "sync" + "time" +) + +const ( + metadataHLCStateKey = "artifact_metadata_hlc" + defaultMetadataHLCMaxDrift = 5 * time.Minute + // hlcWallLayout deliberately omits the ":" separators of RFC3339 so the + // rendered timestamp is safe to embed directly in artifact filenames. + // Windows forbids ":" in path components, and metadata event files are + // named after the HLC. Fixed-width fields keep the result lexicographically + // sortable and round-trippable through ParseHLCTimestamp. + hlcWallLayout = "2006-01-02T150405.000000000Z" + hlcLogicalWidth = 20 +) + +// ErrHLCDrift identifies a valid remote or persisted timestamp that cannot be +// observed yet because it exceeds the configured wall-clock drift bound. +var ErrHLCDrift = errors.New("metadata HLC exceeds drift bound") + +type hlcStateStore interface { + GetSyncState(key string) (string, error) + SetSyncState(key, value string) error +} + +// HLCTimestamp is a hybrid logical clock value for metadata events. +type HLCTimestamp struct { + WallTime time.Time + Logical uint64 +} + +// String formats the timestamp in a lexicographically sortable form. +func (t HLCTimestamp) String() string { + return fmt.Sprintf( + "%s-%0*d", + normalizeHLCWallTime(t.WallTime).Format(hlcWallLayout), + hlcLogicalWidth, + t.Logical, + ) +} + +// ParseHLCTimestamp parses a timestamp produced by HLCTimestamp.String. +func ParseHLCTimestamp(s string) (HLCTimestamp, error) { + idx := strings.LastIndex(s, "-") + if idx < 0 { + return HLCTimestamp{}, fmt.Errorf("invalid HLC timestamp %q: missing logical counter", s) + } + wallPart := s[:idx] + logicalPart := s[idx+1:] + if len(logicalPart) != hlcLogicalWidth || !isDecimal(logicalPart) { + return HLCTimestamp{}, fmt.Errorf("invalid HLC timestamp %q: logical counter must be %d digits", s, hlcLogicalWidth) + } + wall, err := time.Parse(hlcWallLayout, wallPart) + if err != nil { + return HLCTimestamp{}, fmt.Errorf("invalid HLC timestamp %q: %w", s, err) + } + logical, err := strconv.ParseUint(logicalPart, 10, 64) + if err != nil { + return HLCTimestamp{}, fmt.Errorf("invalid HLC timestamp %q: %w", s, err) + } + return HLCTimestamp{ + WallTime: normalizeHLCWallTime(wall), + Logical: logical, + }, nil +} + +// Compare returns -1, 0, or 1 when t is ordered before, equal to, or after other. +func (t HLCTimestamp) Compare(other HLCTimestamp) int { + wall := normalizeHLCWallTime(t.WallTime) + otherWall := normalizeHLCWallTime(other.WallTime) + switch { + case wall.Before(otherWall): + return -1 + case wall.After(otherWall): + return 1 + case t.Logical < other.Logical: + return -1 + case t.Logical > other.Logical: + return 1 + default: + return 0 + } +} + +// OrderingKey appends a deterministic tie-breaker, usually the artifact hash. +func (t HLCTimestamp) OrderingKey(tieBreaker string) string { + return t.String() + "-" + tieBreaker +} + +// HLCClockOptions configures a persisted metadata HLC clock. +type HLCClockOptions struct { + StateKey string + Now func() time.Time + MaxDrift time.Duration +} + +// HLCClock persists a monotonic hybrid logical clock in the sync-state store. +type HLCClock struct { + mu sync.Mutex + store hlcStateStore + stateKey string + now func() time.Time + maxDrift time.Duration +} + +// NewHLCClock returns a persisted metadata HLC clock. +func NewHLCClock(store hlcStateStore, opts HLCClockOptions) *HLCClock { + stateKey := opts.StateKey + if stateKey == "" { + stateKey = metadataHLCStateKey + } + now := opts.Now + if now == nil { + now = time.Now + } + maxDrift := opts.MaxDrift + if maxDrift <= 0 { + maxDrift = defaultMetadataHLCMaxDrift + } + return &HLCClock{ + store: store, + stateKey: stateKey, + now: now, + maxDrift: maxDrift, + } +} + +// Next returns and persists the next local metadata-event timestamp. +func (c *HLCClock) Next() (HLCTimestamp, error) { + c.mu.Lock() + defer c.mu.Unlock() + + last, ok, err := c.load() + if err != nil { + return HLCTimestamp{}, err + } + now := c.currentWallTime() + if ok { + if err := c.checkPersistedDrift(last, now); err != nil { + return HLCTimestamp{}, err + } + if !now.After(last.WallTime) { + next := HLCTimestamp{WallTime: last.WallTime, Logical: last.Logical + 1} + return next, c.persist(next) + } + } + next := HLCTimestamp{WallTime: now} + return next, c.persist(next) +} + +// Observe returns and persists a timestamp that is after the local and remote HLCs. +func (c *HLCClock) Observe(remote HLCTimestamp) (HLCTimestamp, error) { + c.mu.Lock() + defer c.mu.Unlock() + + last, ok, err := c.load() + if err != nil { + return HLCTimestamp{}, err + } + now := c.currentWallTime() + remote = HLCTimestamp{ + WallTime: normalizeHLCWallTime(remote.WallTime), + Logical: remote.Logical, + } + if ok { + if err := c.checkPersistedDrift(last, now); err != nil { + return HLCTimestamp{}, err + } + } + if remote.WallTime.After(now.Add(c.maxDrift)) { + return HLCTimestamp{}, fmt.Errorf( + "%w: remote HLC wall time %s is more than %s ahead of local time %s", + ErrHLCDrift, + remote.WallTime.Format(hlcWallLayout), + c.maxDrift, + now.Format(hlcWallLayout), + ) + } + + next := mergeHLC(HLCTimestamp{WallTime: now}, last, ok, remote) + return next, c.persist(next) +} + +func (c *HLCClock) load() (HLCTimestamp, bool, error) { + if c.store == nil { + return HLCTimestamp{}, false, errors.New("HLC state store is required") + } + raw, err := c.store.GetSyncState(c.stateKey) + if err != nil { + return HLCTimestamp{}, false, fmt.Errorf("reading HLC state: %w", err) + } + if strings.TrimSpace(raw) == "" { + return HLCTimestamp{}, false, nil + } + stamp, err := ParseHLCTimestamp(raw) + if err != nil { + return HLCTimestamp{}, false, fmt.Errorf("reading HLC state: %w", err) + } + return stamp, true, nil +} + +func (c *HLCClock) persist(stamp HLCTimestamp) error { + if err := c.store.SetSyncState(c.stateKey, stamp.String()); err != nil { + return fmt.Errorf("persisting HLC state: %w", err) + } + return nil +} + +func (c *HLCClock) currentWallTime() time.Time { + return normalizeHLCWallTime(c.now()) +} + +func (c *HLCClock) checkPersistedDrift(last HLCTimestamp, now time.Time) error { + if last.WallTime.After(now.Add(c.maxDrift)) { + return fmt.Errorf( + "%w: persisted HLC wall time %s is more than %s ahead of local time %s", + ErrHLCDrift, + last.WallTime.Format(hlcWallLayout), + c.maxDrift, + now.Format(hlcWallLayout), + ) + } + return nil +} + +func mergeHLC(now, last HLCTimestamp, hasLast bool, remote HLCTimestamp) HLCTimestamp { + maxWall := now.WallTime + if hasLast && last.WallTime.After(maxWall) { + maxWall = last.WallTime + } + if remote.WallTime.After(maxWall) { + maxWall = remote.WallTime + } + + next := HLCTimestamp{WallTime: maxWall} + if hasLast && last.WallTime.Equal(maxWall) { + next.Logical = last.Logical + } + if remote.WallTime.Equal(maxWall) && remote.Logical > next.Logical { + next.Logical = remote.Logical + } + if now.WallTime.Equal(maxWall) && next.Logical == 0 { + return next + } + next.Logical++ + return next +} + +func normalizeHLCWallTime(t time.Time) time.Time { + return t.UTC().Round(0) +} + +func isDecimal(s string) bool { + for _, r := range s { + if r < '0' || r > '9' { + return false + } + } + return s != "" +} diff --git a/internal/artifact/hlc_test.go b/internal/artifact/hlc_test.go new file mode 100644 index 000000000..f7fa4e351 --- /dev/null +++ b/internal/artifact/hlc_test.go @@ -0,0 +1,225 @@ +package artifact + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestHLCClockNextPersistsAcrossRestarts(t *testing.T) { + database := testDB(t) + now := fixedHLCTime() + + firstClock := NewHLCClock(database, HLCClockOptions{ + Now: func() time.Time { return now }, + MaxDrift: 5 * time.Minute, + }) + first, err := firstClock.Next() + require.NoError(t, err) + assert.Equal(t, now, first.WallTime) + assert.Equal(t, uint64(0), first.Logical) + + persisted, err := database.GetSyncState(metadataHLCStateKey) + require.NoError(t, err) + assert.Equal(t, "2026-06-14T010203.000000001Z-00000000000000000000", persisted) + + secondClock := NewHLCClock(database, HLCClockOptions{ + Now: func() time.Time { return now }, + MaxDrift: 5 * time.Minute, + }) + second, err := secondClock.Next() + require.NoError(t, err) + assert.Equal(t, now, second.WallTime) + assert.Equal(t, uint64(1), second.Logical) +} + +func TestHLCClockNextMonotonicCases(t *testing.T) { + base := fixedHLCTime() + + tests := []struct { + name string + last HLCTimestamp + now time.Time + want HLCTimestamp + }{ + { + name: "same wall time increments logical counter", + last: HLCTimestamp{WallTime: base, Logical: 7}, + now: base, + want: HLCTimestamp{WallTime: base, Logical: 8}, + }, + { + name: "physical time ahead resets logical counter", + last: HLCTimestamp{WallTime: base, Logical: 7}, + now: base.Add(time.Nanosecond), + want: HLCTimestamp{WallTime: base.Add(time.Nanosecond), Logical: 0}, + }, + { + name: "backward skew within bound increments logical counter", + last: HLCTimestamp{WallTime: base, Logical: 7}, + now: base.Add(-2 * time.Minute), + want: HLCTimestamp{WallTime: base, Logical: 8}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + database := testDB(t) + require.NoError(t, database.SetSyncState(metadataHLCStateKey, tt.last.String())) + clock := NewHLCClock(database, HLCClockOptions{ + Now: func() time.Time { return tt.now }, + MaxDrift: 5 * time.Minute, + }) + + got, err := clock.Next() + require.NoError(t, err) + + assert.Equal(t, 0, tt.want.Compare(got)) + persisted, err := database.GetSyncState(metadataHLCStateKey) + require.NoError(t, err) + assert.Equal(t, tt.want.String(), persisted) + }) + } +} + +func TestHLCClockNextRejectsBackwardSkewBeyondBound(t *testing.T) { + database := testDB(t) + base := fixedHLCTime() + last := HLCTimestamp{WallTime: base, Logical: 7} + require.NoError(t, database.SetSyncState(metadataHLCStateKey, last.String())) + clock := NewHLCClock(database, HLCClockOptions{ + Now: func() time.Time { return base.Add(-10 * time.Minute) }, + MaxDrift: 5 * time.Minute, + }) + + got, err := clock.Next() + require.Error(t, err) + assert.ErrorIs(t, err, ErrHLCDrift) + + assert.Equal(t, HLCTimestamp{}, got) + assert.Contains(t, err.Error(), "persisted HLC wall time") + persisted, err := database.GetSyncState(metadataHLCStateKey) + require.NoError(t, err) + assert.Equal(t, last.String(), persisted) +} + +func TestHLCClockObserveCases(t *testing.T) { + base := fixedHLCTime() + + tests := []struct { + name string + last *HLCTimestamp + now time.Time + remote HLCTimestamp + want HLCTimestamp + }{ + { + name: "remote future within bound is absorbed", + now: base, + remote: HLCTimestamp{WallTime: base.Add(time.Minute), Logical: 3}, + want: HLCTimestamp{WallTime: base.Add(time.Minute), Logical: 4}, + }, + { + name: "same wall uses max logical counter", + last: &HLCTimestamp{WallTime: base, Logical: 5}, + now: base, + remote: HLCTimestamp{WallTime: base, Logical: 7}, + want: HLCTimestamp{WallTime: base, Logical: 8}, + }, + { + name: "local physical time wins when ahead", + last: &HLCTimestamp{WallTime: base, Logical: 7}, + now: base.Add(time.Minute), + remote: HLCTimestamp{WallTime: base.Add(30 * time.Second), Logical: 9}, + want: HLCTimestamp{WallTime: base.Add(time.Minute), Logical: 0}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + database := testDB(t) + if tt.last != nil { + require.NoError(t, database.SetSyncState(metadataHLCStateKey, tt.last.String())) + } + clock := NewHLCClock(database, HLCClockOptions{ + Now: func() time.Time { return tt.now }, + MaxDrift: 5 * time.Minute, + }) + + got, err := clock.Observe(tt.remote) + require.NoError(t, err) + + assert.Equal(t, 0, tt.want.Compare(got)) + persisted, err := database.GetSyncState(metadataHLCStateKey) + require.NoError(t, err) + assert.Equal(t, tt.want.String(), persisted) + }) + } +} + +func TestHLCClockObserveRejectsRemoteFutureBeyondBound(t *testing.T) { + database := testDB(t) + base := fixedHLCTime() + last := HLCTimestamp{WallTime: base, Logical: 7} + require.NoError(t, database.SetSyncState(metadataHLCStateKey, last.String())) + clock := NewHLCClock(database, HLCClockOptions{ + Now: func() time.Time { return base }, + MaxDrift: 5 * time.Minute, + }) + + got, err := clock.Observe(HLCTimestamp{ + WallTime: base.Add(10 * time.Minute), + Logical: 3, + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrHLCDrift) + + assert.Equal(t, HLCTimestamp{}, got) + assert.Contains(t, err.Error(), "remote HLC wall time") + persisted, err := database.GetSyncState(metadataHLCStateKey) + require.NoError(t, err) + assert.Equal(t, last.String(), persisted) +} + +func TestHLCTimestampOrderingKeyAndParse(t *testing.T) { + base := fixedHLCTime() + stamp := HLCTimestamp{WallTime: base, Logical: 42} + + text := stamp.String() + parsed, err := ParseHLCTimestamp(text) + require.NoError(t, err) + + assert.Equal(t, "2026-06-14T010203.000000001Z-00000000000000000042", text) + assert.Equal(t, 0, stamp.Compare(parsed)) + assert.Equal(t, -1, stamp.Compare(HLCTimestamp{WallTime: base, Logical: 43})) + assert.Equal(t, -1, stamp.Compare(HLCTimestamp{WallTime: base.Add(time.Nanosecond), Logical: 0})) + assert.Equal(t, 1, stamp.Compare(HLCTimestamp{WallTime: base.Add(-time.Nanosecond), Logical: 99})) + assert.Less(t, stamp.OrderingKey("a111"), stamp.OrderingKey("b222")) + assert.Less(t, + HLCTimestamp{WallTime: base, Logical: 41}.OrderingKey("ffff"), + stamp.OrderingKey("0000"), + ) +} + +func TestParseHLCTimestampRejectsMalformedValues(t *testing.T) { + tests := []string{ + "", + "2026-06-14T010203Z-00000000000000000000", + "2026-06-14T010203.000000001Z-42", + "2026-06-14T010203.000000001Z-0000000000000000000x", + "2026-06-14T01:02:03.000000001Z-00000000000000000000", + } + + for _, tt := range tests { + t.Run(tt, func(t *testing.T) { + _, err := ParseHLCTimestamp(tt) + require.Error(t, err) + }) + } +} + +func fixedHLCTime() time.Time { + return time.Date(2026, 6, 14, 1, 2, 3, 1, time.UTC) +} diff --git a/internal/artifact/import_exact.go b/internal/artifact/import_exact.go new file mode 100644 index 000000000..be5d1c4f1 --- /dev/null +++ b/internal/artifact/import_exact.go @@ -0,0 +1,320 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "strings" + + "go.kenn.io/agentsview/internal/db" +) + +const artifactImportDrainLimit = 128 + +const ( + artifactImportReasonCheckpoint = "checkpoint dependencies incomplete" + artifactImportReasonMetadata = "metadata target unavailable" +) + +func artifactImportWork(entry Entry, reason string, requiredVersion int) db.ArtifactImportWork { + return db.ArtifactImportWork{ + Origin: entry.Ref.Origin, Kind: string(entry.Ref.Kind), Name: entry.Ref.Name, + SHA256: entry.Identity.SHA256, Size: entry.Identity.Size, + Reason: reason, RequiredFormatVersion: requiredVersion, + } +} + +// RecordChanged records only protocol references that can drive an import. +// Dependency arrivals merely request a bounded retry of already queued work. +func (c *StoreImportCoordinator) RecordChanged(ctx context.Context, entry Entry) error { + if c == nil || c.database == nil || c.store == nil { + return errors.New("artifact import coordinator is required") + } + if err := validateRefIdentity(entry.Ref, entry.Identity); err != nil { + return err + } + if entry.Ref.Origin == c.localOrigin { + return nil + } + switch entry.Ref.Kind { + case KindCheckpoints: + sequence, err := checkpointSequence(entry.Ref.Name) + if err != nil { + return err + } + head := db.ArtifactPeerCheckpointHead{ + Origin: entry.Ref.Origin, Sequence: sequence, + CheckpointSHA256: entry.Identity.SHA256, CheckpointSize: entry.Identity.Size, + } + if err := c.database.RecordArtifactPeerCheckpointHead(ctx, head); err != nil { + current, found, readErr := c.database.GetArtifactPeerCheckpointHead(ctx, entry.Ref.Origin) + if readErr != nil { + return errors.Join(err, readErr) + } + if !found || current.Sequence <= sequence { + return err + } + return c.requestDrain() + } + if err := c.database.EnqueueArtifactImport(ctx, + artifactImportWork(entry, artifactImportReasonCheckpoint, formatVersion)); err != nil { + return err + } + case KindMeta: + if err := c.database.EnqueueArtifactImport(ctx, + artifactImportWork(entry, artifactImportReasonMetadata, formatVersion)); err != nil { + return err + } + } + return c.requestDrain() +} + +func (c *StoreImportCoordinator) drainQueuedImports( + ctx context.Context, +) (ImportResult, error) { + work, err := c.database.PendingArtifactImports(ctx, formatVersion, artifactImportDrainLimit) + if err != nil { + return ImportResult{}, err + } + result := ImportResult{} + clock := NewHLCClock(c.database, HLCClockOptions{Now: c.now}) + // Content checkpoints must land before metadata from the same bounded + // page. Wire ordering places metadata ahead of checkpoints, while metadata + // projections may target sessions introduced by those checkpoints. + for _, phase := range []Kind{KindCheckpoints, KindMeta} { + for _, item := range work { + if Kind(item.Kind) != phase { + continue + } + if err := ctx.Err(); err != nil { + return result, err + } + var ( + itemResult ImportResult + acknowledge bool + ) + switch phase { + case KindCheckpoints: + itemResult, acknowledge, err = c.importQueuedCheckpoint(ctx, item) + case KindMeta: + itemResult.Metadata, acknowledge, err = c.importQueuedMetadata(ctx, clock, item) + } + result.Sessions += itemResult.Sessions + result.Messages += itemResult.Messages + result.Metadata += itemResult.Metadata + if err != nil { + return result, err + } + if acknowledge { + if _, err := c.database.AcknowledgeArtifactImport(ctx, item); err != nil { + return result, err + } + } + } + } + deferred, _, err := c.database.ArtifactImportQueueStats(ctx) + if err != nil { + return result, err + } + result.Deferred = deferred + return result, nil +} + +func queuedImportEntry(item db.ArtifactImportWork) (Entry, error) { + ref, err := NewRef(item.Origin, Kind(item.Kind), item.Name) + if err != nil { + return Entry{}, err + } + identity, err := NewIdentity(item.SHA256, item.Size) + if err != nil { + return Entry{}, err + } + entry := Entry{Ref: ref, Identity: identity} + if err := validateRefIdentity(ref, identity); err != nil { + return Entry{}, err + } + return entry, nil +} + +func (c *StoreImportCoordinator) importQueuedCheckpoint( + ctx context.Context, item db.ArtifactImportWork, +) (ImportResult, bool, error) { + entry, err := queuedImportEntry(item) + if err != nil { + return ImportResult{}, false, err + } + sequence, err := checkpointSequence(entry.Ref.Name) + if err != nil { + return ImportResult{}, false, err + } + head, found, err := c.database.GetArtifactPeerCheckpointHead(ctx, entry.Ref.Origin) + if err != nil { + return ImportResult{}, false, err + } + if found && head.Sequence > sequence { + return ImportResult{}, true, nil + } + landing, landed, err := c.database.GetArtifactCheckpointLandingHead(ctx, entry.Ref.Origin) + if err != nil { + return ImportResult{}, false, err + } + if landed && landing.Sequence == sequence && found && head.Sequence == sequence && + head.CheckpointSHA256 == entry.Identity.SHA256 && head.CheckpointSize == entry.Identity.Size { + return ImportResult{}, true, nil + } + data, err := readVerifiedStoreArtifact( + ctx, c.database, c.store, entry, checkpointDecodedLimit, + ) + if errors.Is(err, errIncompleteArtifact) { + return ImportResult{}, false, nil + } + if err != nil { + if errors.Is(err, ErrArtifactInvalid) { + qerr := c.store.Quarantine(ctx, entry.Ref, err.Error()) + return ImportResult{}, true, qerr + } + return ImportResult{}, false, fmt.Errorf("reading checkpoint %s: %w", entry.Ref.Name, err) + } + var cp checkpoint + if err := json.Unmarshal(data, &cp); err != nil { + qerr := c.store.Quarantine(ctx, entry.Ref, "checkpoint JSON is invalid") + return ImportResult{}, true, qerr + } + if cp.Version > formatVersion { + if err := c.database.EnqueueArtifactImport(ctx, + artifactImportWork(entry, artifactImportReasonCheckpoint, cp.Version)); err != nil { + return ImportResult{}, false, err + } + return ImportResult{}, false, nil + } + canonical, canonicalErr := canonicalJSON(cp) + if canonicalErr != nil || !bytes.Equal(canonical, data) { + qerr := c.store.Quarantine(ctx, entry.Ref, "checkpoint JSON is not canonical") + return ImportResult{}, true, qerr + } + if err := validateCheckpoint(&cp, entry.Ref.Origin); err != nil { + if errors.Is(err, errFutureArtifactVersion) { + return ImportResult{}, false, nil + } + qerr := c.store.Quarantine(ctx, entry.Ref, err.Error()) + return ImportResult{}, true, qerr + } + if err := validateCheckpointSequenceIdentity(cp, entry.Ref.Name); err != nil { + qerr := c.store.Quarantine(ctx, entry.Ref, err.Error()) + return ImportResult{}, true, qerr + } + outcome, err := inspectCheckpointClosureFromStore( + ctx, c.database, c.store, entry.Ref.Origin, cp, + ) + if err != nil { + return ImportResult{}, false, err + } + if outcome != checkpointClosureComplete { + return ImportResult{}, false, nil + } + result, err := importCheckpointFromStore(ctx, c.database, c.store, entry.Ref.Origin, cp) + if err != nil || result.Deferred > 0 { + return result, false, err + } + return result, true, nil +} + +func (c *StoreImportCoordinator) importQueuedMetadata( + ctx context.Context, clock *HLCClock, item db.ArtifactImportWork, +) (int, bool, error) { + entry, err := queuedImportEntry(item) + if err != nil { + return 0, false, err + } + orderKey, err := metadataArtifactOrderKey(entry.Ref.Name) + if err != nil { + qerr := c.store.Quarantine(ctx, entry.Ref, err.Error()) + return 0, true, errors.Join(err, qerr) + } + applied, err := c.database.MetadataEventApplied(ctx, entry.Ref.Origin, orderKey) + if err != nil || applied { + return 0, applied, err + } + data, err := readVerifiedStoreArtifact(ctx, c.database, c.store, entry, manifestDecodedLimit) + if errors.Is(err, errIncompleteArtifact) { + return 0, false, nil + } + if err != nil { + return 0, false, err + } + idx := strings.LastIndex(orderKey, "-") + if idx < 0 { + qerr := c.store.Quarantine(ctx, entry.Ref, "metadata filename lacks hash") + return 0, true, qerr + } + art := metadataArtifact{ + path: entry.Ref.Name, orderKey: orderKey, + hash: orderKey[idx+1:], hlc: orderKey[:idx], + } + var envelope metadataEventEnvelope + if err := json.Unmarshal(data, &envelope); err != nil { + qerr := c.store.Quarantine(ctx, entry.Ref, "metadata JSON is invalid") + return 0, true, qerr + } + if envelope.Version > formatVersion { + if err := c.database.EnqueueArtifactImport(ctx, + artifactImportWork(entry, artifactImportReasonMetadata, envelope.Version)); err != nil { + return 0, false, err + } + return 0, false, nil + } + if err := json.Unmarshal(data, &art.event); err != nil { + qerr := c.store.Quarantine(ctx, entry.Ref, "metadata JSON is invalid") + return 0, true, qerr + } + canonical, canonicalErr := canonicalJSON(art.event) + if canonicalErr != nil || !bytes.Equal(canonical, data) { + qerr := c.store.Quarantine(ctx, entry.Ref, "metadata JSON is not canonical") + return 0, true, qerr + } + stamp, err := ParseHLCTimestamp(art.hlc) + if err != nil { + qerr := markAppliedAndQuarantineMetadata( + ctx, c.database, c.store, entry.Ref, art, "metadata HLC is invalid", + ) + return 0, true, qerr + } + if err := validateMetadataArtifactEvent(art, entry.Ref.Origin); err != nil { + qerr := markAppliedAndQuarantineMetadata( + ctx, c.database, c.store, entry.Ref, art, err.Error(), + ) + return 0, true, qerr + } + if err := validateMetadataOp(art.event.Op); err != nil { + qerr := markAppliedAndQuarantineMetadata( + ctx, c.database, c.store, entry.Ref, art, "metadata operation is unsupported", + ) + return 0, true, qerr + } + projection, err := metadataProjection(art, c.localOrigin) + if err != nil { + qerr := markAppliedAndQuarantineMetadata( + ctx, c.database, c.store, entry.Ref, art, "metadata payload is invalid", + ) + return 0, true, qerr + } + if err := observeMetadataStamp(clock, stamp, art.hlc); err != nil { + if errors.Is(err, ErrHLCDrift) { + return 0, false, nil + } + return 0, false, err + } + result, err := c.database.ApplyMetadataProjection(ctx, projection) + if errors.Is(err, db.ErrMetadataTargetUnavailable) { + return 0, false, nil + } + if err != nil { + return 0, false, fmt.Errorf("replaying metadata event %s: %w", entry.Ref.Name, err) + } + if result.Applied || result.Conflict { + return 1, true, nil + } + return 0, true, nil +} diff --git a/internal/artifact/maintenance.go b/internal/artifact/maintenance.go new file mode 100644 index 000000000..a22c19efb --- /dev/null +++ b/internal/artifact/maintenance.go @@ -0,0 +1,240 @@ +package artifact + +import ( + "context" + "errors" + "fmt" + "sync" + "time" +) + +const ( + defaultPackPassBytes = int64(256 << 20) + defaultPackRetryDelay = 30 * time.Second +) + +type artifactPacker interface { + Pack(context.Context, int64) (PackResult, error) + LooseBacklog(context.Context) (LooseBacklog, error) +} + +type packSchedulerOptions struct { + RetryDelay time.Duration + Logf func(string, ...any) +} + +// packScheduler owns one optional background worker. Notifications are +// constant-work and coalesce in a one-slot channel; packing never runs on the +// caller's goroutine. +type packScheduler struct { + packer artifactPacker + logf func(string, ...any) + retry time.Duration + + ctx context.Context + cancel context.CancelFunc + wake chan struct{} + done chan struct{} + once sync.Once +} + +func newPackScheduler(packer artifactPacker, opts packSchedulerOptions) *packScheduler { + retry := opts.RetryDelay + if retry <= 0 { + retry = defaultPackRetryDelay + } + ctx, cancel := context.WithCancel(context.Background()) + scheduler := &packScheduler{ + packer: packer, + logf: opts.Logf, + retry: retry, + ctx: ctx, + cancel: cancel, + wake: make(chan struct{}, 1), + done: make(chan struct{}), + } + go scheduler.run() + return scheduler +} + +// Recover schedules a pass when Docbank reports any pack-eligible backlog. It +// uses Docbank's aggregate and never enumerates logical nodes or loose files. +func (s *packScheduler) Recover(ctx context.Context) error { + if s == nil || s.packer == nil { + return nil + } + backlog, err := s.packer.LooseBacklog(ctx) + if err != nil { + return err + } + if backlog.EligibleObjects > 0 || backlog.EligibleStoredBytes > 0 { + s.notify() + } + return nil +} + +func (s *packScheduler) Notify(ctx context.Context) { + if s == nil || s.packer == nil || artifactMaintenanceSuppressed(ctx) { + return + } + s.notify() +} + +func (s *packScheduler) notify() { + select { + case <-s.ctx.Done(): + return + default: + } + select { + case s.wake <- struct{}{}: + default: + } +} + +func (s *packScheduler) Close() { + if s == nil { + return + } + s.once.Do(func() { + s.cancel() + <-s.done + }) +} + +func (s *packScheduler) run() { + defer close(s.done) + var timer *time.Timer + var retry <-chan time.Time + for { + select { + case <-s.ctx.Done(): + if timer != nil { + timer.Stop() + } + return + case <-s.wake: + if timer != nil { + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + retry = nil + } + case <-retry: + retry = nil + } + + if err := s.ctx.Err(); err != nil { + return + } + result, err := s.packer.Pack(s.ctx, defaultPackPassBytes) + if err != nil { + if !errors.Is(err, context.Canceled) && s.logf != nil { + s.logf("artifact pack: %v", err) + } + continue + } + if result.More { + if timer == nil { + timer = time.NewTimer(s.retry) + } else { + timer.Reset(s.retry) + } + retry = timer.C + } + } +} + +type suppressArtifactMaintenanceKey struct{} + +// SuppressArtifactMaintenance marks a shutdown-flush context. Required export +// and exchange work proceeds, but successful completion cannot start optional +// physical packing after the flush budget has begun. +func SuppressArtifactMaintenance(ctx context.Context) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, suppressArtifactMaintenanceKey{}, true) +} + +func artifactMaintenanceSuppressed(ctx context.Context) bool { + return ctx != nil && ctx.Value(suppressArtifactMaintenanceKey{}) == true +} + +// ArtifactMaintenanceOptions keeps semantic deletion separate from bounded +// Docbank reclamation. Each stage performs at most one resumable pass. +type ArtifactMaintenanceOptions struct { + TrashGrace time.Duration + EmptyTrash WorkBudget + GC WorkBudget + Repack WorkBudget +} + +// PhysicalMaintenanceResult reports independently resumable physical stages. +type PhysicalMaintenanceResult struct { + EmptyTrash MaintenanceResult + GarbageCollect MaintenanceResult + Repack MaintenanceResult +} + +// ValidateArtifactMaintenanceOptions validates every caller-controlled +// physical maintenance limit before semantic retention can mutate the store. +func ValidateArtifactMaintenanceOptions(opts ArtifactMaintenanceOptions) error { + if opts.TrashGrace < 0 { + return fmt.Errorf( + "%w: physical trash grace must not be negative", ErrArtifactInvalid) + } + for _, budget := range []WorkBudget{opts.EmptyTrash, opts.GC, opts.Repack} { + if err := validateArtifactWorkBudget(budget); err != nil { + return err + } + } + if opts.EmptyTrash.MaxBytes != 0 { + return fmt.Errorf( + "%w: trash emptying supports only an object budget", ErrArtifactInvalid) + } + return nil +} + +type artifactMaintainer interface { + Verify(context.Context, WorkBudget) (MaintenanceResult, error) + EmptyTrash(context.Context, time.Duration, WorkBudget) (MaintenanceResult, error) + GarbageCollect(context.Context, WorkBudget) (MaintenanceResult, error) + Repack(context.Context, WorkBudget) (MaintenanceResult, error) +} + +// runPhysicalMaintenance performs one bounded pass of each physical stage. +// It never decides logical liveness and stops between stages on cancellation. +func runPhysicalMaintenance( + ctx context.Context, + maintainer artifactMaintainer, + opts ArtifactMaintenanceOptions, +) (PhysicalMaintenanceResult, error) { + if maintainer == nil { + return PhysicalMaintenanceResult{}, ErrArtifactUnsupported + } + if err := ValidateArtifactMaintenanceOptions(opts); err != nil { + return PhysicalMaintenanceResult{}, err + } + var result PhysicalMaintenanceResult + var err error + result.EmptyTrash, err = maintainer.EmptyTrash(ctx, opts.TrashGrace, opts.EmptyTrash) + if err != nil { + return result, err + } + if err := ctx.Err(); err != nil { + return result, err + } + result.GarbageCollect, err = maintainer.GarbageCollect(ctx, opts.GC) + if err != nil { + return result, err + } + if err := ctx.Err(); err != nil { + return result, err + } + result.Repack, err = maintainer.Repack(ctx, opts.Repack) + return result, errors.Join(err, ctx.Err()) +} diff --git a/internal/artifact/maintenance_test.go b/internal/artifact/maintenance_test.go new file mode 100644 index 000000000..6969d435a --- /dev/null +++ b/internal/artifact/maintenance_test.go @@ -0,0 +1,324 @@ +package artifact + +import ( + "context" + "errors" + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/docbank" +) + +func TestDocbankPhysicalMaintenanceRunsBoundedPhysicalStages(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + maintainer, ok := any(store).(artifactMaintainer) + require.True(t, ok) + ref := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex([]byte("physical")), + }, []byte("physical")) + + verified, err := maintainer.Verify(t.Context(), WorkBudget{MaxObjects: 1, MaxBytes: 1 << 20}) + require.NoError(t, err) + assert.Equal(t, 1, verified.Processed) + require.NoError(t, store.Trash(t.Context(), ref)) + emptied, err := maintainer.EmptyTrash(t.Context(), 0, WorkBudget{MaxObjects: 1}) + require.NoError(t, err) + assert.Equal(t, 1, emptied.Processed) + collected, err := maintainer.GarbageCollect(t.Context(), WorkBudget{MaxObjects: 1, MaxBytes: 1 << 20}) + require.NoError(t, err) + assert.Positive(t, collected.Processed) + _, err = maintainer.Repack(t.Context(), WorkBudget{MaxObjects: 1, MaxBytes: 1 << 20}) + require.NoError(t, err) +} + +func TestDocbankPhysicalMaintenanceResumesStages(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + maintainer := any(store).(artifactMaintainer) + for i := range 3 { + body := fmt.Appendf(nil, "trash-%d", i) + ref := createRetentionRef(t, store, Ref{ + Origin: retentionOrigin, Kind: KindRaw, Name: hashHex(body), + }, body) + require.NoError(t, store.Trash(t.Context(), ref)) + } + trash, err := maintainer.EmptyTrash(t.Context(), 0, WorkBudget{MaxObjects: 1}) + require.NoError(t, err) + passes := 1 + for trash.More { + trash, err = maintainer.EmptyTrash(t.Context(), 0, WorkBudget{ + MaxObjects: 1, Cursor: trash.NextCursor, + }) + require.NoError(t, err) + passes++ + } + assert.Equal(t, 3, passes) + + collected, err := maintainer.GarbageCollect(t.Context(), WorkBudget{MaxObjects: 1}) + require.NoError(t, err) + for collected.More { + collected, err = maintainer.GarbageCollect(t.Context(), WorkBudget{ + MaxObjects: 1, Cursor: collected.NextCursor, + }) + require.NoError(t, err) + } +} + +func TestDocbankPhysicalMaintenanceRejectsInvalidBudget(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + maintainer := any(store).(artifactMaintainer) + _, err := maintainer.GarbageCollect(t.Context(), WorkBudget{ + MaxObjects: docbank.MaxMaintenanceObjects + 1, + }) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.ErrorIs(t, mapDocbankError(docbank.ErrStaleRevision), ErrArtifactConflict) +} + +func TestDocbankRepackProcessedIncludesEveryMaintenanceAction(t *testing.T) { + report := docbank.RepackReport{ + MappingsPruned: 2, PacksSelected: 3, PacksRewritten: 5, + PacksSealed: 7, PacksRemoved: 11, PacksDeferredOversized: 13, + BlobsRepacked: 17, + } + assert.Equal(t, 58, docbankRepackProcessed(report)) +} + +func TestPackSchedulerCoalescesBatchNotifications(t *testing.T) { + packer := newBlockingPacker() + scheduler := newPackScheduler(packer, packSchedulerOptions{RetryDelay: time.Hour}) + t.Cleanup(scheduler.Close) + + returned := make(chan struct{}) + go func() { + scheduler.Notify(t.Context()) + close(returned) + }() + select { + case <-returned: + case <-time.After(time.Second): + require.Fail(t, "Notify blocked on physical packing") + } + require.Eventually(t, func() bool { return packer.calls.Load() == 1 }, time.Second, time.Millisecond) + for range 1_000 { + scheduler.Notify(t.Context()) + } + assert.Equal(t, int32(1), packer.maxConcurrent.Load()) + packer.release <- PackResult{} + require.Eventually(t, func() bool { return packer.calls.Load() == 2 }, time.Second, time.Millisecond) + packer.release <- PackResult{} + assert.Never(t, func() bool { return packer.calls.Load() > 2 }, 20*time.Millisecond, time.Millisecond) + assert.Equal(t, []int64{defaultPackPassBytes, defaultPackPassBytes}, packer.budgets()) +} + +func TestPackSchedulerRecoversIndexedBacklog(t *testing.T) { + packer := newBlockingPacker() + packer.setBacklog(LooseBacklog{EligibleObjects: 1, EligibleStoredBytes: 1}) + scheduler := newPackScheduler(packer, packSchedulerOptions{RetryDelay: time.Hour}) + t.Cleanup(scheduler.Close) + require.NoError(t, scheduler.Recover(t.Context())) + assert.Equal(t, int32(1), packer.backlogCalls.Load()) + require.Eventually(t, func() bool { return packer.calls.Load() == 1 }, time.Second, time.Millisecond) + packer.release <- PackResult{} +} + +func TestPackSchedulerRetriesMoreAndCloseCancels(t *testing.T) { + packer := newBlockingPacker() + scheduler := newPackScheduler(packer, packSchedulerOptions{RetryDelay: 10 * time.Millisecond}) + scheduler.Notify(t.Context()) + require.Eventually(t, func() bool { return packer.calls.Load() == 1 }, time.Second, time.Millisecond) + packer.release <- PackResult{More: true} + require.Eventually(t, func() bool { return packer.calls.Load() == 2 }, time.Second, time.Millisecond) + scheduler.Close() + assert.Eventually(t, func() bool { return packer.canceled.Load() == 1 }, time.Second, time.Millisecond) +} + +func TestPackSchedulerSuppressedContextDoesNotStartWork(t *testing.T) { + packer := newBlockingPacker() + scheduler := newPackScheduler(packer, packSchedulerOptions{RetryDelay: time.Hour}) + t.Cleanup(scheduler.Close) + scheduler.Notify(SuppressArtifactMaintenance(t.Context())) + assert.Never(t, func() bool { return packer.calls.Load() != 0 }, 20*time.Millisecond, time.Millisecond) +} + +func TestExportNotificationWorkDoesNotScaleWithArchiveCardinality(t *testing.T) { + small := exportNotificationCardinalityStats(t, 10) + large := exportNotificationCardinalityStats(t, 1_000) + want := artifactExportCallStats{claims: 1, sessions: 1, messages: 1, usage: 1} + assert.Equal(t, want, small) + assert.Equal(t, want, large) +} + +type blockingPacker struct { + calls atomic.Int32 + active atomic.Int32 + maxConcurrent atomic.Int32 + canceled atomic.Int32 + release chan PackResult + mu sync.Mutex + seenBudgets []int64 + backlog LooseBacklog + backlogCalls atomic.Int32 +} + +func newBlockingPacker() *blockingPacker { + return &blockingPacker{release: make(chan PackResult, 4)} +} + +func (p *blockingPacker) Pack(ctx context.Context, maxBytes int64) (PackResult, error) { + p.calls.Add(1) + active := p.active.Add(1) + defer p.active.Add(-1) + for maximum := p.maxConcurrent.Load(); active > maximum; maximum = p.maxConcurrent.Load() { + if p.maxConcurrent.CompareAndSwap(maximum, active) { + break + } + } + p.mu.Lock() + p.seenBudgets = append(p.seenBudgets, maxBytes) + p.mu.Unlock() + select { + case result := <-p.release: + return result, nil + case <-ctx.Done(): + p.canceled.Add(1) + return PackResult{}, ctx.Err() + } +} + +func (p *blockingPacker) LooseBacklog(context.Context) (LooseBacklog, error) { + p.backlogCalls.Add(1) + p.mu.Lock() + defer p.mu.Unlock() + return p.backlog, nil +} + +func (p *blockingPacker) budgets() []int64 { + p.mu.Lock() + defer p.mu.Unlock() + return append([]int64(nil), p.seenBudgets...) +} + +func (p *blockingPacker) setBacklog(backlog LooseBacklog) { + p.mu.Lock() + p.backlog = backlog + p.mu.Unlock() +} + +type artifactExportCallStats struct { + pending, claims, sessions, messages, usage, owned int +} + +type countingArtifactExportDB struct { + artifactExportStore + stats artifactExportCallStats +} + +func (d *countingArtifactExportDB) PendingArtifactExports( + ctx context.Context, limit int, +) ([]db.ArtifactExportQueueItem, error) { + d.stats.pending++ + return d.artifactExportStore.PendingArtifactExports(ctx, limit) +} + +func (d *countingArtifactExportDB) ArtifactExportClaims( + ctx context.Context, ids []string, +) ([]db.ArtifactExportQueueItem, error) { + d.stats.claims++ + return d.artifactExportStore.ArtifactExportClaims(ctx, ids) +} + +func (d *countingArtifactExportDB) GetSessionFull(ctx context.Context, id string) (*db.Session, error) { + d.stats.sessions++ + return d.artifactExportStore.GetSessionFull(ctx, id) +} + +func (d *countingArtifactExportDB) GetAllMessages(ctx context.Context, id string) ([]db.Message, error) { + d.stats.messages++ + return d.artifactExportStore.GetAllMessages(ctx, id) +} + +func (d *countingArtifactExportDB) GetUsageEvents(ctx context.Context, id string) ([]db.UsageEvent, error) { + d.stats.usage++ + return d.artifactExportStore.GetUsageEvents(ctx, id) +} + +func (d *countingArtifactExportDB) ListOwnedSessionIDsForExport(ctx context.Context) ([]string, error) { + d.stats.owned++ + return d.artifactExportStore.ListOwnedSessionIDsForExport(ctx) +} + +func exportNotificationCardinalityStats(t *testing.T, archiveSize int) artifactExportCallStats { + t.Helper() + database := testDB(t) + for index := range archiveSize - 1 { + _, err := database.WriteSessionBatch([]db.SessionBatchWrite{{ + Session: db.Session{ + ID: fmt.Sprintf("unrelated-%05d", index), Project: "archive", + Machine: "remote-origin", Agent: "claude", + }, + ReplaceMessages: true, + }}) + require.NoError(t, err) + } + seedSession(t, database, "changed-session", "project") + countingDB := &countingArtifactExportDB{artifactExportStore: database} + _, err := ExportToStore(t.Context(), countingDB, newRetentionStore(t), ExportOptions{ + Origin: retentionOrigin, SessionIDs: []string{"changed-session"}, + }) + require.NoError(t, err) + return countingDB.stats +} + +func TestRunPhysicalMaintenanceStopsOnCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + maintainer := &recordingMaintainer{cancel: cancel} + result, err := runPhysicalMaintenance(ctx, maintainer, ArtifactMaintenanceOptions{ + TrashGrace: time.Hour, + EmptyTrash: WorkBudget{MaxObjects: 2}, + GC: WorkBudget{MaxObjects: 3, MaxBytes: 4}, + Repack: WorkBudget{MaxObjects: 5, MaxBytes: 6}, + }) + require.ErrorIs(t, err, context.Canceled) + assert.Equal(t, []string{"empty"}, maintainer.calls) + assert.Equal(t, MaintenanceResult{Processed: 1}, result.EmptyTrash) +} + +func TestRunPhysicalMaintenanceRejectsInvalidInputs(t *testing.T) { + _, err := runPhysicalMaintenance(t.Context(), nil, ArtifactMaintenanceOptions{}) + assert.ErrorIs(t, err, ErrArtifactUnsupported) + _, err = runPhysicalMaintenance(t.Context(), &recordingMaintainer{}, ArtifactMaintenanceOptions{ + TrashGrace: -time.Second, + }) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +type recordingMaintainer struct { + calls []string + cancel context.CancelFunc +} + +func (m *recordingMaintainer) Verify(context.Context, WorkBudget) (MaintenanceResult, error) { + return MaintenanceResult{}, errors.New("unexpected verify") +} +func (m *recordingMaintainer) EmptyTrash( + context.Context, time.Duration, WorkBudget, +) (MaintenanceResult, error) { + m.calls = append(m.calls, "empty") + if m.cancel != nil { + m.cancel() + } + return MaintenanceResult{Processed: 1}, nil +} +func (m *recordingMaintainer) GarbageCollect(context.Context, WorkBudget) (MaintenanceResult, error) { + m.calls = append(m.calls, "gc") + return MaintenanceResult{Processed: 2}, nil +} +func (m *recordingMaintainer) Repack(context.Context, WorkBudget) (MaintenanceResult, error) { + m.calls = append(m.calls, "repack") + return MaintenanceResult{Processed: 3}, nil +} diff --git a/internal/artifact/manifest_session.go b/internal/artifact/manifest_session.go new file mode 100644 index 000000000..0dc9219af --- /dev/null +++ b/internal/artifact/manifest_session.go @@ -0,0 +1,236 @@ +package artifact + +import "go.kenn.io/agentsview/internal/db" + +// manifestSession is the manifest wire representation of a session row. It is +// deliberately a separate type from db.Session: manifest bytes feed content +// hashes, so this struct is part of the pinned artifact format and its +// serialized output can never change for an existing value. Keeping it apart +// from the internal DB type means adding a db.Session field cannot silently +// re-hash every exported manifest; extending THIS struct is an explicit wire +// format decision (see TestManifestSessionMatchesDBSessionWireFormat). +type manifestSession struct { + ID string `json:"id"` + Project string `json:"project"` + Machine string `json:"machine"` + Agent string `json:"agent"` + AgentLabel string `json:"agent_label,omitempty"` + Entrypoint string `json:"entrypoint,omitempty"` + FirstMessage *string `json:"first_message"` + DisplayName *string `json:"display_name,omitempty"` + StartedAt *string `json:"started_at"` + EndedAt *string `json:"ended_at"` + MessageCount int `json:"message_count"` + UserMessageCount int `json:"user_message_count"` + ParentSessionID *string `json:"parent_session_id,omitempty"` + RelationshipType string `json:"relationship_type,omitempty"` + TotalOutputTokens int `json:"total_output_tokens"` + PeakContextTokens int `json:"peak_context_tokens"` + HasTotalOutputTokens bool `json:"has_total_output_tokens"` + HasPeakContextTokens bool `json:"has_peak_context_tokens"` + IsAutomated bool `json:"is_automated"` + + ToolFailureSignalCount int `json:"tool_failure_signal_count"` + ToolRetryCount int `json:"tool_retry_count"` + EditChurnCount int `json:"edit_churn_count"` + ConsecutiveFailureMax int `json:"consecutive_failure_max"` + Outcome string `json:"outcome"` + OutcomeConfidence string `json:"outcome_confidence"` + EndedWithRole string `json:"ended_with_role"` + FinalFailureStreak int `json:"final_failure_streak"` + SignalsPendingSince *string `json:"signals_pending_since,omitempty"` + CompactionCount int `json:"compaction_count"` + MidTaskCompactionCount int `json:"mid_task_compaction_count"` + ContextPressureMax *float64 `json:"context_pressure_max,omitempty"` + HealthScore *int `json:"health_score,omitempty"` + HealthGrade *string `json:"health_grade,omitempty"` + QualitySignals *manifestQualitySignals `json:"quality_signals,omitempty"` + SecretLeakCount int `json:"secret_leak_count"` + + Cwd string `json:"cwd,omitempty"` + GitBranch string `json:"git_branch,omitempty"` + SourceSessionID string `json:"source_session_id,omitempty"` + SourceVersion string `json:"source_version,omitempty"` + TranscriptFidelity string `json:"transcript_fidelity,omitempty"` + ParserMalformedLines int `json:"parser_malformed_lines,omitempty"` + IsTruncated bool `json:"is_truncated,omitempty"` + + DeletedAt *string `json:"deleted_at,omitempty"` + TerminationStatus *string `json:"termination_status,omitempty"` + FilePath *string `json:"file_path,omitempty"` + FileSize *int64 `json:"file_size,omitempty"` + FileMtime *int64 `json:"file_mtime,omitempty"` + FileInode *int64 `json:"file_inode,omitempty"` + FileDevice *int64 `json:"file_device,omitempty"` + FileHash *string `json:"file_hash,omitempty"` + LocalModifiedAt *string `json:"local_modified_at,omitempty"` + TranscriptRevision *string `json:"transcript_revision,omitempty"` + CreatedAt string `json:"created_at"` +} + +// manifestQualitySignals mirrors db.QualitySignals for the same reason +// manifestSession mirrors db.Session: it appears in hashed manifest bytes. +type manifestQualitySignals struct { + Version int `json:"version"` + ShortPromptCount int `json:"short_prompt_count"` + UnstructuredStart bool `json:"unstructured_start"` + MissingSuccessCriteriaCount int `json:"missing_success_criteria_count"` + MissingVerificationCount int `json:"missing_verification_count"` + DuplicatePromptCount int `json:"duplicate_prompt_count"` + NoCodeContextCount int `json:"no_code_context_count"` + RunawayToolLoopCount int `json:"runaway_tool_loop_count"` +} + +func manifestSessionFromDB(s db.Session) manifestSession { + return manifestSession{ + ID: s.ID, + Project: s.Project, + Machine: s.Machine, + Agent: s.Agent, + AgentLabel: s.AgentLabel, + Entrypoint: s.Entrypoint, + FirstMessage: s.FirstMessage, + DisplayName: s.DisplayName, + StartedAt: s.StartedAt, + EndedAt: s.EndedAt, + MessageCount: s.MessageCount, + UserMessageCount: s.UserMessageCount, + ParentSessionID: s.ParentSessionID, + RelationshipType: s.RelationshipType, + TotalOutputTokens: s.TotalOutputTokens, + PeakContextTokens: s.PeakContextTokens, + HasTotalOutputTokens: s.HasTotalOutputTokens, + HasPeakContextTokens: s.HasPeakContextTokens, + IsAutomated: s.IsAutomated, + + ToolFailureSignalCount: s.ToolFailureSignalCount, + ToolRetryCount: s.ToolRetryCount, + EditChurnCount: s.EditChurnCount, + ConsecutiveFailureMax: s.ConsecutiveFailureMax, + Outcome: s.Outcome, + OutcomeConfidence: s.OutcomeConfidence, + EndedWithRole: s.EndedWithRole, + FinalFailureStreak: s.FinalFailureStreak, + SignalsPendingSince: s.SignalsPendingSince, + CompactionCount: s.CompactionCount, + MidTaskCompactionCount: s.MidTaskCompactionCount, + ContextPressureMax: s.ContextPressureMax, + HealthScore: s.HealthScore, + HealthGrade: s.HealthGrade, + QualitySignals: manifestQualitySignalsFromDB(s.QualitySignals), + SecretLeakCount: s.SecretLeakCount, + + Cwd: s.Cwd, + GitBranch: s.GitBranch, + SourceSessionID: s.SourceSessionID, + SourceVersion: s.SourceVersion, + TranscriptFidelity: s.TranscriptFidelity, + ParserMalformedLines: s.ParserMalformedLines, + IsTruncated: s.IsTruncated, + + DeletedAt: s.DeletedAt, + TerminationStatus: s.TerminationStatus, + FilePath: s.FilePath, + FileSize: s.FileSize, + FileMtime: s.FileMtime, + FileInode: s.FileInode, + FileDevice: s.FileDevice, + FileHash: s.FileHash, + LocalModifiedAt: s.LocalModifiedAt, + TranscriptRevision: s.TranscriptRevision, + CreatedAt: s.CreatedAt, + } +} + +func (m manifestSession) dbSession() db.Session { + return db.Session{ + ID: m.ID, + Project: m.Project, + Machine: m.Machine, + Agent: m.Agent, + AgentLabel: m.AgentLabel, + Entrypoint: m.Entrypoint, + FirstMessage: m.FirstMessage, + DisplayName: m.DisplayName, + StartedAt: m.StartedAt, + EndedAt: m.EndedAt, + MessageCount: m.MessageCount, + UserMessageCount: m.UserMessageCount, + ParentSessionID: m.ParentSessionID, + RelationshipType: m.RelationshipType, + TotalOutputTokens: m.TotalOutputTokens, + PeakContextTokens: m.PeakContextTokens, + HasTotalOutputTokens: m.HasTotalOutputTokens, + HasPeakContextTokens: m.HasPeakContextTokens, + IsAutomated: m.IsAutomated, + + ToolFailureSignalCount: m.ToolFailureSignalCount, + ToolRetryCount: m.ToolRetryCount, + EditChurnCount: m.EditChurnCount, + ConsecutiveFailureMax: m.ConsecutiveFailureMax, + Outcome: m.Outcome, + OutcomeConfidence: m.OutcomeConfidence, + EndedWithRole: m.EndedWithRole, + FinalFailureStreak: m.FinalFailureStreak, + SignalsPendingSince: m.SignalsPendingSince, + CompactionCount: m.CompactionCount, + MidTaskCompactionCount: m.MidTaskCompactionCount, + ContextPressureMax: m.ContextPressureMax, + HealthScore: m.HealthScore, + HealthGrade: m.HealthGrade, + QualitySignals: m.QualitySignals.dbQualitySignals(), + SecretLeakCount: m.SecretLeakCount, + + Cwd: m.Cwd, + GitBranch: m.GitBranch, + SourceSessionID: m.SourceSessionID, + SourceVersion: m.SourceVersion, + TranscriptFidelity: m.TranscriptFidelity, + ParserMalformedLines: m.ParserMalformedLines, + IsTruncated: m.IsTruncated, + + DeletedAt: m.DeletedAt, + TerminationStatus: m.TerminationStatus, + FilePath: m.FilePath, + FileSize: m.FileSize, + FileMtime: m.FileMtime, + FileInode: m.FileInode, + FileDevice: m.FileDevice, + FileHash: m.FileHash, + LocalModifiedAt: m.LocalModifiedAt, + TranscriptRevision: m.TranscriptRevision, + CreatedAt: m.CreatedAt, + } +} + +func manifestQualitySignalsFromDB(qs *db.QualitySignals) *manifestQualitySignals { + if qs == nil { + return nil + } + return &manifestQualitySignals{ + Version: qs.Version, + ShortPromptCount: qs.ShortPromptCount, + UnstructuredStart: qs.UnstructuredStart, + MissingSuccessCriteriaCount: qs.MissingSuccessCriteriaCount, + MissingVerificationCount: qs.MissingVerificationCount, + DuplicatePromptCount: qs.DuplicatePromptCount, + NoCodeContextCount: qs.NoCodeContextCount, + RunawayToolLoopCount: qs.RunawayToolLoopCount, + } +} + +func (m *manifestQualitySignals) dbQualitySignals() *db.QualitySignals { + if m == nil { + return nil + } + return &db.QualitySignals{ + Version: m.Version, + ShortPromptCount: m.ShortPromptCount, + UnstructuredStart: m.UnstructuredStart, + MissingSuccessCriteriaCount: m.MissingSuccessCriteriaCount, + MissingVerificationCount: m.MissingVerificationCount, + DuplicatePromptCount: m.DuplicatePromptCount, + NoCodeContextCount: m.NoCodeContextCount, + RunawayToolLoopCount: m.RunawayToolLoopCount, + } +} diff --git a/internal/artifact/manifest_session_test.go b/internal/artifact/manifest_session_test.go new file mode 100644 index 000000000..7bec4a8f4 --- /dev/null +++ b/internal/artifact/manifest_session_test.go @@ -0,0 +1,96 @@ +package artifact + +import ( + "fmt" + "reflect" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/db" +) + +// TestManifestSessionMatchesDBSessionWireFormat pins the manifest wire DTO to +// the JSON-visible fields of db.Session with a fully populated value: every +// exported field gets a distinct non-zero value, so a missing, extra, or +// transposed DTO field changes the canonical JSON and fails the comparison. +// +// If this test fails after adding a field to db.Session, that is the wire +// format asking for a decision: adding the field to manifestSession changes +// every manifest hash fleet-wide (a re-export/re-import of all sessions), so +// extend the DTO only when the field is genuinely part of the session content +// contract; otherwise leave the DTO alone and update populateWireFixture's +// expectations here. +func TestManifestSessionMatchesDBSessionWireFormat(t *testing.T) { + var sess db.Session + populateWireFixture(t, reflect.ValueOf(&sess).Elem(), 1) + + want, err := canonicalJSON(sess) + require.NoError(t, err) + got, err := canonicalJSON(manifestSessionFromDB(sess)) + require.NoError(t, err) + assert.Equal(t, string(want), string(got), + "manifestSession must serialize byte-identically to db.Session") + + roundTrip, err := canonicalJSON(manifestSessionFromDB(sess).dbSession()) + require.NoError(t, err) + assert.Equal(t, string(want), string(roundTrip), + "converting to the wire DTO and back must preserve every wire-visible field") +} + +func TestManifestQualitySignalsMatchesDBWireFormat(t *testing.T) { + var qs db.QualitySignals + populateWireFixture(t, reflect.ValueOf(&qs).Elem(), 100) + + want, err := canonicalJSON(qs) + require.NoError(t, err) + dto := manifestQualitySignalsFromDB(&qs) + require.NotNil(t, dto) + got, err := canonicalJSON(*dto) + require.NoError(t, err) + assert.Equal(t, string(want), string(got)) + + roundTrip, err := canonicalJSON(*dto.dbQualitySignals()) + require.NoError(t, err) + assert.Equal(t, string(want), string(roundTrip)) + + assert.Nil(t, manifestQualitySignalsFromDB(nil)) +} + +// populateWireFixture fills every exported field of a struct with a distinct +// deterministic non-zero value so field transpositions are detectable. +func populateWireFixture(t *testing.T, v reflect.Value, seed int) { + t.Helper() + for i := 0; i < v.NumField(); i++ { + field := v.Field(i) + if !field.CanSet() { + continue + } + setWireFixtureValue(t, field, seed+i) + } +} + +func setWireFixtureValue(t *testing.T, field reflect.Value, n int) { + t.Helper() + switch field.Kind() { + case reflect.String: + field.SetString(fmt.Sprintf("value-%d", n)) + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + field.SetInt(int64(n)) + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + field.SetUint(uint64(n)) + case reflect.Bool: + field.SetBool(true) + case reflect.Float32, reflect.Float64: + field.SetFloat(float64(n) + 0.5) + case reflect.Pointer: + elem := reflect.New(field.Type().Elem()) + setWireFixtureValue(t, elem.Elem(), n) + field.Set(elem) + case reflect.Struct: + populateWireFixture(t, field, n*10) + default: + t.Fatalf("populateWireFixture: unhandled field kind %s; teach the fixture about it", field.Kind()) + } +} diff --git a/internal/artifact/metadata.go b/internal/artifact/metadata.go new file mode 100644 index 000000000..b098a326e --- /dev/null +++ b/internal/artifact/metadata.go @@ -0,0 +1,789 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "sync" + "time" + + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/parser" +) + +const metadataEventExtension = ".json" + +// errOriginNotAdopted reports that this machine has no artifact origin, i.e. +// it never opted into artifact sync via `sync --init`, a sync run, or a peer +// exchange. Recording paths treat it as "stay local"; explicit sync paths +// treat it as a hard error. +var errOriginNotAdopted = errors.New("artifact origin not adopted") + +type metadataSuppressionKey struct{} + +// Metadata operation names written into the metadata event ledger. +const ( + MetadataOpRename = "rename" + MetadataOpSoftDelete = "soft_delete" + MetadataOpRestore = "restore" + MetadataOpStar = "star" + MetadataOpUnstar = "unstar" + MetadataOpPin = "pin" + MetadataOpUnpin = "unpin" + MetadataOpPurge = "purge" +) + +// MetadataPin identifies a pinned message with stable source coordinates. +type MetadataPin struct { + SourceUUID string `json:"source_uuid,omitempty"` + Ordinal int `json:"ordinal"` + Note *string `json:"note,omitempty"` +} + +// MetadataEventInput describes a local user metadata mutation to append. +type MetadataEventInput struct { + SessionID string + Op string + Value json.RawMessage + Pin *MetadataPin +} + +// MetadataRecord describes a metadata artifact written to disk. +type MetadataRecord struct { + HLC string + Origin string + SessionGID string + Op string + Hash string + Path string + Ref Ref +} + +// MetadataPublishedError reports that the artifact file was durably written, +// but local replay bookkeeping failed afterward. +type MetadataPublishedError struct { + Record MetadataRecord + Err error +} + +func (e *MetadataPublishedError) Error() string { + return fmt.Sprintf("metadata event published but local replay state was not recorded: %v", e.Err) +} + +func (e *MetadataPublishedError) Unwrap() error { + return e.Err +} + +// MetadataRecorderOptions configures metadata event artifact writes. +type MetadataRecorderOptions struct { + Origin string + Store ArtifactStore + Now func() time.Time + MaxDrift time.Duration +} + +// MetadataRecorder appends canonical metadata event artifacts. +type MetadataRecorder struct { + mu sync.Mutex + database *db.DB + origin string + store ArtifactStore + clock *HLCClock +} + +// NewMetadataRecorder creates a metadata event recorder for the local artifact store. +func NewMetadataRecorder(database *db.DB, opts MetadataRecorderOptions) *MetadataRecorder { + return &MetadataRecorder{ + database: database, + origin: strings.TrimSpace(opts.Origin), + store: opts.Store, + clock: NewHLCClock(database, HLCClockOptions{ + Now: opts.Now, + MaxDrift: opts.MaxDrift, + }), + } +} + +// WithMetadataEventSuppression marks a context as replaying metadata events. +func WithMetadataEventSuppression(ctx context.Context) context.Context { + return context.WithValue(ctx, metadataSuppressionKey{}, true) +} + +// MetadataEventsSuppressed reports whether local metadata event writes are disabled. +func MetadataEventsSuppressed(ctx context.Context) bool { + suppressed, _ := ctx.Value(metadataSuppressionKey{}).(bool) + return suppressed +} + +// Append writes one metadata event artifact unless ctx is replay-suppressed. +func (r *MetadataRecorder) Append(ctx context.Context, input MetadataEventInput) (MetadataRecord, error) { + if MetadataEventsSuppressed(ctx) { + return MetadataRecord{}, nil + } + if r == nil { + return MetadataRecord{}, nil + } + if r.database == nil { + return MetadataRecord{}, errors.New("metadata recorder database is required") + } + if r.store == nil { + return MetadataRecord{}, errors.New("metadata recorder artifact store is required") + } + if input.SessionID == "" { + return MetadataRecord{}, errors.New("metadata event session id is required") + } + if err := validateMetadataOp(input.Op); err != nil { + return MetadataRecord{}, err + } + origin, err := r.resolveOrigin() + if errors.Is(err, errOriginNotAdopted) { + // The machine never opted into artifact sync: curation stays local. + // If it joins a fleet later, the `sync --init` baseline snapshot + // publishes the accumulated local curation state. + return MetadataRecord{}, nil + } + if err != nil { + return MetadataRecord{}, err + } + stamp, err := r.clock.Next() + if err != nil { + return MetadataRecord{}, err + } + event := metadataEvent{ + Version: formatVersion, + HLC: stamp.String(), + Origin: origin, + SessionGID: MetadataSessionGID(origin, input.SessionID), + Op: input.Op, + Value: input.Value, + Pin: input.Pin, + } + data, err := canonicalJSON(event) + if err != nil { + return MetadataRecord{}, err + } + hash := hashHex(data) + orderKey := stamp.OrderingKey(hash) + projection, err := metadataProjection(metadataArtifact{ + orderKey: orderKey, + hash: hash, + hlc: event.HLC, + event: event, + }, origin) + if err != nil { + return MetadataRecord{}, err + } + ref, err := NewRef(origin, KindMeta, orderKey+metadataEventExtension) + if err != nil { + return MetadataRecord{}, err + } + record := MetadataRecord{ + HLC: event.HLC, + Origin: origin, + SessionGID: event.SessionGID, + Op: event.Op, + Hash: hash, + Ref: ref, + } + identity, err := NewIdentity(hash, int64(len(data))) + if err != nil { + return MetadataRecord{}, err + } + if _, err := r.store.Create(ctx, ref, identity, + canonicalArtifactMediaType(KindMeta), bytes.NewReader(data)); err != nil { + return MetadataRecord{}, fmt.Errorf("creating metadata event: %w", err) + } + if err := r.database.RecordMetadataArtifactProvenance(ctx, db.MetadataArtifactProvenance{ + Origin: origin, OrderKey: orderKey, ArtifactHash: hash, + SessionGID: event.SessionGID, Op: event.Op, + }); err != nil { + return record, &MetadataPublishedError{ + Record: record, + Err: fmt.Errorf("recording metadata artifact provenance: %w", err), + } + } + // Record the local event in the LWW replay register only after the artifact + // exists. Otherwise a failed publish can leave hidden local state that wins + // future LWW comparisons for an event no peer can import. + if _, err := r.database.RecordLocalMetadataProjection(ctx, projection); err != nil { + return record, &MetadataPublishedError{ + Record: record, + Err: fmt.Errorf("recording local metadata replay state: %w", err), + } + } + return record, nil +} + +// RepairLocalSessionMetadata rebuilds local replay bookkeeping for already +// published local metadata artifacts without re-applying their visible +// mutations. +func (r *MetadataRecorder) RepairLocalSessionMetadata( + ctx context.Context, + sessionID string, + ops ...string, +) (int, error) { + if r == nil { + return 0, nil + } + if r.database == nil { + return 0, errors.New("metadata recorder database is required") + } + if r.store == nil { + return 0, errors.New("metadata recorder artifact store is required") + } + if sessionID == "" { + return 0, errors.New("metadata event session id is required") + } + opSet := make(map[string]struct{}, len(ops)) + for _, op := range ops { + if err := validateMetadataOp(op); err != nil { + return 0, err + } + opSet[op] = struct{}{} + } + origin, err := r.resolveOrigin() + if errors.Is(err, errOriginNotAdopted) { + // No origin means no published local artifacts to repair against. + return 0, nil + } + if err != nil { + return 0, err + } + sessionGID := MetadataSessionGID(origin, sessionID) + provenance, queryErr := r.database.MetadataArtifactProvenanceForSession( + ctx, origin, sessionGID, ops..., + ) + if queryErr != nil { + return 0, queryErr + } + events, err := readMetadataArtifactsFromProvenance( + ctx, r.database, r.store, provenance, + ) + if err != nil { + return 0, err + } + repaired := 0 + for _, art := range events { + if err := ctx.Err(); err != nil { + return repaired, err + } + if art.event.SessionGID != sessionGID { + continue + } + if len(opSet) > 0 { + if _, ok := opSet[art.event.Op]; !ok { + continue + } + } + if err := validateMetadataArtifactEvent(art, origin); err != nil { + if errors.Is(err, errFutureArtifactVersion) { + continue + } + return repaired, err + } + if err := validateMetadataOp(art.event.Op); err != nil { + return repaired, err + } + projection, err := metadataProjection(art, origin) + if err != nil { + return repaired, err + } + if _, err := r.database.RecordLocalMetadataProjection(ctx, projection); err != nil { + return repaired, fmt.Errorf("repairing local metadata replay state: %w", err) + } + repaired++ + } + return repaired, nil +} + +func readMetadataArtifactsFromProvenance( + ctx context.Context, + database *db.DB, + store ArtifactStore, + provenance []db.MetadataArtifactProvenance, +) (events []metadataArtifact, retErr error) { + for _, indexed := range provenance { + ref, err := NewRef(indexed.Origin, KindMeta, indexed.OrderKey+metadataEventExtension) + if err != nil { + return nil, err + } + entry, err := store.Stat(ctx, ref) + if errors.Is(err, ErrArtifactNotFound) { + // Provenance survives deliberate removal or vault recovery. A + // missing immutable event cannot repair replay state; callers may + // publish a replacement event with a later ordering key. + continue + } + if err != nil { + return nil, err + } + if entry.Identity.SHA256 != indexed.ArtifactHash { + return nil, fmt.Errorf("%w: metadata provenance identity changed", ErrArtifactCorrupt) + } + data, err := readVerifiedStoreArtifact( + ctx, database, store, entry, manifestDecodedLimit, + ) + if err != nil { + return nil, err + } + var event metadataEvent + if err := json.Unmarshal(data, &event); err != nil { + return nil, fmt.Errorf("decoding metadata artifact %s: %w", ref.Name, err) + } + idx := strings.LastIndex(indexed.OrderKey, "-") + if idx < 0 { + return nil, fmt.Errorf("metadata artifact %s missing hash suffix", ref.Name) + } + events = append(events, metadataArtifact{ + path: ref.Name, orderKey: indexed.OrderKey, hash: indexed.ArtifactHash, + hlc: indexed.OrderKey[:idx], event: event, + }) + } + return events, nil +} + +// AppendBaseline writes metadata events for existing local curation that +// predates artifact metadata recording. +func (r *MetadataRecorder) AppendBaseline(ctx context.Context) (int, error) { + if r == nil { + return 0, nil + } + if r.database == nil { + return 0, errors.New("metadata recorder database is required") + } + snap, err := r.database.MetadataBaselineSnapshot(ctx) + if err != nil { + return 0, err + } + return r.AppendBaselineSnapshot(ctx, snap) +} + +// AppendBaselineSnapshot writes metadata events from a previously captured +// curation snapshot. Callers that import peer artifacts before initialization +// should capture the snapshot before that import so newly imported rows cannot +// be re-published as local baseline metadata. +func (r *MetadataRecorder) AppendBaselineSnapshot( + ctx context.Context, + snap db.MetadataBaselineSnapshot, +) (int, error) { + if r == nil { + return 0, nil + } + if r.database == nil { + return 0, errors.New("metadata recorder database is required") + } + origin, err := r.resolveOrigin() + if err != nil { + return 0, err + } + written := 0 + for _, rename := range snap.Renames { + covered, err := r.baselineFieldCovered(ctx, origin, rename.SessionID, "display_name") + if err != nil { + return written, err + } + if covered { + continue + } + value, err := metadataRenameValue(rename.DisplayName) + if err != nil { + return written, err + } + if _, err := r.Append(ctx, MetadataEventInput{ + SessionID: rename.SessionID, + Op: MetadataOpRename, + Value: value, + }); err != nil { + return written, fmt.Errorf("writing baseline rename metadata: %w", err) + } + written++ + } + for _, sessionID := range snap.StarredSessionIDs { + covered, err := r.baselineFieldCovered(ctx, origin, sessionID, "starred") + if err != nil { + return written, err + } + if covered { + continue + } + if _, err := r.Append(ctx, MetadataEventInput{ + SessionID: sessionID, + Op: MetadataOpStar, + }); err != nil { + return written, fmt.Errorf("writing baseline star metadata: %w", err) + } + written++ + } + for _, sessionID := range snap.SoftDeletedIDs { + covered, err := r.baselineFieldCovered(ctx, origin, sessionID, "deleted_at") + if err != nil { + return written, err + } + if covered { + continue + } + if _, err := r.Append(ctx, MetadataEventInput{ + SessionID: sessionID, + Op: MetadataOpSoftDelete, + }); err != nil { + return written, fmt.Errorf("writing baseline soft-delete metadata: %w", err) + } + written++ + } + for _, pin := range snap.Pins { + metadataPin := MetadataPin{ + SourceUUID: pin.SourceUUID, + Ordinal: pin.Ordinal, + Note: pin.Note, + } + covered, err := r.baselineFieldCovered( + ctx, origin, pin.SessionID, "pin:"+metadataPinAnchor(metadataPin), + ) + if err != nil { + return written, err + } + if covered { + continue + } + if _, err := r.Append(ctx, MetadataEventInput{ + SessionID: pin.SessionID, + Op: MetadataOpPin, + Pin: &metadataPin, + }); err != nil { + return written, fmt.Errorf("writing baseline pin metadata: %w", err) + } + written++ + } + return written, nil +} + +// materializeCurrentState rebuilds the local origin's canonical metadata +// publication into an empty replacement store. It republishes only replay +// winners authored by this origin, including negative/default operations, then +// baselines positive curation fields that have no replay winner. Foreign +// winners remain authoritative in SQLite and are never reattributed locally. +func (r *MetadataRecorder) materializeCurrentState(ctx context.Context) (int, error) { + if r == nil { + return 0, nil + } + if r.database == nil { + return 0, errors.New("metadata recorder database is required") + } + origin, err := r.resolveOrigin() + if err != nil { + return 0, err + } + written := 0 + err = r.database.VisitMetadataReplayWinnersAuthoredBy(ctx, origin, func( + winner db.MetadataProjection, + ) error { + input := MetadataEventInput{ + SessionID: winner.SessionGID, + Op: winner.Op, + } + switch winner.Op { + case MetadataOpRename: + input.Value = json.RawMessage(winner.Value) + case MetadataOpPin, MetadataOpUnpin: + if winner.Pin == nil { + return fmt.Errorf("materializing %s metadata without pin payload", winner.Op) + } + input.Pin = &MetadataPin{ + SourceUUID: winner.Pin.SourceUUID, + Ordinal: winner.Pin.Ordinal, + Note: winner.Pin.Note, + } + } + if _, err := r.Append(ctx, input); err != nil { + return fmt.Errorf("materializing current %s metadata: %w", winner.Op, err) + } + written++ + return nil + }) + if err != nil { + return written, err + } + baselined, err := r.AppendBaseline(ctx) + return written + baselined, err +} + +// materializeCurrentStateAtHLC reconstructs reset metadata idempotently. Replay +// winners retain their original immutable identity; positive baseline fields +// without a replay winner use one pre-persisted HLC so a crash after Create but +// before SQLite bookkeeping retries the same bytes and reference. +func (r *MetadataRecorder) materializeCurrentStateAtHLC( + ctx context.Context, baselineHLC string, +) (int, error) { + if r == nil { + return 0, nil + } + if r.database == nil { + return 0, errors.New("metadata recorder database is required") + } + if r.store == nil { + return 0, errors.New("metadata reset materialization requires an artifact store") + } + if _, err := ParseHLCTimestamp(baselineHLC); err != nil { + return 0, fmt.Errorf("parsing reset baseline HLC: %w", err) + } + origin, err := r.resolveOrigin() + if err != nil { + return 0, err + } + written := 0 + err = r.database.VisitMetadataReplayWinnersAuthoredBy(ctx, origin, func( + winner db.MetadataProjection, + ) error { + event, err := metadataEventFromProjection(winner) + if err != nil { + return err + } + if err := r.createExactMetadataEvent( + ctx, event, winner.OrderKey, winner.ArtifactHash, false, + ); err != nil { + return fmt.Errorf("materializing current %s metadata: %w", winner.Op, err) + } + written++ + return nil + }) + if err != nil { + return written, err + } + baselined := 0 + err = r.database.VisitMetadataBaselinePages(ctx, func(page db.MetadataBaselineSnapshot) error { + pageWritten, err := r.materializeBaselineSnapshotAtHLC( + ctx, origin, baselineHLC, page, + ) + baselined += pageWritten + return err + }) + return written + baselined, err +} + +func metadataEventFromProjection(winner db.MetadataProjection) (metadataEvent, error) { + event := metadataEvent{ + Version: formatVersion, + HLC: winner.HLC, + Origin: winner.EventOrigin, + SessionGID: winner.SessionGID, + Op: winner.Op, + } + switch winner.Op { + case MetadataOpRename: + event.Value = json.RawMessage(winner.Value) + case MetadataOpPin, MetadataOpUnpin: + if winner.Pin == nil { + return metadataEvent{}, fmt.Errorf("materializing %s metadata without pin payload", winner.Op) + } + event.Pin = &MetadataPin{ + SourceUUID: winner.Pin.SourceUUID, + Ordinal: winner.Pin.Ordinal, + Note: winner.Pin.Note, + } + } + return event, nil +} + +func (r *MetadataRecorder) createExactMetadataEvent( + ctx context.Context, + event metadataEvent, + expectedOrderKey string, + expectedHash string, + recordProjection bool, +) error { + if err := validateMetadataOp(event.Op); err != nil { + return err + } + stamp, err := ParseHLCTimestamp(event.HLC) + if err != nil { + return err + } + data, err := canonicalJSON(event) + if err != nil { + return err + } + hash := hashHex(data) + orderKey := stamp.OrderingKey(hash) + if expectedHash != "" && hash != expectedHash { + return fmt.Errorf("metadata projection hash mismatch: reconstructed %s, expected %s", + hash, expectedHash) + } + if expectedOrderKey != "" && orderKey != expectedOrderKey { + return fmt.Errorf("metadata projection order key mismatch: reconstructed %s, expected %s", + orderKey, expectedOrderKey) + } + ref, err := NewRef(event.Origin, KindMeta, orderKey+metadataEventExtension) + if err != nil { + return err + } + identity, err := NewIdentity(hash, int64(len(data))) + if err != nil { + return err + } + if _, err := r.store.Create(ctx, ref, identity, + canonicalArtifactMediaType(KindMeta), bytes.NewReader(data)); err != nil { + return fmt.Errorf("creating exact metadata event: %w", err) + } + if !recordProjection { + return nil + } + projection, err := metadataProjection(metadataArtifact{ + orderKey: orderKey, + hash: hash, + hlc: event.HLC, + event: event, + }, event.Origin) + if err != nil { + return err + } + if err := r.database.RecordMetadataArtifactProvenance(ctx, db.MetadataArtifactProvenance{ + Origin: event.Origin, OrderKey: orderKey, ArtifactHash: hash, + SessionGID: event.SessionGID, Op: event.Op, + }); err != nil { + return fmt.Errorf("recording exact metadata artifact provenance: %w", err) + } + if _, err := r.database.RecordLocalMetadataProjection(ctx, projection); err != nil { + return fmt.Errorf("recording exact local metadata replay state: %w", err) + } + return nil +} + +func (r *MetadataRecorder) materializeBaselineSnapshotAtHLC( + ctx context.Context, + origin string, + hlc string, + snapshot db.MetadataBaselineSnapshot, +) (int, error) { + written := 0 + create := func(sessionID, field string, event metadataEvent) error { + covered, err := r.baselineFieldCovered(ctx, origin, sessionID, field) + if err != nil { + return err + } + if covered { + return nil + } + event.Version = formatVersion + event.HLC = hlc + event.Origin = origin + event.SessionGID = MetadataSessionGID(origin, sessionID) + if err := r.createExactMetadataEvent(ctx, event, "", "", true); err != nil { + return err + } + written++ + return nil + } + for _, rename := range snapshot.Renames { + value, err := metadataRenameValue(rename.DisplayName) + if err != nil { + return written, err + } + if err := create(rename.SessionID, "display_name", metadataEvent{ + Op: MetadataOpRename, Value: value, + }); err != nil { + return written, fmt.Errorf("materializing baseline rename metadata: %w", err) + } + } + for _, sessionID := range snapshot.StarredSessionIDs { + if err := create(sessionID, "starred", metadataEvent{Op: MetadataOpStar}); err != nil { + return written, fmt.Errorf("materializing baseline star metadata: %w", err) + } + } + for _, sessionID := range snapshot.SoftDeletedIDs { + if err := create(sessionID, "deleted_at", metadataEvent{Op: MetadataOpSoftDelete}); err != nil { + return written, fmt.Errorf("materializing baseline soft-delete metadata: %w", err) + } + } + for _, pin := range snapshot.Pins { + metadataPin := &MetadataPin{ + SourceUUID: pin.SourceUUID, + Ordinal: pin.Ordinal, + Note: pin.Note, + } + if err := create( + pin.SessionID, + "pin:"+metadataPinAnchor(*metadataPin), + metadataEvent{Op: MetadataOpPin, Pin: metadataPin}, + ); err != nil { + return written, fmt.Errorf("materializing baseline pin metadata: %w", err) + } + } + return written, nil +} + +func (r *MetadataRecorder) baselineFieldCovered( + ctx context.Context, + origin string, + sessionID string, + field string, +) (bool, error) { + _, ok, err := r.database.MetadataReplayStateOp( + ctx, MetadataSessionGID(origin, sessionID), field, + ) + if err != nil { + return false, fmt.Errorf("checking baseline metadata field %s: %w", field, err) + } + return ok, nil +} + +// MetadataSessionGID returns the global metadata target ID for a session. +func MetadataSessionGID(origin, sessionID string) string { + if host, _ := parser.StripHostPrefix(sessionID); host != "" { + return sessionID + } + return origin + "~" + sessionID +} + +// resolveOrigin returns the recorder's origin without ever creating one: the +// explicit option wins, then the origin persisted in DB sync state. A machine +// with no origin anywhere has not opted into artifact sync and gets +// errOriginNotAdopted. The empty result is not cached, so a recorder built +// before opt-in starts resolving the origin as soon as it is adopted. +func (r *MetadataRecorder) resolveOrigin() (string, error) { + r.mu.Lock() + defer r.mu.Unlock() + if r.origin != "" { + if err := validateOriginID(r.origin); err != nil { + return "", fmt.Errorf("metadata recorder origin: %w", err) + } + return r.origin, nil + } + origin, err := StoredOrigin(r.database) + if err != nil { + return "", err + } + if origin == "" { + return "", errOriginNotAdopted + } + r.origin = origin + return origin, nil +} + +func metadataRenameValue(displayName *string) (json.RawMessage, error) { + data, err := json.Marshal(struct { + DisplayName *string `json:"display_name"` + }{DisplayName: displayName}) + if err != nil { + return nil, err + } + return json.RawMessage(data), nil +} + +func validateMetadataOp(op string) error { + switch op { + case MetadataOpRename, + MetadataOpSoftDelete, + MetadataOpRestore, + MetadataOpStar, + MetadataOpUnstar, + MetadataOpPin, + MetadataOpUnpin, + MetadataOpPurge: + return nil + default: + return fmt.Errorf("unsupported metadata event op %q", op) + } +} diff --git a/internal/artifact/metadata_test.go b/internal/artifact/metadata_test.go new file mode 100644 index 000000000..dc0750f82 --- /dev/null +++ b/internal/artifact/metadata_test.go @@ -0,0 +1,266 @@ +package artifact + +import ( + "context" + "encoding/json" + "fmt" + "strconv" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type metadataPointCountingStore struct { + ArtifactStore + listCalls atomic.Int32 + statCalls atomic.Int32 + openCalls atomic.Int32 +} + +func (s *metadataPointCountingStore) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + s.listCalls.Add(1) + return s.ArtifactStore.Entries(ctx, origin, kind) +} + +func (s *metadataPointCountingStore) Stat(ctx context.Context, ref Ref) (Entry, error) { + s.statCalls.Add(1) + return s.ArtifactStore.Stat(ctx, ref) +} + +func (s *metadataPointCountingStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + s.openCalls.Add(1) + return s.ArtifactStore.Open(ctx, ref) +} + +func TestMetadataRecorderAppendWritesCanonicalEvent(t *testing.T) { + database := testDB(t) + store := newTestArtifactStore(t) + now := fixedHLCTime() + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: "laptop-a1b2c3", + Store: store, + Now: func() time.Time { return now }, + }) + + value := json.RawMessage(`{"display_name":"Renamed session"}`) + record, err := recorder.Append(context.Background(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpRename, + Value: value, + }) + require.NoError(t, err) + + assert.Equal(t, "2026-06-14T010203.000000001Z-00000000000000000000", record.HLC) + assert.Equal(t, "laptop-a1b2c3", record.Origin) + assert.Equal(t, "laptop-a1b2c3~sess-1", record.SessionGID) + assert.Equal(t, MetadataOpRename, record.Op) + + data := readContractArtifact(t, store, record.Ref) + assert.Equal(t, + "{\"hlc\":\"2026-06-14T010203.000000001Z-00000000000000000000\",\"op\":\"rename\",\"origin\":\"laptop-a1b2c3\",\"session_gid\":\"laptop-a1b2c3~sess-1\",\"v\":1,\"value\":{\"display_name\":\"Renamed session\"}}\n", + string(data), + ) + assert.Equal(t, hashHex(data), record.Hash) + assert.Equal(t, record.HLC+"-"+record.Hash+".json", record.Ref.Name) + // The metadata event filename must be safe on every supported OS, + // including Windows, which forbids these characters in path components. + assert.NotContains(t, record.Ref.Name, ":", + "metadata filename must not contain ':' (invalid on Windows)") + for _, c := range `<>:"/\|?*` { + assert.NotContainsf(t, record.HLC, string(c), + "HLC %q must not contain %q (invalid in Windows filenames)", record.HLC, string(c)) + } +} + +func TestMetadataRecorderAppendUsesCanonicalStoreIdentity(t *testing.T) { + database := testDB(t) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + store := &recordingArtifactStore{ArtifactStore: filesystem} + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Store: store, + Origin: "laptop-a1b2c3", + Now: fixedHLCTime, + }) + + record, err := recorder.Append(t.Context(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + require.Len(t, store.creates, 1) + created := store.creates[0] + assert.Equal(t, Kind(KindMeta), created.Ref.Kind) + assert.Equal(t, record.HLC+"-"+record.Hash+".json", created.Ref.Name) + assert.Equal(t, record.Hash, created.Identity.SHA256) + assert.Equal(t, int64(len(readContractArtifact(t, store, created.Ref))), created.Identity.Size) + assert.Equal(t, created.Ref, record.Ref) +} + +func TestMetadataRecorderStoreModeDoesNotExposeLegacyPath(t *testing.T) { + database := testDB(t) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Store: filesystem, + Origin: "laptop-a1b2c3", + Now: fixedHLCTime, + }) + + record, err := recorder.Append(t.Context(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + assert.Empty(t, record.Path, "ArtifactStore takes precedence over the legacy filesystem path") + assert.Equal(t, Kind(KindMeta), record.Ref.Kind) +} + +func TestMetadataRepairStoreWorkIsBoundedByTargetSessionProvenance(t *testing.T) { + for _, unrelated := range []int{20, 1_000} { + t.Run(strconv.Itoa(unrelated), func(t *testing.T) { + database := testDB(t) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + store := &metadataPointCountingStore{ArtifactStore: filesystem} + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Store: store, Origin: "desk-a1b2c3", Now: fixedHLCTime, + }) + for index := range unrelated { + _, err := recorder.Append(t.Context(), MetadataEventInput{ + SessionID: fmt.Sprintf("unrelated-%05d", index), Op: MetadataOpStar, + }) + require.NoError(t, err) + } + _, err = recorder.Append(t.Context(), MetadataEventInput{ + SessionID: "target", Op: MetadataOpStar, + }) + require.NoError(t, err) + _, err = recorder.Append(t.Context(), MetadataEventInput{ + SessionID: "target", Op: MetadataOpSoftDelete, + }) + require.NoError(t, err) + + store.listCalls.Store(0) + store.statCalls.Store(0) + store.openCalls.Store(0) + repaired, err := recorder.RepairLocalSessionMetadata( + t.Context(), "target", MetadataOpStar, + ) + require.NoError(t, err) + assert.Equal(t, 1, repaired) + assert.Zero(t, store.listCalls.Load(), + "targeted repair must never scan the metadata ledger") + assert.Equal(t, int32(1), store.statCalls.Load()) + assert.Equal(t, int32(1), store.openCalls.Load()) + }) + } +} + +func TestImportObservesRemoteHLCForLaterLocalEdits(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + localOrigin := "desktop-d4e5f6" + peerOrigin := "laptop-a1b2c3" + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + peerGID := localOrigin + "~sess-1" + + recorderNow := fixedHLCTime() + // A peer event whose wall time is ahead of the recorder's local clock but + // within the drift bound. + remoteStamp := HLCTimestamp{WallTime: recorderNow.Add(2 * time.Minute), Logical: 5} + ref := createMetadataArtifactInStore( + t, store, replayRenameEvent(t, peerOrigin, peerGID, remoteStamp.String(), "Peer name"), + ) + + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: localOrigin, + Store: store, + Now: func() time.Time { return recorderNow }, + }) + + entry, err := store.Stat(ctx, ref) + require.NoError(t, err) + coordinator := NewStoreImportCoordinator(database, store, localOrigin) + coordinator.now = func() time.Time { return recorderNow } + require.NoError(t, coordinator.RecordChanged(ctx, entry)) + imported, err := coordinator.Finalize(ctx) + require.NoError(t, err) + assert.Equal(t, 1, imported.Metadata) + + // The next local edit must receive an HLC strictly after the observed + // remote HLC even though the local wall clock is behind it. + rec, err := recorder.Append(ctx, MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + localStamp, err := ParseHLCTimestamp(rec.HLC) + require.NoError(t, err) + assert.Equal(t, 1, localStamp.Compare(remoteStamp), + "local HLC %s must be after observed remote HLC %s", rec.HLC, remoteStamp.String()) +} + +func TestMetadataRecorderWithoutOriginIsNoOp(t *testing.T) { + database := testDB(t) + store := newTestArtifactStore(t) + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{Store: store}) + + record, err := recorder.Append(context.Background(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + assert.Equal(t, MetadataRecord{}, record) + + origin, err := StoredOrigin(database) + require.NoError(t, err) + assert.Empty(t, origin, + "recorder must not mint an origin for a machine that never opted into artifact sync") + assertNoPublishedArtifacts(t, store, "desk-a1b2c3") + + repaired, err := recorder.RepairLocalSessionMetadata(context.Background(), "sess-1") + require.NoError(t, err) + assert.Zero(t, repaired) +} + +func TestMetadataRecorderRecordsAfterOriginAdopted(t *testing.T) { + database := testDB(t) + store := newTestArtifactStore(t) + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{Store: store}) + + record, err := recorder.Append(context.Background(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + assert.Empty(t, record.Origin) + + require.NoError(t, AdoptOrigin(database, "desk-a1b2c3")) + + record, err = recorder.Append(context.Background(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + assert.Equal(t, "desk-a1b2c3", record.Origin) + _, err = store.Stat(t.Context(), record.Ref) + require.NoError(t, err, + "recorder must start writing events once the origin is adopted, without reconstruction") +} + +func TestMetadataSessionGID(t *testing.T) { + assert.Equal(t, "desk-a1b2c3~sess-1", MetadataSessionGID("desk-a1b2c3", "sess-1")) + assert.Equal(t, "laptop-d4e5f6~sess-1", MetadataSessionGID("desk-a1b2c3", "laptop-d4e5f6~sess-1")) +} diff --git a/internal/artifact/nested_limits_test.go b/internal/artifact/nested_limits_test.go new file mode 100644 index 000000000..b2466a82d --- /dev/null +++ b/internal/artifact/nested_limits_test.go @@ -0,0 +1,430 @@ +package artifact + +import ( + "bytes" + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +func exportSessionToTestStoreWithLimits( + t *testing.T, + ctx context.Context, + database *db.DB, + store ArtifactStore, + origin string, + limits artifactLimits, +) (string, bool, error) { + t.Helper() + sess, err := database.GetSessionFull(ctx, "sess-1") + require.NoError(t, err) + require.NotNil(t, sess) + messages, err := database.GetAllMessages(ctx, "sess-1") + require.NoError(t, err) + usage, err := database.GetUsageEvents(ctx, "sess-1") + require.NoError(t, err) + return exportLoadedSessionToStore(ctx, store, origin, sess, messages, usage, limits) +} + +func TestDecodeSegmentRejectsAggregateNestedLimitsWithSmallLimits(t *testing.T) { + tests := []struct { + name string + records []segmentMessage + configure func(*artifactLimits) + wantError string + }{ + { + name: "tool calls per segment", + records: []segmentMessage{ + {ToolCalls: []segmentToolCall{{}}}, + {ToolCalls: []segmentToolCall{{}}}, + }, + configure: func(limits *artifactLimits) { + limits.segmentToolCalls = 1 + }, + wantError: "segment tool call limit", + }, + { + name: "result events per segment", + records: []segmentMessage{ + {ToolCalls: []segmentToolCall{{ResultEvents: []segmentResultEvent{{}}}}}, + {ToolCalls: []segmentToolCall{{ResultEvents: []segmentResultEvent{{}}}}}, + }, + configure: func(limits *artifactLimits) { + limits.segmentResultEvents = 1 + }, + wantError: "segment result event limit", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + limits := productionArtifactLimits() + tt.configure(&limits) + data := nestedSegmentData(t, tt.records...) + + _, err := decodeSegmentWithLimits(data, limits) + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantError) + }) + } +} + +func TestPeerArtifactDefersFutureNestedSchema(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + data := []byte(`{"v":2,"ordinal":{"future_shape":true},"tool_calls":{"future_shape":true}}` + "\n") + hash := hashHex(data) + compressed := compressPeerTestData(t, data) + + ref, err := createCompressedTestArtifact(t, store, origin, KindSegments, hash, compressed) + require.NoError(t, err) + assert.Equal(t, hash+".ndjson", ref.Name) + var stored bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, + bytes.NewReader(readContractArtifact(t, store, ref)), &stored)) + assert.Equal(t, compressed, stored.Bytes()) +} + +func TestDecodeSegmentAcceptsCanonicalTrailingNewlineAndEmptySession(t *testing.T) { + record := nestedSegmentData(t, segmentMessage{}) + tests := []struct { + name string + data []byte + want int + }{ + {name: "canonical trailing newline", data: record, want: 1}, + {name: "zero byte empty segment", data: nil, want: 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + msgs, err := decodeSegment(tt.data) + require.NoError(t, err) + assert.Len(t, msgs, tt.want) + }) + } +} + +func TestImportQuarantinesNestedAmplificationWithoutAdvancingState(t *testing.T) { + tests := []struct { + name string + record segmentMessage + }{ + { + name: "too many tool calls", + record: segmentMessage{ + ToolCalls: make([]segmentToolCall, 257), + }, + }, + { + name: "too many result events", + record: segmentMessage{ + ToolCalls: []segmentToolCall{{ + ResultEvents: make([]segmentResultEvent, 1_025), + }}, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + localOrigin := "desktop-d4e5f6" + importDB := testDB(t) + gid, segmentRef := writeNestedImportFixture( + t, store, origin, tt.record, + ) + + res, err := importResultFromTestStore(ctx, importDB, store, localOrigin) + require.NoError(t, err) + assert.False(t, res.Changed()) + state, err := importDB.GetSyncState(importStateKey(origin, gid)) + require.NoError(t, err) + assert.Empty(t, state) + got, err := importDB.GetSessionFull(ctx, gid) + require.NoError(t, err) + assert.Nil(t, got) + _, err = store.Stat(ctx, segmentRef) + assert.ErrorIs(t, err, ErrArtifactNotFound) + }) + } +} + +func TestExportRejectsNestedAmplificationBeforePublication(t *testing.T) { + tests := []struct { + name string + message db.Message + wantError string + }{ + { + name: "too many tool calls in one message", + message: db.Message{ + ToolCalls: make([]db.ToolCall, 257), + }, + wantError: "tool call limit exceeded for message ordinal 0", + }, + { + name: "too many result events in one tool call", + message: db.Message{ + ToolCalls: []db.ToolCall{{ + ResultEvents: make([]db.ToolResultEvent, 1_025), + }}, + }, + wantError: "result event limit exceeded for tool call 0 in message ordinal 0", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + message := tt.message + message.SessionID = "sess-1" + message.Ordinal = 0 + message.Role = "assistant" + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{message})) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantError) + assertNoPublishedArtifacts(t, store, origin) + }) + } +} + +func TestExportChunksOnAggregateNestedLimitsWithSmallLimits(t *testing.T) { + tests := []struct { + name string + resultEvents []db.ToolResultEvent + configure func(*artifactLimits) + }{ + { + name: "tool calls per segment", + configure: func(limits *artifactLimits) { + limits.segmentToolCalls = 2 + }, + }, + { + name: "result events per segment", + resultEvents: []db.ToolResultEvent{{EventIndex: 0}}, + configure: func(limits *artifactLimits) { + limits.segmentResultEvents = 2 + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + msgs := make([]db.Message, 3) + for ordinal := range msgs { + msgs[ordinal] = db.Message{ + SessionID: "sess-1", + Ordinal: ordinal, + Role: "assistant", + ToolCalls: []db.ToolCall{{ + ResultEvents: tt.resultEvents, + }}, + } + } + require.NoError(t, database.ReplaceSessionMessages("sess-1", msgs)) + limits := productionArtifactLimits() + tt.configure(&limits) + + manifestHash, changed, err := exportSessionToTestStoreWithLimits( + t, ctx, database, store, origin, limits, + ) + require.NoError(t, err) + assert.True(t, changed) + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + m, err := decodeManifestWithLimits( + readContractArtifact(t, store, manifestRef), productionArtifactLimits(), + ) + require.NoError(t, err) + require.Len(t, m.Segments, 2) + got := testStoreManifestMessages(t, store, origin, m) + require.Len(t, got, 3) + for ordinal := range got { + assert.Equal(t, ordinal, got[ordinal].Ordinal) + require.Len(t, got[ordinal].ToolCalls, 1) + assert.Len(t, got[ordinal].ToolCalls[0].ResultEvents, len(tt.resultEvents)) + } + }) + } +} + +func TestExportRejectsMessageThatCannotFitNestedSegmentLimits(t *testing.T) { + tests := []struct { + name string + message db.Message + configure func(*artifactLimits) + wantError string + }{ + { + name: "tool calls cannot split across segments", + message: db.Message{ + ToolCalls: []db.ToolCall{{}, {}}, + }, + configure: func(limits *artifactLimits) { + limits.segmentToolCalls = 1 + }, + wantError: "2 tool calls", + }, + { + name: "one tool result history cannot split across segments", + message: db.Message{ + ToolCalls: []db.ToolCall{{ + ResultEvents: []db.ToolResultEvent{{EventIndex: 0}, {EventIndex: 1}}, + }}, + }, + configure: func(limits *artifactLimits) { + limits.segmentResultEvents = 1 + }, + wantError: "2 result events", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + message := tt.message + message.SessionID = "sess-1" + message.Ordinal = 0 + message.Role = "assistant" + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{message})) + limits := productionArtifactLimits() + tt.configure(&limits) + + _, _, err := exportSessionToTestStoreWithLimits( + t, ctx, database, store, origin, limits, + ) + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot fit in one segment") + assert.Contains(t, err.Error(), tt.wantError) + assertNoPublishedArtifacts(t, store, origin) + }) + } +} + +func TestExportRejectsSessionNestedLimitsBeforeWritingWithSmallLimits(t *testing.T) { + tests := []struct { + name string + configure func(*artifactLimits) + wantError string + }{ + { + name: "tool calls per session", + configure: func(limits *artifactLimits) { + limits.sessionToolCalls = 1 + }, + wantError: "session tool call limit", + }, + { + name: "result events per session", + configure: func(limits *artifactLimits) { + limits.sessionResultEvents = 1 + }, + wantError: "session result event limit", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + msgs := make([]db.Message, 2) + for ordinal := range msgs { + msgs[ordinal] = db.Message{ + SessionID: "sess-1", + Ordinal: ordinal, + Role: "assistant", + ToolCalls: []db.ToolCall{{ + ResultEvents: []db.ToolResultEvent{{EventIndex: 0}}, + }}, + } + } + require.NoError(t, database.ReplaceSessionMessages("sess-1", msgs)) + limits := productionArtifactLimits() + tt.configure(&limits) + + _, _, err := exportSessionToTestStoreWithLimits( + t, ctx, database, store, origin, limits, + ) + require.Error(t, err) + assert.Contains(t, err.Error(), tt.wantError) + assertNoPublishedArtifacts(t, store, origin) + }) + } +} + +func nestedSegmentData(t *testing.T, records ...segmentMessage) []byte { + t.Helper() + var data bytes.Buffer + for ordinal := range records { + records[ordinal].Version = formatVersion + records[ordinal].Ordinal = ordinal + records[ordinal].Role = "assistant" + encoded, err := canonicalJSON(records[ordinal]) + require.NoError(t, err) + _, err = data.Write(encoded) + require.NoError(t, err) + } + return data.Bytes() +} + +func writeNestedImportFixture( + t *testing.T, + store ArtifactStore, + origin string, + record segmentMessage, +) (string, Ref) { + t.Helper() + data := nestedSegmentData(t, record) + segmentHash := hashHex(data) + segmentRef, err := NewRef(origin, KindSegments, segmentHash+".ndjson") + require.NoError(t, err) + createContractArtifact(t, store, segmentRef, data) + + gid := origin + "~sess-1" + m := manifest{ + Version: formatVersion, + Origin: origin, + NativeSessionID: "sess-1", + Session: manifestSession{ + ID: "sess-1", + Machine: origin, + }, + Segments: []string{segmentHash}, + } + manifestData, err := canonicalJSON(m) + require.NoError(t, err) + manifestHash := hashHex(manifestData) + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + createContractArtifact(t, store, manifestRef, manifestData) + cpData, err := canonicalJSON(checkpoint{ + Version: formatVersion, + Origin: origin, + Sequence: 1, + Sessions: map[string]string{gid: manifestHash}, + }) + require.NoError(t, err) + cpRef, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + createContractArtifact(t, store, cpRef, cpData) + return gid, segmentRef +} diff --git a/internal/artifact/peer.go b/internal/artifact/peer.go new file mode 100644 index 000000000..79a1c3e80 --- /dev/null +++ b/internal/artifact/peer.go @@ -0,0 +1,519 @@ +package artifact + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "strings" + "time" + + "go.kenn.io/agentsview/internal/db" +) + +const ( + KindCheckpoints = "checkpoints" + KindManifests = "manifests" + KindSegments = "segments" + KindMeta = "meta" + KindRaw = "raw" +) + +var ( + ErrArtifactInvalid = errors.New("invalid artifact") + ErrArtifactNotFound = errors.New("artifact not found") + ErrArtifactConflict = errors.New("artifact conflict") +) + +// PeerArtifact is one immutable artifact file served through the peer API. +type PeerArtifact struct { + Origin string + Kind string + Name string + Hash string + ContentType string + Data []byte +} + +// PeerArtifactWrite describes the result of a peer artifact write. +type PeerArtifactWrite struct { + Origin string + Kind string + Name string + Hash string + Size int64 + Duplicate bool +} + +// PeerArtifactSpool is a fully verified, seekable wire response. Close removes +// its private backing file. +type PeerArtifactSpool struct { + Origin string + Kind string + Name string + Hash string + ContentType string + Size int64 + File *os.File +} + +// OriginArtifactIndex groups an origin's artifact names by protocol kind. +type OriginArtifactIndex struct { + Origin string `json:"origin"` + Checkpoints []string `json:"checkpoints"` + Manifests []string `json:"manifests"` + Segments []string `json:"segments"` + Meta []string `json:"meta"` + Raw []string `json:"raw"` +} + +// OriginCheckpointSummary describes the latest checkpoint published by one origin. +type OriginCheckpointSummary struct { + Sequence int + SessionCount int + ModTime time.Time + Found bool +} + +// OriginCheckpointLanding adds durable local provenance to a checkpoint summary. +type OriginCheckpointLanding struct { + OriginCheckpointSummary + LandedSessionCount int +} + +func (s *PeerArtifactSpool) Close() error { + if s == nil || s.File == nil { + return nil + } + file := s.File + s.File = nil + return closeAndRemoveTransportSpool(file) +} + +func CheckpointLandingStatusFromStore( + ctx context.Context, + store ArtifactStore, + origin string, + database any, + isLocal bool, +) (OriginCheckpointLanding, error) { + if err := ctx.Err(); err != nil { + return OriginCheckpointLanding{}, err + } + if store == nil { + return OriginCheckpointLanding{}, fmt.Errorf("%w: artifact store is required", ErrArtifactInvalid) + } + if err := validateOriginID(origin); err != nil { + return OriginCheckpointLanding{}, fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + summary, cp, err := latestStoreCheckpointSummary(ctx, store, origin) + if err != nil { + return OriginCheckpointLanding{}, err + } + return checkpointLandingFromSummary(ctx, summary, cp, origin, database, isLocal) +} + +// CheckpointLandingStatusAtStoreHead reads exactly one provenance-selected +// checkpoint. It never enumerates checkpoint history, so status work is +// bounded by the requested origin page and the selected checkpoint's session +// map. A missing or corrupt selected head is returned to the caller instead of +// silently falling back to older history. +func CheckpointLandingStatusAtStoreHead( + ctx context.Context, + store ArtifactStore, + origin string, + sequence int, + expected Identity, + database any, + isLocal bool, +) (OriginCheckpointLanding, error) { + if err := ctx.Err(); err != nil { + return OriginCheckpointLanding{}, err + } + if store == nil || sequence < 1 { + return OriginCheckpointLanding{}, fmt.Errorf("%w: artifact store and checkpoint sequence are required", + ErrArtifactInvalid) + } + ref, err := NewRef(origin, KindCheckpoints, fmt.Sprintf("cp-%010d.json", sequence)) + if err != nil { + return OriginCheckpointLanding{}, err + } + listed, err := store.Stat(ctx, ref) + if err != nil { + return OriginCheckpointLanding{}, err + } + if expected.SHA256 != "" && listed.Identity != expected { + return OriginCheckpointLanding{}, fmt.Errorf("%w: recorded checkpoint identity changed", + ErrArtifactCorrupt) + } + if listed.Identity.Size > checkpointDecodedLimit { + return OriginCheckpointLanding{}, fmt.Errorf("%w: checkpoint exceeds decoded limit", + ErrArtifactInvalid) + } + opened, reader, err := store.Open(ctx, ref) + if err != nil { + return OriginCheckpointLanding{}, err + } + if opened != listed { + _ = reader.Close() + return OriginCheckpointLanding{}, fmt.Errorf("%w: checkpoint catalog identity changed", + ErrArtifactCorrupt) + } + data, readErr := io.ReadAll(io.LimitReader(reader, checkpointDecodedLimit+1)) + verifyErr := reader.Verify() + closeErr := reader.Close() + if err := errors.Join(readErr, verifyErr, closeErr); err != nil { + return OriginCheckpointLanding{}, fmt.Errorf("%w: reading checkpoint: %v", + ErrArtifactCorrupt, err) + } + if int64(len(data)) != listed.Identity.Size || int64(len(data)) > checkpointDecodedLimit { + return OriginCheckpointLanding{}, fmt.Errorf("%w: checkpoint size mismatch", + ErrArtifactCorrupt) + } + var cp checkpoint + if err := json.Unmarshal(data, &cp); err != nil { + return OriginCheckpointLanding{}, fmt.Errorf("%w: decoding checkpoint: %v", + ErrArtifactCorrupt, err) + } + if err := validateCheckpoint(&cp, origin); err != nil { + return OriginCheckpointLanding{}, err + } + if err := validateCheckpointSequenceIdentity(cp, ref.Name); err != nil { + return OriginCheckpointLanding{}, err + } + summary := OriginCheckpointSummary{ + Sequence: cp.Sequence, SessionCount: len(cp.Sessions), + ModTime: listed.Modified, Found: true, + } + return checkpointLandingFromSummary(ctx, summary, &cp, origin, database, isLocal) +} + +func latestStoreCheckpointSummary( + ctx context.Context, store ArtifactStore, origin string, +) (_ OriginCheckpointSummary, _ *checkpoint, retErr error) { + iterator, err := openStoreEntryIterator(ctx, store, origin, KindCheckpoints) + if err != nil { + return OriginCheckpointSummary{}, nil, err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + var latest OriginCheckpointSummary + var latestCheckpoint *checkpoint + for { + entries, nextErr := iterator.Next(ctx, checkpointFloorPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return OriginCheckpointSummary{}, nil, nextErr + } + for _, entry := range entries { + if entry.Identity.Size > checkpointDecodedLimit { + continue + } + opened, reader, err := store.Open(ctx, entry.Ref) + if errors.Is(err, ErrArtifactNotFound) || errors.Is(err, ErrArtifactCorrupt) { + continue + } + if err != nil { + return OriginCheckpointSummary{}, nil, err + } + if opened.Ref != entry.Ref || opened.Identity != entry.Identity { + _ = reader.Close() + continue + } + data, readErr := io.ReadAll(io.LimitReader(reader, checkpointDecodedLimit+1)) + verifyErr := reader.Verify() + closeErr := reader.Close() + if err := ctx.Err(); err != nil { + return OriginCheckpointSummary{}, nil, err + } + if readErr != nil || verifyErr != nil || closeErr != nil || + int64(len(data)) > checkpointDecodedLimit { + continue + } + var candidate checkpoint + if err := json.Unmarshal(data, &candidate); err != nil { + continue + } + if err := validateCheckpoint(&candidate, origin); err != nil { + continue + } + if err := validateCheckpointSequenceIdentity(candidate, entry.Ref.Name); err != nil { + continue + } + if !latest.Found || candidate.Sequence > latest.Sequence { + copy := candidate + latestCheckpoint = © + latest = OriginCheckpointSummary{ + Sequence: candidate.Sequence, SessionCount: len(candidate.Sessions), + ModTime: entry.Modified, Found: true, + } + } + } + if errors.Is(nextErr, io.EOF) { + return latest, latestCheckpoint, nil + } + } +} + +func checkpointLandingFromSummary( + ctx context.Context, + summary OriginCheckpointSummary, + cp *checkpoint, + origin string, + database any, + isLocal bool, +) (OriginCheckpointLanding, error) { + status := OriginCheckpointLanding{OriginCheckpointSummary: summary} + if cp == nil || len(cp.Sessions) == 0 { + return status, nil + } + if isLocal { + reader, ok := database.(interface { + GetArtifactCheckpointHead(context.Context, string) (db.ArtifactCheckpointHead, bool, error) + StreamArtifactPublications(context.Context, string, func(db.ArtifactPublication) error) (int64, error) + }) + if !ok { + return status, nil + } + head, found, err := reader.GetArtifactCheckpointHead(ctx, origin) + if err != nil || !found || head.Sequence != cp.Sequence { + return status, err + } + landed := 0 + revision, err := reader.StreamArtifactPublications(ctx, origin, + func(publication db.ArtifactPublication) error { + if cp.Sessions[origin+"~"+publication.SessionID] == publication.ManifestHash { + landed++ + } + return nil + }) + if err != nil || revision != head.PublicationRevision { + return status, err + } + status.LandedSessionCount = landed + return status, nil + } + reader, ok := database.(interface { + StreamArtifactCheckpointLanding( + context.Context, string, func(string, string) error, + ) (db.ArtifactCheckpointLanding, bool, error) + }) + if !ok { + return status, nil + } + landed := 0 + landing, found, err := reader.StreamArtifactCheckpointLanding(ctx, origin, + func(gid, manifestHash string) error { + if cp.Sessions[gid] == manifestHash { + landed++ + } + return nil + }) + if err != nil || !found || landing.Sequence != cp.Sequence { + return status, err + } + status.LandedSessionCount = landed + return status, nil +} + +func validateArtifactName(name string) error { + if strings.TrimSpace(name) == "" { + return fmt.Errorf("%w: artifact name is required", ErrArtifactInvalid) + } + if strings.Contains(name, "/") || strings.Contains(name, "\\") || strings.Contains(name, "..") { + return fmt.Errorf("%w: invalid artifact name", ErrArtifactInvalid) + } + return nil +} + +func normalizeMetadataName(name string) (filename, hash string, err error) { + base := strings.TrimSuffix(name, metadataEventExtension) + idx := strings.LastIndex(base, "-") + if idx < 0 { + return "", "", fmt.Errorf("%w: metadata artifact missing hash suffix", ErrArtifactInvalid) + } + hash = base[idx+1:] + if err := validateHashHex(hash); err != nil { + return "", "", err + } + return base + metadataEventExtension, hash, nil +} + +func normalizeCheckpointName(name string) (string, error) { + base := strings.TrimSuffix(name, ".json") + if _, err := checkpointSequence(base + ".json"); err != nil { + return "", err + } + return base + ".json", nil +} + +func checkpointSequence(filename string) (int, error) { + base := strings.TrimSuffix(filename, ".json") + if len(base) != len("cp-0000000000") || !strings.HasPrefix(base, "cp-") { + return 0, fmt.Errorf("%w: invalid checkpoint name", ErrArtifactInvalid) + } + seq := 0 + for _, r := range base[len("cp-"):] { + if r < '0' || r > '9' { + return 0, fmt.Errorf("%w: invalid checkpoint name", ErrArtifactInvalid) + } + seq = seq*10 + int(r-'0') + } + if seq <= 0 { + return 0, fmt.Errorf("%w: invalid checkpoint sequence", ErrArtifactInvalid) + } + return seq, nil +} + +func validateHashHex(hash string) error { + if len(hash) != 64 { + return fmt.Errorf("%w: invalid artifact hash", ErrArtifactInvalid) + } + for _, r := range hash { + if (r < '0' || r > '9') && (r < 'a' || r > 'f') { + return fmt.Errorf("%w: invalid artifact hash", ErrArtifactInvalid) + } + } + return nil +} + +func validateCheckpointData(data []byte, origin, filename string) error { + var cp checkpoint + if err := json.Unmarshal(data, &cp); err != nil { + return fmt.Errorf("%w: decoding checkpoint: %v", ErrArtifactInvalid, err) + } + if cp.Version > formatVersion { + if cp.Origin != origin { + return fmt.Errorf( + "%w: checkpoint origin mismatch for %s: got %q", + ErrArtifactInvalid, origin, cp.Origin, + ) + } + if err := validateCheckpointSequenceIdentity(cp, filename); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + if err := validateCheckpointReferences(&cp, origin); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + return nil + } + if err := validateCheckpoint(&cp, origin); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + if err := validateCheckpointSequenceIdentity(cp, filename); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + return nil +} + +func validateCheckpointSequenceIdentity(cp checkpoint, filename string) error { + seq, err := checkpointSequence(filename) + if err != nil { + return err + } + if cp.Sequence != seq { + return fmt.Errorf( + "checkpoint sequence mismatch: name has %d, body has %d", + seq, cp.Sequence, + ) + } + return nil +} + +func validateCanonicalManifestArtifactData(decoded []byte, origin string) error { + m, err := decodeManifestWithLimits(decoded, productionArtifactLimits()) + if err != nil { + return fmt.Errorf("%w: decoding manifest: %v", ErrArtifactInvalid, err) + } + if m.Version > formatVersion { + if m.Origin != origin { + return fmt.Errorf("%w: manifest origin mismatch for %s: got %q", ErrArtifactInvalid, origin, m.Origin) + } + return nil + } + if m.Origin != origin { + return fmt.Errorf("%w: manifest origin mismatch for %s: got %q", ErrArtifactInvalid, origin, m.Origin) + } + if m.NativeSessionID == "" || m.Session.ID != m.NativeSessionID || m.Session.Machine != origin { + return fmt.Errorf("%w: manifest session identity mismatch", ErrArtifactInvalid) + } + if len(m.Segments) == 0 { + return fmt.Errorf("%w: manifest has no message segments", ErrArtifactInvalid) + } + if err := validateManifestReferences(m); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + return nil +} + +func validateCanonicalSegmentArtifactData(decoded []byte) error { + if _, err := decodeSegment(decoded); err != nil { + if errors.Is(err, errFutureArtifactVersion) { + if ferr := validateFutureSegmentData(decoded); ferr != nil { + return fmt.Errorf("%w: decoding segment: %v", ErrArtifactInvalid, ferr) + } + return nil + } + return fmt.Errorf("%w: decoding segment: %v", ErrArtifactInvalid, err) + } + return nil +} + +func validateFutureSegmentData(data []byte) error { + records, err := segmentRecords(data, maxSegmentMessages) + if err != nil { + return err + } + for _, line := range records { + var record struct { + Version int `json:"v"` + } + if err := json.Unmarshal(line, &record); err != nil { + return err + } + if record.Version <= formatVersion { + return fmt.Errorf("message segment has unsupported artifact version %d", record.Version) + } + } + return nil +} + +func validateMetadataArtifactData(data []byte, origin, filename, hash string) error { + if got := hashHex(data); got != hash { + return fmt.Errorf("%w: metadata artifact hash mismatch: got %s", ErrArtifactInvalid, got) + } + var envelope metadataEventEnvelope + if err := json.Unmarshal(data, &envelope); err != nil { + return fmt.Errorf("%w: decoding metadata event: %v", ErrArtifactInvalid, err) + } + if envelope.Version > formatVersion { + return nil + } + base := strings.TrimSuffix(filename, metadataEventExtension) + hlc := strings.TrimSuffix(base, "-"+hash) + var event metadataEvent + if err := json.Unmarshal(data, &event); err != nil { + return fmt.Errorf("%w: decoding metadata event: %v", ErrArtifactInvalid, err) + } + art := metadataArtifact{ + path: filename, + hlc: hlc, + hash: hash, + event: event, + } + if err := validateMetadataArtifactEvent(art, origin); err != nil { + if errors.Is(err, errFutureArtifactVersion) { + return nil + } + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + // Unknown ops stay accepted for forward compatibility (replay marks + // them applied and skips them), but a known op must carry a payload + // that projects cleanly or replay could never apply it. + if err := validateMetadataOp(event.Op); err == nil { + if _, _, _, _, err := metadataProjectionFields(event); err != nil { + return fmt.Errorf("%w: metadata event payload: %v", ErrArtifactInvalid, err) + } + } + return nil +} diff --git a/internal/artifact/peer_test.go b/internal/artifact/peer_test.go new file mode 100644 index 000000000..f7cbadf42 --- /dev/null +++ b/internal/artifact/peer_test.go @@ -0,0 +1,227 @@ +package artifact + +import ( + "bytes" + "encoding/json" + "strings" + "testing" + + "github.com/klauspost/compress/zstd" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestPeerArtifactRejectsInvalidReferences(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + + checkpointData := []byte(`{"origin":"peer-a1b2c3","seq":1,"sessions":{"peer-a1b2c3~sess-1":"../outside"},"v":1}` + "\n") + _, err := createCompressedTestArtifact( + t, store, origin, KindCheckpoints, "cp-0000000001.json", checkpointData, + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + + manifestData, err := canonicalJSON(manifest{ + Version: formatVersion, + Origin: origin, + NativeSessionID: "sess-1", + Session: manifestSession{ + ID: "sess-1", + Machine: origin, + Agent: "claude", + Project: "alpha", + CreatedAt: "2026-06-14T01:02:03Z", + }, + Segments: []string{"../outside"}, + }) + require.NoError(t, err) + manifestHash := hashHex(manifestData) + _, err = createCompressedTestArtifact( + t, store, origin, KindManifests, manifestHash, compressPeerTestData(t, manifestData), + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func TestPeerArtifactMetadataMustMatchOrigin(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + body := []byte(`{"hlc":"2026-06-14T010203.000000001Z-other-b2c3d4","op":"rename","origin":"other-b2c3d4","session_gid":"other-b2c3d4~sess-1","v":1,"value":{"display_name":"Remote"}}` + "\n") + hash := hashHex(body) + name := "2026-06-14T010203.000000001Z-other-b2c3d4-" + hash + + _, err := createCompressedTestArtifact(t, store, origin, KindMeta, name, body) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func TestPeerArtifactMetadataRejectsMalformedKnownOpPayload(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + hlc := "2026-06-14T010203.000000001Z-peer-a1b2c3" + + tests := []struct { + name string + op string + value json.RawMessage + }{ + {name: "pin missing payload", op: MetadataOpPin}, + {name: "unpin missing payload", op: MetadataOpUnpin}, + {name: "rename non-object value", op: MetadataOpRename, value: json.RawMessage(`[1,2]`)}, + {name: "rename missing value", op: MetadataOpRename}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + data, err := canonicalJSON(metadataEvent{ + Version: formatVersion, + HLC: hlc, + Origin: origin, + SessionGID: origin + "~sess-1", + Op: tt.op, + Value: tt.value, + }) + require.NoError(t, err) + name := hlc + "-" + hashHex(data) + _, err = createCompressedTestArtifact(t, store, origin, KindMeta, name, data) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + }) + } +} + +func TestPeerArtifactStoresFutureVersionArtifacts(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + + checkpointData, err := canonicalJSON(checkpoint{ + Version: formatVersion + 1, + Origin: origin, + Sequence: 1, + Sessions: map[string]string{ + origin + "~sess-1": strings64("a"), + }, + }) + require.NoError(t, err) + checkpointRef, err := createCompressedTestArtifact( + t, store, origin, KindCheckpoints, "cp-0000000001.json", checkpointData, + ) + require.NoError(t, err) + assert.Equal(t, checkpointData, readContractArtifact(t, store, checkpointRef)) + + segmentData, err := canonicalJSON(segmentMessage{ + Version: formatVersion + 1, + Ordinal: 0, + Role: "user", + Content: "future segment", + }) + require.NoError(t, err) + compressedSegment := compressPeerTestData(t, segmentData) + segmentHash := hashHex(segmentData) + segmentRef, err := createCompressedTestArtifact( + t, store, origin, KindSegments, segmentHash, compressedSegment, + ) + require.NoError(t, err) + assert.Equal(t, segmentHash+".ndjson", segmentRef.Name) + assert.Equal(t, segmentData, readContractArtifact(t, store, segmentRef)) + + futureManifestData, err := canonicalJSON(struct { + Version int `json:"v"` + Origin string `json:"origin"` + Future string `json:"future"` + }{ + Version: formatVersion + 1, + Origin: origin, + Future: "schema-owned-by-newer-peer", + }) + require.NoError(t, err) + compressedManifest := compressPeerTestData(t, futureManifestData) + manifestHash := hashHex(futureManifestData) + manifestRef, err := createCompressedTestArtifact( + t, store, origin, KindManifests, manifestHash, compressedManifest, + ) + require.NoError(t, err) + assert.Equal(t, manifestHash+".json", manifestRef.Name) + assert.Equal(t, futureManifestData, readContractArtifact(t, store, manifestRef)) + + hlc := "2026-06-14T010203.000000001Z-peer-a1b2c3" + metadataData, err := canonicalJSON(metadataEvent{ + Version: formatVersion + 1, + HLC: hlc, + Origin: origin, + SessionGID: origin + "~sess-1", + Op: "future_op", + Value: json.RawMessage(`{"future":true}`), + }) + require.NoError(t, err) + metadataHash := hashHex(metadataData) + metadataName := hlc + "-" + metadataHash + metadataRef, err := createCompressedTestArtifact( + t, store, origin, KindMeta, metadataName, metadataData, + ) + require.NoError(t, err) + assert.Equal(t, metadataName+metadataEventExtension, metadataRef.Name) + assert.Equal(t, metadataData, readContractArtifact(t, store, metadataRef)) +} + +func TestPeerArtifactStoresFutureMetadataWithoutCurrentEnvelope(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + data, err := canonicalJSON(struct { + Version int `json:"v"` + FutureClock string `json:"future_clock"` + FutureKey string `json:"future_key"` + }{ + Version: formatVersion + 1, + FutureClock: "schema-owned-by-newer-peer", + FutureKey: origin + "~sess-1", + }) + require.NoError(t, err) + hash := hashHex(data) + name := "2026-06-14T010203.000000001Z-peer-a1b2c3-" + hash + + ref, err := createCompressedTestArtifact(t, store, origin, KindMeta, name, data) + require.NoError(t, err) + assert.Equal(t, data, readContractArtifact(t, store, ref)) +} + +func TestPeerArtifactRejectsMixedFutureSegmentVersions(t *testing.T) { + store := newTestArtifactStore(t) + origin := "peer-a1b2c3" + futureLine, err := canonicalJSON(segmentMessage{ + Version: formatVersion + 1, + Ordinal: 0, + Role: "user", + Content: "future", + }) + require.NoError(t, err) + currentLine, err := canonicalJSON(segmentMessage{ + Version: formatVersion, + Ordinal: 1, + Role: "assistant", + Content: "current", + }) + require.NoError(t, err) + segmentData := append(futureLine, currentLine...) + compressed := compressPeerTestData(t, segmentData) + hash := hashHex(segmentData) + + _, err = createCompressedTestArtifact(t, store, origin, KindSegments, hash, compressed) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func compressPeerTestData(t *testing.T, data []byte) []byte { + t.Helper() + var buf bytes.Buffer + enc, err := zstd.NewWriter(&buf) + require.NoError(t, err) + _, err = enc.Write(data) + require.NoError(t, err) + require.NoError(t, enc.Close()) + return buf.Bytes() +} + +func strings64(ch string) string { + return strings.Repeat(ch, 64) +} diff --git a/internal/artifact/replay.go b/internal/artifact/replay.go new file mode 100644 index 000000000..13705f9b5 --- /dev/null +++ b/internal/artifact/replay.go @@ -0,0 +1,184 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "path/filepath" + "strings" + + "go.kenn.io/agentsview/internal/db" +) + +type metadataArtifact struct { + path string + orderKey string + hash string + hlc string + event metadataEvent +} + +type metadataEventEnvelope struct { + Version int `json:"v"` + HLC json.RawMessage `json:"hlc"` +} + +func observeMetadataStamp(clock *HLCClock, stamp HLCTimestamp, hlc string) error { + if clock == nil { + return nil + } + if _, err := clock.Observe(stamp); err != nil { + return fmt.Errorf("observing metadata HLC %s: %w", hlc, err) + } + return nil +} + +func markAppliedAndQuarantineMetadata( + ctx context.Context, + database *db.DB, + store ArtifactStore, + ref Ref, + art metadataArtifact, + reason string, +) error { + if err := database.MarkMetadataEventApplied(ctx, ref.Origin, art.orderKey, art.hash); err != nil { + return err + } + return store.Quarantine(ctx, ref, reason) +} + +func metadataArtifactOrderKey(path string) (string, error) { + base := filepath.Base(path) + name := strings.TrimSuffix(base, metadataEventExtension) + if name == base { + return "", fmt.Errorf("metadata artifact %s missing %s extension", base, metadataEventExtension) + } + return name, nil +} + +func validateMetadataArtifactEvent(art metadataArtifact, origin string) error { + if art.event.HLC != art.hlc { + return fmt.Errorf("metadata event %s HLC mismatch: got %q", art.path, art.event.HLC) + } + if art.event.Origin != origin { + return fmt.Errorf( + "metadata event %s origin mismatch for %s: got %q", + art.path, origin, art.event.Origin, + ) + } + if art.event.SessionGID == "" { + return fmt.Errorf("metadata event %s has empty session GID", art.path) + } + if art.event.Version > formatVersion { + return fmt.Errorf( + "%w: metadata event %s has artifact version %d", + errFutureArtifactVersion, art.path, art.event.Version, + ) + } + if art.event.Version != formatVersion { + return fmt.Errorf( + "metadata event %s has unsupported artifact version %d", + art.path, art.event.Version, + ) + } + // Checked after the version gate: a future format may change the HLC + // shape, but a current-version event with an unparseable HLC would + // poison raw order-key LWW comparison and must never be accepted. + if _, err := ParseHLCTimestamp(art.hlc); err != nil { + return fmt.Errorf("metadata event %s has invalid HLC: %v", art.path, err) + } + return nil +} + +func metadataProjection(art metadataArtifact, localOrigin string) (db.MetadataProjection, error) { + event := art.event + field, value, displayName, pin, err := metadataProjectionFields(event) + if err != nil { + return db.MetadataProjection{}, err + } + return db.MetadataProjection{ + EventOrigin: event.Origin, + OrderKey: art.orderKey, + HLC: event.HLC, + ArtifactHash: art.hash, + SessionGID: event.SessionGID, + LocalSessionID: metadataLocalSessionID(localOrigin, event.SessionGID), + Field: field, + Op: event.Op, + Value: value, + DisplayName: displayName, + Pin: pin, + }, nil +} + +func metadataProjectionFields( + event metadataEvent, +) (field string, value string, displayName *string, pin *db.MetadataPinProjection, err error) { + switch event.Op { + case MetadataOpRename: + var payload struct { + DisplayName *string `json:"display_name"` + } + if err := json.Unmarshal(event.Value, &payload); err != nil { + return "", "", nil, nil, fmt.Errorf("decoding rename metadata value: %w", err) + } + value, err := metadataCanonicalValue(event.Value) + return "display_name", value, payload.DisplayName, nil, err + case MetadataOpSoftDelete, MetadataOpRestore: + return "deleted_at", event.Op, nil, nil, nil + case MetadataOpStar, MetadataOpUnstar: + return "starred", event.Op, nil, nil, nil + case MetadataOpPin, MetadataOpUnpin: + if event.Pin == nil { + return "", "", nil, nil, fmt.Errorf("%s metadata event missing pin payload", event.Op) + } + value, err := metadataCanonicalPin(*event.Pin) + if err != nil { + return "", "", nil, nil, err + } + return "pin:" + metadataPinAnchor(*event.Pin), value, nil, &db.MetadataPinProjection{ + SourceUUID: event.Pin.SourceUUID, + Ordinal: event.Pin.Ordinal, + Note: event.Pin.Note, + }, nil + case MetadataOpPurge: + return "purge", event.Op, nil, nil, nil + default: + return "", "", nil, nil, fmt.Errorf("unsupported metadata event op %q", event.Op) + } +} + +func metadataLocalSessionID(localOrigin, gid string) string { + prefix := localOrigin + "~" + if after, ok := strings.CutPrefix(gid, prefix); ok { + return after + } + return gid +} + +func metadataPinAnchor(pin MetadataPin) string { + if pin.SourceUUID != "" { + return "source_uuid:" + pin.SourceUUID + } + return fmt.Sprintf("ordinal:%d", pin.Ordinal) +} + +func metadataCanonicalValue(raw json.RawMessage) (string, error) { + if len(bytes.TrimSpace(raw)) == 0 { + return "", nil + } + data, err := canonicalJSON(raw) + if err != nil { + return "", err + } + return string(bytes.TrimSpace(data)), nil +} + +func metadataCanonicalPin(pin MetadataPin) (string, error) { + data, err := canonicalJSON(pin) + if err != nil { + return "", err + } + return string(bytes.TrimSpace(data)), nil +} diff --git a/internal/artifact/replay_helpers_test.go b/internal/artifact/replay_helpers_test.go new file mode 100644 index 000000000..38f1008a5 --- /dev/null +++ b/internal/artifact/replay_helpers_test.go @@ -0,0 +1,134 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +// importFromTestStore exercises the production exact-reference importer with +// the newest checkpoint in a test fixture. Initial transports report this same +// reference after transferring its dependencies. +func importFromTestStore( + ctx context.Context, database *db.DB, store ArtifactStore, localOrigin string, +) (int, int, error) { + result, err := importResultFromTestStore(ctx, database, store, localOrigin) + return result.Sessions, result.Messages, err +} + +func importResultFromTestStore( + ctx context.Context, database *db.DB, store ArtifactStore, localOrigin string, +) (ImportResult, error) { + origins, err := store.Origins(ctx) + if err != nil { + return ImportResult{}, err + } + defer origins.Close() + coordinator := NewStoreImportCoordinator(database, store, localOrigin) + for { + page, nextErr := origins.Next(ctx, artifactImportPageSize) + for _, origin := range page { + if origin == localOrigin { + continue + } + for _, kind := range []Kind{KindCheckpoints, KindMeta} { + entries, err := testStoreEntries(ctx, store, origin, kind) + if err != nil { + return ImportResult{}, err + } + if kind == KindCheckpoints && len(entries) > 1 { + entries = entries[len(entries)-1:] + } + for _, entry := range entries { + if err := coordinator.RecordChanged(ctx, entry); err != nil { + return ImportResult{}, err + } + } + } + } + if errors.Is(nextErr, io.EOF) { + break + } + if nextErr != nil { + return ImportResult{}, nextErr + } + } + return coordinator.Finalize(ctx) +} + +func testStoreEntries( + ctx context.Context, store ArtifactStore, origin string, kind Kind, +) ([]Entry, error) { + iterator, err := store.Entries(ctx, origin, kind) + if err != nil { + return nil, err + } + defer iterator.Close() + var entries []Entry + for { + page, nextErr := iterator.Next(ctx, artifactImportPageSize) + entries = append(entries, page...) + if errors.Is(nextErr, io.EOF) { + return entries, nil + } + if nextErr != nil { + return nil, nextErr + } + } +} + +func createStoreMetadataEvent( + t *testing.T, store ArtifactStore, origin string, event metadataEvent, +) Ref { + t.Helper() + data, err := canonicalJSON(event) + require.NoError(t, err) + hash := hashHex(data) + stamp, err := ParseHLCTimestamp(event.HLC) + require.NoError(t, err) + ref, err := NewRef(origin, KindMeta, stamp.OrderingKey(hash)+metadataEventExtension) + require.NoError(t, err) + identity, err := NewIdentity(hash, int64(len(data))) + require.NoError(t, err) + _, err = store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(KindMeta), bytes.NewReader(data)) + require.NoError(t, err) + return ref +} + +func replayRenameEvent( + t *testing.T, origin, gid, hlc, displayName string, +) metadataEvent { + t.Helper() + value, err := json.Marshal(struct { + DisplayName string `json:"display_name"` + }{DisplayName: displayName}) + require.NoError(t, err) + return metadataEvent{ + Version: formatVersion, HLC: hlc, Origin: origin, + SessionGID: gid, Op: MetadataOpRename, Value: value, + } +} + +func replayTestHLC(offset time.Duration, logical uint64) string { + return HLCTimestamp{WallTime: fixedHLCTime().Add(offset), Logical: logical}.String() +} + +func assertMetadataConflictCount(t *testing.T, database *db.DB, gid, field string, want int) { + t.Helper() + var got int + err := database.Reader().QueryRowContext(context.Background(), + `SELECT COUNT(*) FROM metadata_conflicts WHERE session_gid = ? AND field = ?`, + gid, field, + ).Scan(&got) + require.NoError(t, err) + assert.Equal(t, want, got) +} diff --git a/internal/artifact/repository.go b/internal/artifact/repository.go new file mode 100644 index 000000000..b9122e16b --- /dev/null +++ b/internal/artifact/repository.go @@ -0,0 +1,322 @@ +package artifact + +import ( + "context" + "errors" + "fmt" + "io/fs" + "os" + "path/filepath" + "strings" + "sync" + + "go.kenn.io/docbank" + docsqlite "go.kenn.io/docbank/pkg/sqlite" +) + +const ( + repositoryDirectory = "artifacts" + repositoryLooseCompressionBytes = int64(4 << 10) + repositoryLooseCompressionSaving = 10 +) + +// Repository owns the process's local Docbank vault and the AgentsView pack +// scheduler layered on its physical-write receipts. Docbank owns root +// canonicalization, hierarchy locking, catalog validation, and vault reset. +type Repository struct { + mu sync.Mutex + + content *docbankStore + packer *packScheduler + root *os.Root + rootIdentity fs.FileInfo + rootPath string + driver docsqlite.Driver + closed bool +} + +// OpenRepository opens the dedicated Docbank vault below dataDir. +func OpenRepository(ctx context.Context, dataDir string) (*Repository, error) { + return openRepository(ctx, dataDir, nil) +} + +func openRepository( + ctx context.Context, dataDir string, driver docsqlite.Driver, +) (*Repository, error) { + if ctx == nil { + return nil, fmt.Errorf("%w: repository context is required", ErrArtifactInvalid) + } + if strings.TrimSpace(dataDir) == "" { + return nil, fmt.Errorf("%w: repository data directory is required", ErrArtifactInvalid) + } + rootPath, err := canonicalRepositoryRoot(dataDir) + if err != nil { + return nil, err + } + if err := rejectLegacyRepositoryLayout(rootPath); err != nil { + return nil, err + } + vault, err := docbank.New(ctx, repositoryDocbankConfig(rootPath, driver)) + if err != nil { + return nil, fmt.Errorf("opening artifact repository: %w", err) + } + return repositoryFromVault(vault, rootPath, driver) +} + +func repositoryDocbankConfig(rootPath string, driver docsqlite.Driver) docbank.Config { + return docbank.Config{ + Root: rootPath, + SQLite: driver, + LooseCompression: docbank.LooseCompressionOptions{ + Enabled: true, + MinBytes: repositoryLooseCompressionBytes, + MinSavingsPercent: repositoryLooseCompressionSaving, + }, + } +} + +func repositoryFromVault( + vault *docbank.Vault, rootPath string, driver docsqlite.Driver, +) (_ *Repository, retErr error) { + if vault == nil { + return nil, errors.New("artifact repository requires an open Docbank vault") + } + defer func() { + if retErr != nil { + retErr = errors.Join(retErr, vault.Close()) + } + }() + root, err := os.OpenRoot(rootPath) + if err != nil { + return nil, fmt.Errorf("retaining artifact repository root: %w", err) + } + rootOpen := true + defer func() { + if rootOpen { + retErr = errors.Join(retErr, root.Close()) + } + }() + identity, err := root.Stat(".") + if err != nil { + return nil, fmt.Errorf("stating artifact repository root: %w", err) + } + current, err := os.Stat(rootPath) + if err != nil || !os.SameFile(identity, current) { + return nil, errors.New("artifact repository root changed while opening") + } + + content := newDocbankContent(vault) + repository := &Repository{ + content: content, + root: root, + rootIdentity: identity, + rootPath: rootPath, + driver: driver, + } + repository.packer = newPackScheduler(content, packSchedulerOptions{}) + rootOpen = false + return repository, nil +} + +// rejectLegacyRepositoryLayout prevents an unreleased loose-artifact tree from +// being silently adopted as a Docbank vault. Existing Docbank roots are left to +// Docbank's catalog and layout validation. +func rejectLegacyRepositoryLayout(rootPath string) error { + info, err := os.Lstat(rootPath) + switch { + case errors.Is(err, fs.ErrNotExist): + return nil + case err != nil: + return fmt.Errorf("checking artifact repository layout: %w", err) + case !info.IsDir(): + return fmt.Errorf( + "artifact repository path is not a Docbank vault; move or remove %s and retry", + rootPath, + ) + } + entries, err := os.ReadDir(rootPath) + if err != nil { + return fmt.Errorf("checking artifact repository layout: %w", err) + } + if len(entries) == 0 { + return nil + } + for _, entry := range entries { + if entry.Name() == "docbank.db" || entry.Name() == "vault.lock" { + return nil + } + } + return fmt.Errorf( + "artifact repository contains the old loose artifact layout; move or remove %s and retry", + rootPath, + ) +} + +// Content returns the logical artifact boundary owned by this repository. +func (r *Repository) Content() ArtifactStore { + if r == nil { + return nil + } + return r.content +} + +// RecoverPacking seeds the repository's bounded pack scheduler from Docbank's +// indexed loose-object backlog. +func (r *Repository) RecoverPacking(ctx context.Context) error { + if r == nil || r.content == nil { + return fmt.Errorf("%w: artifact repository is required", ErrArtifactInvalid) + } + if r.packer == nil { + return ErrArtifactUnsupported + } + return r.packer.Recover(ctx) +} + +// NotifyBatch coalesces successful repository work into the asynchronous pack +// scheduler without scanning the vault. +func (r *Repository) NotifyBatch(ctx context.Context) { + if r == nil || r.content == nil || artifactMaintenanceSuppressed(ctx) { + return + } + if r.packer != nil { + r.packer.Notify(ctx) + } +} + +// RunMaintenance performs one bounded pass of each physical Docbank stage. +func (r *Repository) RunMaintenance( + ctx context.Context, opts ArtifactMaintenanceOptions, +) (PhysicalMaintenanceResult, error) { + if r == nil || r.content == nil { + return PhysicalMaintenanceResult{}, ErrArtifactUnsupported + } + return runPhysicalMaintenance(ctx, r.content, opts) +} + +// Closed reports whether repository ownership has already been released or +// transferred by an explicit reset. +func (r *Repository) Closed() bool { + if r == nil { + return true + } + r.mu.Lock() + defer r.mu.Unlock() + return r.closed +} + +// NewFolderTransport rejects a wire directory that overlaps the retained +// Docbank root before returning an opened transport. +func (r *Repository) NewFolderTransport(target string) (Transport, error) { + transport, err := openFolderTransport(target) + if err != nil { + return nil, err + } + identity, err := transport.root.Stat(".") + if err != nil { + return nil, errors.Join(err, transport.Close()) + } + if err := r.validateOpenedExternalTarget(transport.target, identity); err != nil { + return nil, errors.Join(err, transport.Close()) + } + return transport, nil +} + +// Close waits for active verified readers through Docbank, then releases the +// retained root identity. It is safe to call more than once. +func (r *Repository) Close() error { + if r == nil { + return nil + } + r.mu.Lock() + defer r.mu.Unlock() + if r.closed { + return nil + } + r.closed = true + if r.packer != nil { + r.packer.Close() + } + return errors.Join(r.content.Close(), r.root.Close()) +} + +func (r *Repository) validateOpenedExternalTarget( + targetCanonical string, targetIdentity fs.FileInfo, +) error { + if r == nil { + return fmt.Errorf("%w: artifact repository is required", ErrArtifactInvalid) + } + if strings.TrimSpace(targetCanonical) == "" || targetIdentity == nil { + return fmt.Errorf("%w: opened external artifact target is required", ErrArtifactInvalid) + } + r.mu.Lock() + defer r.mu.Unlock() + if r.closed { + return fs.ErrClosed + } + return validateOpenedRootDisjoint( + r.rootPath, r.rootIdentity, targetCanonical, targetIdentity, + ) +} + +func validateOpenedRootDisjoint( + rootPath string, + rootIdentity fs.FileInfo, + targetCanonical string, + targetIdentity fs.FileInfo, +) error { + currentTarget, err := os.Stat(targetCanonical) + if err != nil || !os.SameFile(targetIdentity, currentTarget) { + return fmt.Errorf( + "external artifact target %s changed while validating its opened identity", + targetCanonical, + ) + } + if rootsOverlap(rootPath, targetCanonical) || + targetIdentityChainContains(targetCanonical, rootIdentity) || + targetIdentityChainContains(rootPath, targetIdentity) || + os.SameFile(rootIdentity, targetIdentity) { + return fmt.Errorf( + "external artifact target %s must not overlap local artifact repository", + targetCanonical, + ) + } + currentRoot, err := os.Stat(rootPath) + if err != nil || !os.SameFile(rootIdentity, currentRoot) { + return fmt.Errorf( + "external artifact target %s cannot prove non-overlap because the local artifact repository is no longer reachable at its canonical path", + targetCanonical, + ) + } + return nil +} + +func canonicalRepositoryRoot(dataDir string) (string, error) { + dataAbs, err := filepath.Abs(dataDir) + if err != nil { + return "", fmt.Errorf("resolving artifact repository root: %w", err) + } + rootPath, err := canonicalArtifactPath(filepath.Join(dataAbs, repositoryDirectory)) + if err != nil { + return "", fmt.Errorf("resolving artifact repository root symlinks: %w", err) + } + return filepath.Clean(rootPath), nil +} + +func targetIdentityChainContains(target string, identity fs.FileInfo) bool { + current := target + for { + info, err := os.Stat(current) + if err == nil { + if os.SameFile(identity, info) { + return true + } + } else if !errors.Is(err, fs.ErrNotExist) { + return false + } + parent := filepath.Dir(current) + if parent == current { + return false + } + current = parent + } +} diff --git a/internal/artifact/repository_test.go b/internal/artifact/repository_test.go new file mode 100644 index 000000000..ee3ba4c64 --- /dev/null +++ b/internal/artifact/repository_test.go @@ -0,0 +1,158 @@ +package artifact + +import ( + "bytes" + "io" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + docsqlite "go.kenn.io/docbank/pkg/sqlite" + "go.kenn.io/docbank/pkg/sqlite/mattn" + "go.kenn.io/docbank/pkg/sqlite/modernc" +) + +func TestOpenRepositoryOwnsCompressedDocbankContent(t *testing.T) { + dataDir := t.TempDir() + repository, err := OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + body := bytes.Repeat([]byte("repository compression policy\n"), 200) + body = body[:4<<10] + result := createCheckpointBody(t, repository.Content(), 1, body) + assert.Equal(t, "zstd", result.Physical.Encoding) + assert.DirExists(t, filepath.Join(dataDir, "artifacts")) + assert.FileExists(t, filepath.Join(dataDir, "artifacts", "docbank.db")) +} + +func TestOpenRepositorySupportsDocbankSQLiteDrivers(t *testing.T) { + for _, driver := range []docsqlite.Driver{mattn.Driver{}, modernc.Driver{}} { + t.Run(driver.Name(), func(t *testing.T) { + repository, err := openRepository(t.Context(), t.TempDir(), driver) + require.NoError(t, err) + result := createCheckpointBody(t, repository.Content(), 1, []byte("driver parity")) + assert.True(t, result.Created) + require.NoError(t, repository.Close()) + }) + } +} + +func TestOpenRepositoryRejectsLegacyLooseLayoutWithoutMutation(t *testing.T) { + dataDir := t.TempDir() + artifactDir := filepath.Join(dataDir, "artifacts") + legacy := filepath.Join(artifactDir, contractOrigin, string(KindCheckpoints), "cp-0000000001.json") + require.NoError(t, os.MkdirAll(filepath.Dir(legacy), 0o755)) + original := []byte("disposable old loose artifact") + require.NoError(t, os.WriteFile(legacy, original, 0o644)) + + repository, err := OpenRepository(t.Context(), dataDir) + assert.Nil(t, repository) + assert.ErrorContains(t, err, "old loose artifact layout") + got, readErr := os.ReadFile(legacy) + require.NoError(t, readErr) + assert.Equal(t, original, got) + assert.NoFileExists(t, filepath.Join(artifactDir, "docbank.db")) +} + +func TestOpenRepositoryUsesDocbankHierarchyLock(t *testing.T) { + dataDir := t.TempDir() + repository, err := OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + second, err := OpenRepository(t.Context(), dataDir) + assert.Nil(t, second) + assert.ErrorContains(t, err, "vault is locked") + + overlapping, err := OpenRepository(t.Context(), filepath.Join(dataDir, "artifacts")) + assert.Nil(t, overlapping) + assert.ErrorContains(t, err, "vault is locked") +} + +func TestOpenRepositoryFollowsFinalRootSymlink(t *testing.T) { + realDataDir := t.TempDir() + realRoot := filepath.Join(realDataDir, "artifacts") + realRepository, err := openRepository(t.Context(), realDataDir, modernc.Driver{}) + require.NoError(t, err) + require.NoError(t, realRepository.Close()) + + aliasDataDir := t.TempDir() + require.NoError(t, os.Symlink(realRoot, filepath.Join(aliasDataDir, "artifacts"))) + aliased, err := openRepository(t.Context(), aliasDataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, aliased.Close()) }) + canonicalRoot, err := filepath.EvalSymlinks(realRoot) + require.NoError(t, err) + assert.Equal(t, canonicalRoot, aliased.rootPath) +} + +func TestRepositoryRejectsOverlappingFolderTargets(t *testing.T) { + dataDir := t.TempDir() + repository, err := OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + vaultRoot := filepath.Join(dataDir, "artifacts") + + for name, target := range map[string]string{ + "same root": vaultRoot, + "ancestor": dataDir, + "descendant": filepath.Join(vaultRoot, "blobs"), + } { + t.Run(name, func(t *testing.T) { + transport, err := repository.NewFolderTransport(target) + assert.Nil(t, transport) + assert.ErrorContains(t, err, "must not overlap") + }) + } + + transport, err := repository.NewFolderTransport(t.TempDir()) + require.NoError(t, err) + require.NoError(t, transport.(*folderTransport).Close()) +} + +func TestOpenRepositoryRetainsAbsoluteCanonicalRoot(t *testing.T) { + dataDir := t.TempDir() + workingDirectory, err := os.Getwd() + require.NoError(t, err) + relativeDataDir, err := filepath.Rel(workingDirectory, dataDir) + require.NoError(t, err) + repository, err := OpenRepository(t.Context(), relativeDataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + canonical, err := filepath.EvalSymlinks(filepath.Join(dataDir, "artifacts")) + require.NoError(t, err) + assert.True(t, filepath.IsAbs(repository.rootPath)) + assert.Equal(t, canonical, repository.rootPath) +} + +func TestRepositoryCloseWaitsForReaderAndIsIdempotent(t *testing.T) { + repository, err := OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + createContractArtifact(t, repository.Content(), ref, []byte("reader lease")) + _, reader, err := repository.Content().Open(t.Context(), ref) + require.NoError(t, err) + one := make([]byte, 1) + _, err = reader.Read(one) + require.NoError(t, err) + + closeResult := make(chan error, 1) + go func() { closeResult <- repository.Close() }() + select { + case err := <-closeResult: + require.Fail(t, "repository close returned before reader close", "error: %v", err) + case <-time.After(100 * time.Millisecond): + } + + _, err = io.Copy(io.Discard, reader) + require.NoError(t, err) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + require.NoError(t, <-closeResult) + assert.NoError(t, repository.Close()) +} diff --git a/internal/artifact/reset.go b/internal/artifact/reset.go new file mode 100644 index 000000000..28568cc26 --- /dev/null +++ b/internal/artifact/reset.go @@ -0,0 +1,506 @@ +package artifact + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io/fs" + "os" + "strconv" + "strings" + "time" + + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/docbank" + docsqlite "go.kenn.io/docbank/pkg/sqlite" +) + +const ( + // ArtifactResetManualCleanupWarning is returned verbatim by every reset + // surface so operators cannot mistake diagnostic retention for cleanup. + ArtifactResetManualCleanupWarning = "Cleanup is manual: after diagnosis, remove the moved-aside vault yourself. Repeated resets may accumulate diagnostic vaults and consume disk; AgentsView never deletes them automatically." + // ArtifactResetForeignRelayWarning explains the one class of artifacts that + // cannot be reconstructed from the authoritative local SQLite archive. + ArtifactResetForeignRelayWarning = "Foreign relay artifacts are unavailable until peers resend them; sessions already stored in SQLite are unchanged." +) + +// RepositoryResetResult names both preserved vault paths and the local +// publication recreated from SQLite. +type RepositoryResetResult struct { + VaultRoot string `json:"vault_root"` + DiagnosticRoot string `json:"diagnostic_root"` + Export ExportResult `json:"export"` +} + +type repositoryResetHooks struct { + now func() time.Time + resetVault func(context.Context, docbank.Config, docbank.ResetOptions) (*docbank.Vault, error) + driver docsqlite.Driver + export func(context.Context, *db.DB, ArtifactStore, ExportOptions) (ExportResult, error) + beforeRelease func() error +} + +// ResetRepository explicitly moves the owned Docbank vault aside, creates a +// fresh vault, and republishes this machine's artifacts from SQLite. current is +// the daemon-owned repository when reset runs in-process; direct callers pass +// nil and must already hold the SQLite write-owner lock. +func ResetRepository( + ctx context.Context, + dataDir string, + database *db.DB, + origin string, + current *Repository, +) (*Repository, RepositoryResetResult, error) { + return resetRepositoryWith(ctx, dataDir, database, origin, current, repositoryResetHooks{}) +} + +// BeginRepositoryReset validates and moves the current vault aside, then opens +// its fresh replacement. beforeRelease runs at Docbank's last boundary before +// it releases and moves the current vault; returning an error leaves the +// current vault in place. +func BeginRepositoryReset( + ctx context.Context, + dataDir string, + origin string, + current *Repository, + beforeRelease func() error, +) (*Repository, RepositoryResetResult, error) { + return beginRepositoryResetWith(ctx, dataDir, origin, current, repositoryResetHooks{ + beforeRelease: beforeRelease, + }) +} + +// RepublishRepositoryReset reconstructs the local publication in a fresh vault +// after BeginRepositoryReset. It leaves the fresh repository open on failure +// so an in-process owner can retain it and retry publication later. +func RepublishRepositoryReset( + ctx context.Context, + dataDir string, + database *db.DB, + origin string, + fresh *Repository, + result RepositoryResetResult, +) (RepositoryResetResult, error) { + return republishRepositoryResetWith( + ctx, dataDir, database, origin, fresh, result, repositoryResetHooks{}, + ) +} + +func resetRepositoryWith( + ctx context.Context, + dataDir string, + database *db.DB, + origin string, + current *Repository, + hooks repositoryResetHooks, +) (returned *Repository, result RepositoryResetResult, retErr error) { + if ctx == nil { + return nil, result, fmt.Errorf("%w: repository reset context is required", ErrArtifactInvalid) + } + if database == nil || database.ReadOnly() { + return nil, result, fmt.Errorf("%w: writable repository reset database is required", ErrArtifactInvalid) + } + originalBeforeRelease := hooks.beforeRelease + var pending db.ArtifactResetRepublishPending + markerPrepared := false + hooks.beforeRelease = func() error { + if strings.TrimSpace(origin) != "" { + var err error + pending, err = PrepareRepositoryResetRepublish(ctx, database, dataDir, origin) + if err != nil { + return err + } + markerPrepared = true + } + if originalBeforeRelease != nil { + return originalBeforeRelease() + } + return nil + } + returned, result, retErr = beginRepositoryResetWith( + ctx, dataDir, origin, current, hooks, + ) + if retErr != nil { + if markerPrepared { + if _, statErr := os.Lstat(result.DiagnosticRoot); errors.Is(statErr, fs.ErrNotExist) { + _, clearErr := database.ClearArtifactResetRepublishPending( + context.WithoutCancel(ctx), pending, + ) + retErr = errors.Join(retErr, clearErr) + } + } + return nil, result, retErr + } + result, retErr = republishRepositoryResetWith( + ctx, dataDir, database, origin, returned, result, hooks, + ) + if retErr != nil { + retErr = errors.Join(retErr, returned.Close()) + return nil, result, retErr + } + return returned, result, nil +} + +func beginRepositoryResetWith( + ctx context.Context, + dataDir string, + origin string, + current *Repository, + hooks repositoryResetHooks, +) (returned *Repository, result RepositoryResetResult, retErr error) { + if ctx == nil { + return nil, result, fmt.Errorf("%w: repository reset context is required", ErrArtifactInvalid) + } + origin = strings.TrimSpace(origin) + if origin != "" { + if err := validateOriginID(origin); err != nil { + return nil, result, err + } + } + rootPath, err := canonicalRepositoryRoot(dataDir) + if err != nil { + return nil, result, err + } + if hooks.now == nil { + hooks.now = time.Now + } + if hooks.resetVault == nil { + hooks.resetVault = docbank.ResetVault + } + if current != nil { + current.mu.Lock() + defer current.mu.Unlock() + if current.closed || current.root == nil { + return nil, result, errors.New("artifact repository is not available for reset") + } + if current.rootPath != rootPath { + return nil, result, errors.New("artifact repository reset target does not match the current owner") + } + } + + diagnosticRoot, err := nextRepositoryDiagnosticRoot(ctx, rootPath, hooks.now()) + if err != nil { + return nil, result, err + } + result = RepositoryResetResult{VaultRoot: rootPath, DiagnosticRoot: diagnosticRoot} + driver := hooks.driver + if driver == nil && current != nil { + driver = current.driver + } + config := repositoryDocbankConfig(rootPath, driver) + releaseCurrent := hooks.beforeRelease + if current != nil { + releaseCurrent = func() error { + if hooks.beforeRelease != nil { + if err := hooks.beforeRelease(); err != nil { + return err + } + } + current.closed = true + if current.packer != nil { + current.packer.Close() + } + return errors.Join(current.content.Close(), current.root.Close()) + } + } + vault, err := hooks.resetVault(ctx, config, docbank.ResetOptions{ + DiagnosticRoot: diagnosticRoot, + ReleaseCurrent: releaseCurrent, + }) + if err != nil { + return nil, result, repositoryResetErrorIfMoved(result, err) + } + + fresh, err := repositoryFromVault(vault, rootPath, driver) + if err != nil { + return nil, result, repositoryResetAfterMoveError( + result, err, + ) + } + returned = fresh + return returned, result, nil +} + +func republishRepositoryResetWith( + ctx context.Context, + dataDir string, + database *db.DB, + origin string, + fresh *Repository, + result RepositoryResetResult, + hooks repositoryResetHooks, +) (RepositoryResetResult, error) { + if ctx == nil { + return result, repositoryResetAfterMoveError( + result, fmt.Errorf("%w: repository reset context is required", ErrArtifactInvalid), + ) + } + // Cancellation must be observed before touching SQLite. Shutdown can cancel + // this phase after the filesystem commit and then tear the database down. + if err := ctx.Err(); err != nil { + return result, repositoryResetAfterMoveError(result, err) + } + if database == nil || database.ReadOnly() { + return result, repositoryResetAfterMoveError( + result, fmt.Errorf("%w: writable repository reset database is required", ErrArtifactInvalid), + ) + } + origin = strings.TrimSpace(origin) + if origin == "" { + return result, nil + } + if err := validateOriginID(origin); err != nil { + return result, repositoryResetAfterMoveError(result, err) + } + if fresh == nil || fresh.Closed() { + return result, repositoryResetAfterMoveError( + result, errors.New("fresh artifact repository is not available for republish"), + ) + } + if hooks.export == nil { + hooks.export = func( + ctx context.Context, database *db.DB, store ArtifactStore, opts ExportOptions, + ) (ExportResult, error) { + return ExportToStore(ctx, database, store, opts) + } + } + var ( + recovered bool + err error + ) + result.Export, recovered, err = recoverRepositoryResetRepublishWith( + ctx, database, fresh, origin, hooks.export, + ) + if err != nil { + return result, repositoryResetAfterMoveError(result, err) + } + if recovered { + fresh.NotifyBatch(ctx) + return result, nil + } + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, + Store: fresh.Content(), + }) + if _, err := recorder.materializeCurrentState(ctx); err != nil { + return result, repositoryResetAfterMoveError(result, err) + } + result.Export, err = hooks.export( + ctx, database, fresh.Content(), ExportOptions{Origin: origin, Full: true}, + ) + if err != nil { + return result, repositoryResetAfterMoveError(result, err) + } + fresh.NotifyBatch(ctx) + return result, nil +} + +type repositoryResetExportFunc func( + context.Context, *db.DB, ArtifactStore, ExportOptions, +) (ExportResult, error) + +// RecoverRepositoryResetRepublish completes a durable interrupted repository +// reset before ordinary publication or exchange begins. Only the store owned by +// the matching repository root may consume and clear the marker. +func RecoverRepositoryResetRepublish( + ctx context.Context, + database *db.DB, + repository *Repository, + origin string, +) (ExportResult, bool, error) { + return recoverRepositoryResetRepublishWith( + ctx, database, repository, origin, + func( + ctx context.Context, database *db.DB, store ArtifactStore, opts ExportOptions, + ) (ExportResult, error) { + return ExportToStore(ctx, database, store, opts) + }, + ) +} + +// PublishRepositoryArtifacts completes any durable interrupted reset before +// performing an ordinary export. A completed full recovery already satisfies +// the requested publication and is returned directly. +func PublishRepositoryArtifacts( + ctx context.Context, + database *db.DB, + repository *Repository, + opts ExportOptions, +) (ExportResult, error) { + result, recovered, err := RecoverRepositoryResetRepublish( + ctx, database, repository, opts.Origin, + ) + if err != nil || recovered { + return result, err + } + return ExportToStore(ctx, database, repository.Content(), opts) +} + +func recoverRepositoryResetRepublishWith( + ctx context.Context, + database *db.DB, + repository *Repository, + origin string, + export repositoryResetExportFunc, +) (ExportResult, bool, error) { + if ctx == nil { + return ExportResult{}, false, errors.New("artifact reset recovery context is required") + } + if err := ctx.Err(); err != nil { + return ExportResult{}, false, err + } + if database == nil || database.ReadOnly() { + return ExportResult{}, false, errors.New("writable artifact reset database is required") + } + origin = strings.TrimSpace(origin) + if err := validateOriginID(origin); err != nil { + return ExportResult{}, false, err + } + if repository == nil || repository.Closed() { + return ExportResult{}, false, errors.New("artifact reset repository is required") + } + store := repository.Content() + pending, found, err := database.ArtifactResetRepublishPending(ctx) + if err != nil || !found { + return ExportResult{}, false, err + } + fingerprint, err := repositoryFingerprint(repository) + if err != nil { + return ExportResult{}, false, err + } + if pending.RootFingerprint != fingerprint || pending.Origin != origin { + return ExportResult{}, false, + errors.New("artifact reset republish state does not match the repository") + } + if export == nil { + return ExportResult{}, false, errors.New("artifact reset recovery exporter is required") + } + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, + Store: store, + }) + if _, err := recorder.materializeCurrentStateAtHLC(ctx, pending.BaselineHLC); err != nil { + return ExportResult{}, false, err + } + result, err := export(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + if err != nil { + return ExportResult{}, false, err + } + cleared, err := database.ClearArtifactResetRepublishPending(ctx, pending) + if err != nil { + return ExportResult{}, false, err + } + if !cleared { + return ExportResult{}, false, + errors.New("artifact reset republish state changed before completion") + } + return result, true, nil +} + +func repositoryFingerprint(repository *Repository) (string, error) { + if repository == nil { + return "", errors.New("artifact repository is required") + } + repository.mu.Lock() + defer repository.mu.Unlock() + if repository.closed { + return "", errors.New("artifact repository is closed") + } + digest := sha256.Sum256([]byte(repository.rootPath)) + return hex.EncodeToString(digest[:]), nil +} + +// PrepareRepositoryResetRepublish advances the local metadata clock and +// persists the durable reset intent before Docbank releases the current vault. +// Callers must invoke their short lifecycle commit only after this returns. +func PrepareRepositoryResetRepublish( + ctx context.Context, + database *db.DB, + dataDir string, + origin string, +) (db.ArtifactResetRepublishPending, error) { + if ctx == nil { + return db.ArtifactResetRepublishPending{}, errors.New("artifact reset republish context is required") + } + if err := ctx.Err(); err != nil { + return db.ArtifactResetRepublishPending{}, err + } + if database == nil || database.ReadOnly() { + return db.ArtifactResetRepublishPending{}, errors.New("writable artifact reset database is required") + } + origin = strings.TrimSpace(origin) + if err := validateOriginID(origin); err != nil { + return db.ArtifactResetRepublishPending{}, err + } + fingerprint, err := repositoryRootFingerprint(dataDir) + if err != nil { + return db.ArtifactResetRepublishPending{}, err + } + stamp, err := NewHLCClock(database, HLCClockOptions{}).Next() + if err != nil { + return db.ArtifactResetRepublishPending{}, err + } + tokenBytes := make([]byte, 32) + if _, err := rand.Read(tokenBytes); err != nil { + return db.ArtifactResetRepublishPending{}, fmt.Errorf("generating artifact reset token: %w", err) + } + pending := db.ArtifactResetRepublishPending{ + Version: 1, + RootFingerprint: fingerprint, + Origin: origin, + Token: hex.EncodeToString(tokenBytes), + BaselineHLC: stamp.String(), + } + if err := database.SetArtifactResetRepublishPending(ctx, pending); err != nil { + return db.ArtifactResetRepublishPending{}, err + } + return pending, nil +} + +func repositoryRootFingerprint(dataDir string) (string, error) { + root, err := canonicalRepositoryRoot(dataDir) + if err != nil { + return "", err + } + digest := sha256.Sum256([]byte(root)) + return hex.EncodeToString(digest[:]), nil +} + +func nextRepositoryDiagnosticRoot( + ctx context.Context, rootPath string, now time.Time, +) (string, error) { + base := rootPath + ".reset-" + now.UTC().Format("20060102T150405.000000000Z") + for sequence := 0; ; sequence++ { + if err := ctx.Err(); err != nil { + return "", err + } + candidate := base + if sequence > 0 { + candidate += "." + strconv.Itoa(sequence) + } + _, err := os.Lstat(candidate) + if errors.Is(err, fs.ErrNotExist) { + return candidate, nil + } + if err != nil { + return "", fmt.Errorf("checking artifact reset diagnostic path: %w", err) + } + } +} + +func repositoryResetErrorIfMoved(result RepositoryResetResult, err error) error { + if _, statErr := os.Lstat(result.DiagnosticRoot); statErr == nil { + return repositoryResetAfterMoveError(result, err) + } + return err +} + +func repositoryResetAfterMoveError(result RepositoryResetResult, err error) error { + return fmt.Errorf( + "artifact reset failed after moving the original vault to %s; the fresh vault path is %s; preserve both paths for manual recovery: %w", + result.DiagnosticRoot, result.VaultRoot, err, + ) +} diff --git a/internal/artifact/reset_test.go b/internal/artifact/reset_test.go new file mode 100644 index 000000000..eff8478ee --- /dev/null +++ b/internal/artifact/reset_test.go @@ -0,0 +1,882 @@ +package artifact + +import ( + "bytes" + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/docbank" + docsqlite "go.kenn.io/docbank/pkg/sqlite" + "go.kenn.io/docbank/pkg/sqlite/modernc" +) + +type cancelAfterMetadataCreateStore struct { + ArtifactStore + cancel context.CancelFunc + remaining int +} + +type resetRecoveryObservingTransport struct { + database *db.DB + origin string + exchanges int +} + +func (t *resetRecoveryObservingTransport) Prepare(context.Context, ArtifactStore) error { + return nil +} + +func (t *resetRecoveryObservingTransport) Exchange(ctx context.Context, store ArtifactStore) error { + t.exchanges++ + _, pending, err := t.database.ArtifactResetRepublishPending(ctx) + if err != nil { + return err + } + if pending { + return errors.New("transport exchange observed pending reset recovery") + } + page, err := firstStoreEntryPage(ctx, store, t.origin, KindMeta, 10) + if err != nil { + return err + } + if len(page.Items) != 1 { + return fmt.Errorf("transport exchange observed %d metadata events, want 1", len(page.Items)) + } + return nil +} + +func (s *cancelAfterMetadataCreateStore) Create( + ctx context.Context, + ref Ref, + identity Identity, + mediaType string, + body io.Reader, +) (CreateResult, error) { + result, err := s.ArtifactStore.Create(ctx, ref, identity, mediaType, body) + if err == nil && ref.Kind == KindMeta && s.remaining > 0 { + s.remaining-- + if s.remaining == 0 { + s.cancel() + } + } + return result, err +} + +func TestArtifactResetCorruptCatalogNormalStartupFailsClosed(t *testing.T) { + dataDir := t.TempDir() + database, err := db.Open(filepath.Join(dataDir, "sessions.db")) + require.NoError(t, err) + seedSession(t, database, "local-session", "project-a") + database.Close() + + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + require.NoError(t, repository.Close()) + catalog := filepath.Join(dataDir, repositoryDirectory, "docbank.db") + require.NoError(t, os.WriteFile(catalog, []byte("corrupt docbank catalog"), 0o600)) + before := snapshotDirectory(t, dataDir) + + reopened, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + + assert.Nil(t, reopened) + require.Error(t, err) + assert.Equal(t, before, snapshotDirectory(t, dataDir)) +} + +func TestArtifactResetStoppedDaemonMovesAsideAndRepublishesFromSQLiteFloor(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + seedSession(t, database, "local-session", "project-a") + localDisplayName := "Recovered local title" + require.NoError(t, database.RenameSession("local-session", &localDisplayName)) + origin := "desktop-d4e5f6" + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + first, err := ExportToStore(t.Context(), database, repository.Content(), ExportOptions{ + Origin: origin, Full: true, + }) + require.NoError(t, err) + foreignBody := []byte("foreign relay bytes") + foreignRef := requireContractRef(t, "peer-a1b2c3", KindRaw, hashHex(foreignBody)) + _, err = repository.Content().Create(t.Context(), foreignRef, identityForBytes(t, foreignBody), + canonicalArtifactMediaType(KindRaw), bytes.NewReader(foreignBody)) + require.NoError(t, err) + require.NoError(t, repository.Close()) + + fresh, result, err := resetRepositoryWith( + t.Context(), dataDir, database, origin, nil, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, fresh.Close()) }) + + wantRoot, err := canonicalRepositoryRoot(dataDir) + require.NoError(t, err) + assert.Equal(t, wantRoot, result.VaultRoot) + assert.Equal(t, result.VaultRoot+".reset-20260722T123456.123456789Z", result.DiagnosticRoot) + assert.DirExists(t, result.DiagnosticRoot) + assert.GreaterOrEqual(t, result.Export.CheckpointSequence, first.CheckpointSequence) + assert.Equal(t, []string{origin}, listAllContractOrigins(t, fresh.Content(), 10)) + assert.NotEmpty(t, listAllContractEntries(t, fresh.Content(), origin, KindMeta, 10), + "reset must publish local curation that existed only in SQLite") + _, err = fresh.Content().Stat(t.Context(), foreignRef) + require.ErrorIs(t, err, ErrArtifactNotFound) + replayDB := testDB(t) + replay, err := importResultFromTestStore(t.Context(), replayDB, fresh.Content(), "receiver-a1b2c3") + require.NoError(t, err) + assert.Positive(t, replay.Metadata) + replayed, err := replayDB.GetSession(t.Context(), origin+"~local-session") + require.NoError(t, err) + require.NotNil(t, replayed) + require.NotNil(t, replayed.DisplayName) + assert.Equal(t, localDisplayName, *replayed.DisplayName) + + diagnostic, err := docbank.New(t.Context(), docbank.Config{ + Root: result.DiagnosticRoot, SQLite: modernc.Driver{}, + }) + require.NoError(t, err) + diagnosticStore := newDocbankContent(diagnostic) + t.Cleanup(func() { require.NoError(t, diagnosticStore.Close()) }) + _, err = diagnosticStore.Stat(t.Context(), foreignRef) + require.NoError(t, err) +} + +func TestArtifactResetRematerializesCoveredLocalCurationBeforeFreshCheckpoint(t *testing.T) { + ctx := t.Context() + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + origin := "desktop-d4e5f6" + seedSession(t, database, "local-session", "project-a") + require.NoError(t, database.ReplaceSessionMessages("local-session", []db.Message{ + {SessionID: "local-session", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5, SourceUUID: "uuid-question"}, + {SessionID: "local-session", Ordinal: 1, Role: "assistant", Content: "world", ContentLength: 5, SourceUUID: "uuid-answer"}, + })) + + bootstrap, err := openRepository(ctx, t.TempDir(), modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, bootstrap.Close()) }) + _, err = ExportToStore(ctx, database, bootstrap.Content(), ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + + displayName := "Recovered curated title" + require.NoError(t, database.RenameSession("local-session", &displayName)) + starred, err := database.StarSession("local-session") + require.NoError(t, err) + assert.True(t, starred) + messages, err := database.GetAllMessages(ctx, "local-session") + require.NoError(t, err) + require.Len(t, messages, 2) + note := "keep this answer" + _, err = database.PinMessage("local-session", messages[1].ID, ¬e) + require.NoError(t, err) + require.NoError(t, database.SoftDeleteSession("local-session")) + + current, err := openRepository(ctx, dataDir, modernc.Driver{}) + require.NoError(t, err) + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, + Store: current.Content(), + }) + written, err := recorder.AppendBaseline(ctx) + require.NoError(t, err) + assert.Equal(t, 4, written) + + materializedBeforeExport := false + fresh, _, err := resetRepositoryWith( + ctx, dataDir, database, origin, current, + repositoryResetHooks{ + now: fixedArtifactResetTime, + export: func( + ctx context.Context, database *db.DB, store ArtifactStore, opts ExportOptions, + ) (ExportResult, error) { + materializedBeforeExport = len(listAllContractEntries( + t, store, origin, KindMeta, 10, + )) == 4 + return ExportToStore(ctx, database, store, opts) + }, + }, + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, fresh.Close()) }) + assert.True(t, materializedBeforeExport, + "fresh metadata must exist before ExportToStore creates the checkpoint") + + replayDB := testDB(t) + _, err = importResultFromTestStore(ctx, replayDB, bootstrap.Content(), "receiver-a1b2c3") + require.NoError(t, err) + replayed, err := replayDB.GetSessionFull(ctx, origin+"~local-session") + require.NoError(t, err) + require.NotNil(t, replayed) + require.NotNil(t, replayed.DisplayName) + assert.NotEqual(t, displayName, *replayed.DisplayName) + + replayedResult, err := importResultFromTestStore( + ctx, replayDB, fresh.Content(), "receiver-a1b2c3", + ) + require.NoError(t, err) + assert.Equal(t, 4, replayedResult.Metadata) + replayed, err = replayDB.GetSessionFull(ctx, origin+"~local-session") + require.NoError(t, err) + require.NotNil(t, replayed) + require.NotNil(t, replayed.DisplayName) + assert.Equal(t, displayName, *replayed.DisplayName) + require.NotNil(t, replayed.DeletedAt) + stars, err := replayDB.ListStarredSessionIDs(ctx) + require.NoError(t, err) + assert.Equal(t, []string{origin + "~local-session"}, stars) + pins, err := replayDB.ListPinnedMessages(ctx, origin+"~local-session", "") + require.NoError(t, err) + require.Len(t, pins, 1) + assert.Equal(t, 1, pins[0].Ordinal) + require.NotNil(t, pins[0].Note) + assert.Equal(t, note, *pins[0].Note) +} + +func TestArtifactResetRematerializesLocalNegativeWinnersWithoutReattributingForeign(t *testing.T) { + ctx := t.Context() + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + localOrigin := "desktop-d4e5f6" + peerOrigin := "laptop-a1b2c3" + seedSession(t, database, "cleared-session", "project-a") + seedSession(t, database, "foreign-session", "project-a") + require.NoError(t, database.ReplaceSessionMessages("cleared-session", []db.Message{ + {SessionID: "cleared-session", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5, SourceUUID: "uuid-question"}, + {SessionID: "cleared-session", Ordinal: 1, Role: "assistant", Content: "world", ContentLength: 5, SourceUUID: "uuid-answer"}, + })) + + bootstrap, err := openRepository(ctx, t.TempDir(), modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, bootstrap.Close()) }) + _, err = ExportToStore(ctx, database, bootstrap.Content(), ExportOptions{ + Origin: localOrigin, + Full: true, + }) + require.NoError(t, err) + + current, err := openRepository(ctx, dataDir, modernc.Driver{}) + require.NoError(t, err) + localRecorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: localOrigin, + Store: current.Content(), + Now: fixedHLCTime, + }) + staleName := "stale local title" + require.NoError(t, database.RenameSession("cleared-session", &staleName)) + renameValue, err := metadataRenameValue(&staleName) + require.NoError(t, err) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpRename, Value: renameValue, + }) + require.NoError(t, err) + _, err = database.StarSession("cleared-session") + require.NoError(t, err) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpStar, + }) + require.NoError(t, err) + messages, err := database.GetAllMessages(ctx, "cleared-session") + require.NoError(t, err) + require.Len(t, messages, 2) + note := "stale pin" + _, err = database.PinMessage("cleared-session", messages[1].ID, ¬e) + require.NoError(t, err) + pin := &MetadataPin{SourceUUID: "uuid-answer", Ordinal: 1, Note: ¬e} + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpPin, Pin: pin, + }) + require.NoError(t, err) + require.NoError(t, database.SoftDeleteSession("cleared-session")) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpSoftDelete, + }) + require.NoError(t, err) + + require.NoError(t, database.RenameSession("cleared-session", nil)) + renameValue, err = metadataRenameValue(nil) + require.NoError(t, err) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpRename, Value: renameValue, + }) + require.NoError(t, err) + _, err = database.UnstarSession("cleared-session") + require.NoError(t, err) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpUnstar, + }) + require.NoError(t, err) + require.NoError(t, database.UnpinMessage("cleared-session", messages[1].ID)) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpUnpin, Pin: pin, + }) + require.NoError(t, err) + _, err = database.RestoreSession("cleared-session") + require.NoError(t, err) + _, err = localRecorder.Append(ctx, MetadataEventInput{ + SessionID: "cleared-session", Op: MetadataOpRestore, + }) + require.NoError(t, err) + + peerDB := testDB(t) + peerRepository, err := openRepository(ctx, t.TempDir(), modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, peerRepository.Close()) }) + peerRecorder := NewMetadataRecorder(peerDB, MetadataRecorderOptions{ + Origin: peerOrigin, + Store: peerRepository.Content(), + Now: func() time.Time { return fixedHLCTime().Add(time.Hour) }, + }) + peerName := "peer-authored winner" + peerValue, err := metadataRenameValue(&peerName) + require.NoError(t, err) + _, err = peerRecorder.Append(ctx, MetadataEventInput{ + SessionID: localOrigin + "~foreign-session", Op: MetadataOpRename, Value: peerValue, + }) + require.NoError(t, err) + _, err = importResultFromTestStore(ctx, database, peerRepository.Content(), localOrigin) + require.NoError(t, err) + + fresh, _, err := resetRepositoryWith( + ctx, dataDir, database, localOrigin, current, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, fresh.Close()) }) + + metadataEntries := listAllContractEntries(t, fresh.Content(), localOrigin, KindMeta, 10) + require.Len(t, metadataEntries, 4) + wantOps := map[string]bool{ + MetadataOpRename: true, + MetadataOpUnstar: true, + MetadataOpUnpin: true, + MetadataOpRestore: true, + } + for _, entry := range metadataEntries { + data := readContractArtifact(t, fresh.Content(), entry.Ref) + var event metadataEvent + require.NoError(t, json.Unmarshal(data, &event)) + assert.Equal(t, localOrigin+"~cleared-session", event.SessionGID) + assert.True(t, wantOps[event.Op], "unexpected materialized op %q", event.Op) + delete(wantOps, event.Op) + } + assert.Empty(t, wantOps) + + replayDB := testDB(t) + _, err = importResultFromTestStore(ctx, replayDB, bootstrap.Content(), "receiver-a1b2c3") + require.NoError(t, err) + require.NoError(t, replayDB.RenameSession(localOrigin+"~cleared-session", &staleName)) + _, err = replayDB.StarSession(localOrigin + "~cleared-session") + require.NoError(t, err) + replayMessages, err := replayDB.GetAllMessages(ctx, localOrigin+"~cleared-session") + require.NoError(t, err) + require.Len(t, replayMessages, 2) + _, err = replayDB.PinMessage(localOrigin+"~cleared-session", replayMessages[1].ID, ¬e) + require.NoError(t, err) + require.NoError(t, replayDB.SoftDeleteSession(localOrigin+"~cleared-session")) + + result, err := importResultFromTestStore(ctx, replayDB, fresh.Content(), "receiver-a1b2c3") + require.NoError(t, err) + assert.Equal(t, 4, result.Metadata) + replayed, err := replayDB.GetSessionFull(ctx, localOrigin+"~cleared-session") + require.NoError(t, err) + require.NotNil(t, replayed) + assert.Nil(t, replayed.DeletedAt) + require.NotNil(t, replayed.DisplayName) + assert.NotEqual(t, staleName, *replayed.DisplayName) + stars, err := replayDB.ListStarredSessionIDs(ctx) + require.NoError(t, err) + assert.Empty(t, stars) + pins, err := replayDB.ListPinnedMessages(ctx, localOrigin+"~cleared-session", "") + require.NoError(t, err) + assert.Empty(t, pins) +} + +func TestArtifactResetRepublishCrashAcrossWinnerPageIsIdempotent(t *testing.T) { + const winnerPageSize = 128 + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + localOrigin := "desktop-d4e5f6" + peerOrigin := "laptop-a1b2c3" + for index := range winnerPageSize + 1 { + recordResetMetadataWinner( + t, database, localOrigin, + fmt.Sprintf("%s~session-%03d", localOrigin, index), + MetadataOpUnstar, index, + ) + } + recordResetMetadataWinner( + t, database, peerOrigin, localOrigin+"~foreign-session", + MetadataOpRestore, winnerPageSize+2, + ) + + current, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + fresh, _, err := beginRepositoryResetWith( + t.Context(), dataDir, localOrigin, current, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + require.NoError(t, err) + + firstCtx, cancel := context.WithCancel(t.Context()) + interruptingStore := &cancelAfterMetadataCreateStore{ + ArtifactStore: fresh.Content(), + cancel: cancel, + remaining: winnerPageSize + 1, + } + firstRecorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: localOrigin, + Store: interruptingStore, + }) + baselineHLC := HLCTimestamp{WallTime: fixedHLCTime()}.String() + _, err = firstRecorder.materializeCurrentStateAtHLC(firstCtx, baselineHLC) + require.ErrorIs(t, err, context.Canceled) + require.NoError(t, fresh.Close()) + + reopened, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, reopened.Close()) }) + retryRecorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: localOrigin, + Store: reopened.Content(), + }) + _, err = retryRecorder.materializeCurrentStateAtHLC(t.Context(), baselineHLC) + require.NoError(t, err) + + entries := listAllContractEntries(t, reopened.Content(), localOrigin, KindMeta, 64) + assert.Len(t, entries, winnerPageSize+1, + "retry must immutable-create the same winner artifacts instead of appending replacements") + assert.Empty(t, listAllContractEntries(t, reopened.Content(), peerOrigin, KindMeta, 64), + "reset recovery must not reauthor foreign winners") +} + +func TestArtifactResetRepublishCrashAfterBaselineCreateDoesNotGrowEvents(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + seedSession(t, database, "local-session", "project-a") + _, err := database.StarSession("local-session") + require.NoError(t, err) + localOrigin := "desktop-d4e5f6" + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + + firstCtx, cancel := context.WithCancel(t.Context()) + interruptingStore := &cancelAfterMetadataCreateStore{ + ArtifactStore: repository.Content(), cancel: cancel, remaining: 1, + } + firstRecorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: localOrigin, + Store: interruptingStore, + Now: fixedHLCTime, + }) + baselineHLC := HLCTimestamp{WallTime: fixedHLCTime()}.String() + _, err = firstRecorder.materializeCurrentStateAtHLC(firstCtx, baselineHLC) + require.ErrorIs(t, err, context.Canceled) + require.NoError(t, repository.Close()) + + reopened, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, reopened.Close()) }) + retryRecorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: localOrigin, + Store: reopened.Content(), + Now: fixedHLCTime, + }) + _, err = retryRecorder.materializeCurrentStateAtHLC(t.Context(), baselineHLC) + require.NoError(t, err) + + entries := listAllContractEntries(t, reopened.Content(), localOrigin, KindMeta, 10) + assert.Len(t, entries, 1, + "retry after Create-before-projection crash must reuse one deterministic baseline event") +} + +func TestArtifactResetRepublishCrashAcrossUncoveredBaselinePageIsIdempotent(t *testing.T) { + const baselineRows = 129 + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + origin := "desktop-d4e5f6" + for index := range baselineRows { + sessionID := fmt.Sprintf("session-%03d", index) + seedSession(t, database, sessionID, "project-a") + _, err := database.StarSession(sessionID) + require.NoError(t, err) + } + pending, err := PrepareRepositoryResetRepublish(t.Context(), database, dataDir, origin) + require.NoError(t, err) + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + + firstCtx, cancel := context.WithCancel(t.Context()) + interruptingStore := &cancelAfterMetadataCreateStore{ + ArtifactStore: repository.Content(), cancel: cancel, remaining: 128, + } + firstRecorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, + Store: interruptingStore, + }) + _, err = firstRecorder.materializeCurrentStateAtHLC(firstCtx, pending.BaselineHLC) + require.ErrorIs(t, err, context.Canceled) + _, markerFound, err := database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.True(t, markerFound) + require.NoError(t, repository.Close()) + + reopened, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, reopened.Close()) }) + _, recovered, err := RecoverRepositoryResetRepublish( + t.Context(), database, reopened, origin, + ) + require.NoError(t, err) + assert.True(t, recovered) + assert.Len(t, listAllContractEntries(t, reopened.Content(), origin, KindMeta, 64), baselineRows, + "retry must complete uncovered baseline pages without duplicate event growth") + _, markerFound, err = database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.False(t, markerFound) +} + +func TestRecoverRepositoryResetRepublishClearsMarkerAfterCheckpoint(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + origin := "desktop-d4e5f6" + recordResetMetadataWinner( + t, database, origin, origin+"~local-session", MetadataOpUnstar, 1, + ) + _, err := PrepareRepositoryResetRepublish(t.Context(), database, dataDir, origin) + require.NoError(t, err) + + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + result, recovered, err := RecoverRepositoryResetRepublish( + t.Context(), database, repository, origin, + ) + require.NoError(t, err) + assert.True(t, recovered) + assert.Positive(t, result.CheckpointSequence) + assert.Len(t, listAllContractEntries(t, repository.Content(), origin, KindMeta, 10), 1) + _, found, err := database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.False(t, found, "checkpoint completion must CAS-clear the durable marker") + + second, recovered, err := RecoverRepositoryResetRepublish( + t.Context(), database, repository, origin, + ) + require.NoError(t, err) + assert.False(t, recovered) + assert.Zero(t, second.CheckpointSequence) + assert.Len(t, listAllContractEntries(t, repository.Content(), origin, KindMeta, 10), 1) +} + +func TestRepublishRepositoryResetNotifiesOwnedPacking(t *testing.T) { + repository, err := OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + repository.packer.Close() + packer := newBlockingPacker() + repository.packer = newPackScheduler(packer, packSchedulerOptions{RetryDelay: time.Hour}) + + _, err = republishRepositoryResetWith( + t.Context(), t.TempDir(), testDB(t), "desktop-d4e5f6", repository, + RepositoryResetResult{}, repositoryResetHooks{ + export: func( + context.Context, *db.DB, ArtifactStore, ExportOptions, + ) (ExportResult, error) { + return ExportResult{}, nil + }, + }, + ) + require.NoError(t, err) + require.Eventually(t, func() bool { return packer.calls.Load() == 1 }, + time.Second, time.Millisecond) +} + +func TestSyncRepositoryRecoversResetBeforeFirstExchange(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + origin := "desktop-d4e5f6" + recordResetMetadataWinner( + t, database, origin, origin+"~local-session", MetadataOpUnstar, 1, + ) + _, err := PrepareRepositoryResetRepublish(t.Context(), database, dataDir, origin) + require.NoError(t, err) + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + transport := &resetRecoveryObservingTransport{database: database, origin: origin} + + _, err = syncRepositoryWithTransport( + t.Context(), database, repository, SyncOptions{Origin: origin}, transport, + ) + require.NoError(t, err) + assert.Equal(t, 1, transport.exchanges) +} + +func recordResetMetadataWinner( + t *testing.T, + database *db.DB, + origin string, + sessionGID string, + op string, + index int, +) { + t.Helper() + stamp := HLCTimestamp{ + WallTime: fixedHLCTime().Add(time.Duration(index) * time.Nanosecond), + } + event := metadataEvent{ + Version: formatVersion, + HLC: stamp.String(), + Origin: origin, + SessionGID: sessionGID, + Op: op, + } + data, err := canonicalJSON(event) + require.NoError(t, err) + hash := hashHex(data) + projection, err := metadataProjection(metadataArtifact{ + orderKey: stamp.OrderingKey(hash), + hash: hash, + hlc: event.HLC, + event: event, + }, origin) + require.NoError(t, err) + _, err = database.RecordLocalMetadataProjection(t.Context(), projection) + require.NoError(t, err) +} + +func TestArtifactResetDaemonOwnerTransfersRepositoryReservation(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + seedSession(t, database, "local-session", "project-a") + current, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + oldStore := current.Content() + + fresh, result, err := resetRepositoryWith( + t.Context(), dataDir, database, "desktop-d4e5f6", current, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, fresh.Close()) }) + + assert.DirExists(t, result.DiagnosticRoot) + _, err = firstStoreOriginPage(t.Context(), oldStore, 10) + require.Error(t, err) + contender, err := OpenRepository(t.Context(), dataDir) + assert.Nil(t, contender) + require.ErrorContains(t, err, "vault is locked") + require.NoError(t, current.Close(), "the transferred owner must be inert") +} + +func TestArtifactResetLockConflictDoesNotMutateVault(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + current, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, current.Close()) }) + before := snapshotDirectory(t, filepath.Join(dataDir, repositoryDirectory)) + + fresh, result, err := resetRepositoryWith( + t.Context(), dataDir, database, "desktop-d4e5f6", nil, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + wantRoot, rootErr := canonicalRepositoryRoot(dataDir) + require.NoError(t, rootErr) + + assert.Nil(t, fresh) + assert.Equal(t, wantRoot, result.VaultRoot) + assert.NoDirExists(t, result.DiagnosticRoot) + require.ErrorContains(t, err, "vault is locked") + assert.Equal(t, before, snapshotDirectory(t, filepath.Join(dataDir, repositoryDirectory))) +} + +func TestArtifactResetRejectsInvalidOriginBeforeReleasingCurrentVault(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + current, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, current.Close()) }) + + fresh, result, err := resetRepositoryWith( + t.Context(), dataDir, database, "invalid origin", current, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + + assert.Nil(t, fresh) + assert.Empty(t, result) + require.ErrorContains(t, err, "invalid artifact origin") + assert.False(t, current.Closed()) + assert.NoDirExists(t, filepath.Join(dataDir, repositoryDirectory)+".reset-20260722T123456.123456789Z") +} + +func TestArtifactResetRepeatedTimestampAccumulatesDiagnosticVaults(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + repository, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + require.NoError(t, repository.Close()) + + first, firstResult, err := resetRepositoryWith( + t.Context(), dataDir, database, "", nil, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + require.NoError(t, err) + require.NoError(t, first.Close()) + second, secondResult, err := resetRepositoryWith( + t.Context(), dataDir, database, "", nil, + repositoryResetHooks{now: fixedArtifactResetTime}, + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, second.Close()) }) + + assert.NotEqual(t, firstResult.DiagnosticRoot, secondResult.DiagnosticRoot) + assert.Equal(t, firstResult.DiagnosticRoot+".1", secondResult.DiagnosticRoot) + assert.DirExists(t, firstResult.DiagnosticRoot) + assert.DirExists(t, secondResult.DiagnosticRoot) +} + +func TestArtifactResetFailureBoundariesPreserveRecoveryPaths(t *testing.T) { + for _, tt := range []struct { + name string + hooks func(*testing.T, error) repositoryResetHooks + wantMoved bool + wantSource bool + }{ + { + name: "move", + hooks: func(t *testing.T, injected error) repositoryResetHooks { + return repositoryResetHooks{ + now: fixedArtifactResetTime, + resetVault: func(_ context.Context, _ docbank.Config, opts docbank.ResetOptions) (*docbank.Vault, error) { + require.NoError(t, opts.ReleaseCurrent()) + return nil, injected + }, + } + }, + wantSource: true, + }, + { + name: "fresh initialization", + hooks: func(_ *testing.T, injected error) repositoryResetHooks { + return repositoryResetHooks{now: fixedArtifactResetTime, driver: artifactResetFailDriver{err: injected}} + }, + wantMoved: true, wantSource: true, + }, + { + name: "local republish", + hooks: func(_ *testing.T, injected error) repositoryResetHooks { + return repositoryResetHooks{ + now: fixedArtifactResetTime, + export: func(context.Context, *db.DB, ArtifactStore, ExportOptions) (ExportResult, error) { + return ExportResult{}, injected + }, + } + }, + wantMoved: true, wantSource: true, + }, + } { + t.Run(tt.name, func(t *testing.T) { + dataDir := t.TempDir() + database := openArtifactResetDB(t, dataDir) + seedSession(t, database, "local-session", "project-a") + beforeSession, err := database.GetSession(t.Context(), "local-session") + require.NoError(t, err) + _, beforeFloorFound, err := database.GetArtifactCheckpointFloor( + t.Context(), "desktop-d4e5f6", + ) + require.NoError(t, err) + assert.False(t, beforeFloorFound) + current, err := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, err) + vaultRoot := filepath.Join(dataDir, repositoryDirectory) + injected := errors.New("injected " + tt.name + " failure") + + fresh, result, err := resetRepositoryWith( + t.Context(), dataDir, database, "desktop-d4e5f6", current, tt.hooks(t, injected), + ) + + assert.Nil(t, fresh) + require.ErrorIs(t, err, injected) + if tt.wantMoved { + require.ErrorContains(t, err, result.VaultRoot) + require.ErrorContains(t, err, result.DiagnosticRoot) + assert.DirExists(t, result.DiagnosticRoot) + } else { + assert.NoDirExists(t, result.DiagnosticRoot) + recovered, openErr := openRepository(t.Context(), dataDir, modernc.Driver{}) + require.NoError(t, openErr) + require.NoError(t, recovered.Close()) + } + if tt.wantSource { + assert.DirExists(t, vaultRoot) + } + afterSession, dbErr := database.GetSession(t.Context(), "local-session") + require.NoError(t, dbErr) + assert.Equal(t, beforeSession, afterSession) + _, afterFloorFound, dbErr := database.GetArtifactCheckpointFloor( + t.Context(), "desktop-d4e5f6", + ) + require.NoError(t, dbErr) + assert.False(t, afterFloorFound, "failed reset must not advance SQLite publication floor") + }) + } +} + +func fixedArtifactResetTime() time.Time { + return time.Date(2026, 7, 22, 12, 34, 56, 123456789, time.UTC) +} + +func openArtifactResetDB(t *testing.T, dataDir string) *db.DB { + t.Helper() + database, err := db.Open(filepath.Join(dataDir, "sessions.db")) + require.NoError(t, err) + t.Cleanup(func() { database.Close() }) + return database +} + +func snapshotDirectory(t *testing.T, root string) map[string][]byte { + t.Helper() + snapshot := make(map[string][]byte) + require.NoError(t, filepath.WalkDir(root, func(path string, entry os.DirEntry, err error) error { + if err != nil { + return err + } + relative, err := filepath.Rel(root, path) + if err != nil { + return err + } + if entry.IsDir() { + snapshot[relative+string(filepath.Separator)] = nil + return nil + } + content, err := os.ReadFile(path) + if err != nil { + return err + } + snapshot[relative] = content + return nil + })) + return snapshot +} + +type artifactResetFailDriver struct{ err error } + +func (d artifactResetFailDriver) Name() string { return "artifact reset failing driver" } +func (d artifactResetFailDriver) Open(string, docsqlite.OpenOptions) (*sql.DB, error) { + return nil, d.err +} +func (artifactResetFailDriver) IsBusy(error) bool { return false } +func (artifactResetFailDriver) IsUniqueViolation(error) bool { return false } diff --git a/internal/artifact/store.go b/internal/artifact/store.go new file mode 100644 index 000000000..f2636480f --- /dev/null +++ b/internal/artifact/store.go @@ -0,0 +1,332 @@ +package artifact + +import ( + "context" + "errors" + "fmt" + "io" + "strings" + "time" +) + +var ( + ErrArtifactCorrupt = errors.New("artifact corrupt") + ErrArtifactUnsupported = errors.New("artifact unsupported") +) + +const maxArtifactListPageSize = 5000 + +func validateStoreRef(ref Ref) error { + canonical, err := NewRef(ref.Origin, ref.Kind, ref.Name) + if err != nil { + return err + } + if canonical != ref { + return fmt.Errorf("%w: noncanonical artifact reference", ErrArtifactInvalid) + } + return nil +} + +func validateStoreCollection(origin string, kind Kind) error { + if err := validateOriginID(origin); err != nil { + return fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + switch kind { + case KindCheckpoints, KindManifests, KindSegments, KindMeta, KindRaw: + return nil + default: + return fmt.Errorf("%w: unsupported artifact kind %q", ErrArtifactInvalid, kind) + } +} + +func validateStoreIdentity(identity Identity) error { + canonical, err := NewIdentity(identity.SHA256, identity.Size) + if err != nil { + return err + } + if canonical != identity { + return fmt.Errorf("%w: noncanonical artifact identity", ErrArtifactInvalid) + } + return nil +} + +func validateRefIdentity(ref Ref, identity Identity) error { + sha, err := refIdentitySHA(ref) + if err != nil { + return err + } + if sha != "" && sha != identity.SHA256 { + return fmt.Errorf("%w: reference and content identity differ", ErrArtifactInvalid) + } + return nil +} + +func refIdentitySHA(ref Ref) (string, error) { + switch ref.Kind { + case KindCheckpoints: + return "", nil + case KindManifests: + return strings.TrimSuffix(ref.Name, ".json"), nil + case KindSegments: + return strings.TrimSuffix(ref.Name, ".ndjson"), nil + case KindMeta: + _, sha, err := normalizeMetadataName(ref.Name) + return sha, err + case KindRaw: + return ref.Name, nil + default: + return "", fmt.Errorf("%w: unsupported artifact kind %q", ErrArtifactInvalid, ref.Kind) + } +} + +func canonicalArtifactMediaType(kind Kind) string { + switch kind { + case KindCheckpoints, KindManifests, KindMeta: + return "application/json" + case KindSegments: + return "application/x-ndjson" + case KindRaw: + return "application/octet-stream" + default: + return "" + } +} + +type contextArtifactReader struct { + ctx context.Context + reader io.Reader +} + +func (r *contextArtifactReader) Read(p []byte) (int, error) { + if err := r.ctx.Err(); err != nil { + return 0, err + } + return r.reader.Read(p) +} + +func artifactStoreError(op string, ref Ref, err error) error { + return &ArtifactOpError{Op: op, Ref: ref, Err: err} +} + +// Kind identifies one artifact protocol collection. +type Kind string + +// Ref is a canonical logical artifact reference. Name never includes a wire +// compression extension. +type Ref struct { + Origin string + Kind Kind + Name string +} + +// NewRef validates and constructs a canonical logical artifact reference. +func NewRef(origin string, kind Kind, name string) (Ref, error) { + if err := validateOriginID(origin); err != nil { + return Ref{}, fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + if err := validateCanonicalArtifactName(kind, name); err != nil { + return Ref{}, err + } + return Ref{Origin: origin, Kind: kind, Name: name}, nil +} + +func validateCanonicalArtifactName(kind Kind, name string) error { + if err := validateArtifactName(name); err != nil { + return err + } + switch kind { + case KindCheckpoints: + canonical, err := normalizeCheckpointName(name) + if err != nil { + return err + } + if canonical != name { + return fmt.Errorf("%w: checkpoint name is not canonical", ErrArtifactInvalid) + } + case KindManifests: + if err := validateCanonicalHashName(name, ".json"); err != nil { + return err + } + case KindSegments: + if err := validateCanonicalHashName(name, ".ndjson"); err != nil { + return err + } + case KindMeta: + canonical, _, err := normalizeMetadataName(name) + if err != nil { + return err + } + if canonical != name { + return fmt.Errorf("%w: metadata name is not canonical", ErrArtifactInvalid) + } + case KindRaw: + if err := validateHashHex(name); err != nil { + return err + } + default: + return fmt.Errorf("%w: unsupported artifact kind %q", ErrArtifactInvalid, kind) + } + return nil +} + +func validateCanonicalHashName(name, extension string) error { + if !strings.HasSuffix(name, extension) { + return fmt.Errorf("%w: artifact name must end in %s", ErrArtifactInvalid, extension) + } + hash := strings.TrimSuffix(name, extension) + if err := validateHashHex(hash); err != nil { + return err + } + return nil +} + +// Identity is the canonical uncompressed content identity of an artifact. +type Identity struct { + SHA256 string + Size int64 +} + +// NewIdentity validates and constructs a canonical content identity. +func NewIdentity(sha256 string, size int64) (Identity, error) { + if err := validateHashHex(sha256); err != nil { + return Identity{}, err + } + if size < 0 { + return Identity{}, fmt.Errorf("%w: artifact size must not be negative", ErrArtifactInvalid) + } + return Identity{SHA256: sha256, Size: size}, nil +} + +// Entry describes one live logical artifact. +type Entry struct { + Ref Ref + Identity Identity + Modified time.Time +} + +// Cursor is an opaque stable-enumeration continuation token. +type Cursor string + +// Page is one bounded page of logical artifacts. +type Page struct { + Items []Entry + Next Cursor +} + +// ArtifactOpError adds logical operation context without hiding its cause. +type ArtifactOpError struct { + Op string + Ref Ref + Err error +} + +func (e *ArtifactOpError) Error() string { + if e == nil { + return "artifact operation failed" + } + target := e.Ref.Origin + if e.Ref.Kind != "" { + target += "/" + string(e.Ref.Kind) + } + if e.Ref.Name != "" { + target += "/" + e.Ref.Name + } + if target == "" { + return fmt.Sprintf("artifact %s: %v", e.Op, e.Err) + } + return fmt.Sprintf("artifact %s %s: %v", e.Op, target, e.Err) +} + +func (e *ArtifactOpError) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} + +// VerifiedReader yields authoritative bytes only after terminal io.EOF or a +// successful Verify. Closing an incomplete read does not drain it implicitly. +type VerifiedReader interface { + io.ReadCloser + Verify() error +} + +// CreateResult describes an immutable logical create or identical retry. +type CreateResult struct { + Entry Entry + Created bool + Physical PhysicalWrite +} + +// PhysicalWrite reports the physical effect of a logical create. +type PhysicalWrite struct { + Kind string + Encoding string + LogicalBytes int64 + StoredBytes int64 + PackEligible bool +} + +// LooseBacklog describes unpacked content eligible for bounded packing. +type LooseBacklog struct { + EligibleObjects int64 + EligibleBytes int64 + EligibleStoredBytes int64 +} + +// PackResult describes one bounded physical packing pass. +type PackResult struct { + PackedObjects int + LogicalBytes int64 + More bool +} + +// ArtifactStore stores canonical artifact bytes behind logical references. +type ArtifactStore interface { + Create(context.Context, Ref, Identity, string, io.Reader) (CreateResult, error) + Stat(context.Context, Ref) (Entry, error) + Open(context.Context, Ref) (Entry, VerifiedReader, error) + Origins(context.Context) (OriginIterator, error) + Entries(context.Context, string, Kind) (EntryIterator, error) + Quarantine(context.Context, Ref, string) error + Trash(context.Context, Ref) error + RepairContent(context.Context, Identity, io.Reader) error +} + +// OriginIterator owns one stable origin traversal until EOF or Close. Next may +// use a different bounded page size on each call and returns io.EOF with the +// final non-empty page. +type OriginIterator interface { + Next(context.Context, int) ([]string, error) + Close() error +} + +// EntryIterator owns one stable collection traversal until EOF or Close. Next +// may use a different bounded page size on each call and returns io.EOF with +// the final non-empty page. +type EntryIterator interface { + Next(context.Context, int) ([]Entry, error) + Close() error +} + +// QuarantineIterator owns one stable traversal of hidden artifacts until EOF +// or Close. +type QuarantineIterator interface { + Next(context.Context, int) ([]QuarantinedEntry, error) + Close() error +} + +// WorkBudget bounds one archive-wide maintenance pass. +type WorkBudget struct { + MaxObjects int + MaxBytes int64 + Cursor string +} + +// MaintenanceResult reports bounded maintenance progress and continuation. +type MaintenanceResult struct { + Processed int + Bytes int64 + NextCursor string + More bool +} diff --git a/internal/artifact/store_contract_test.go b/internal/artifact/store_contract_test.go new file mode 100644 index 000000000..1fa27388c --- /dev/null +++ b/internal/artifact/store_contract_test.go @@ -0,0 +1,662 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "io/fs" + "strings" + "sync" + "syscall" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/docbank" + docsqlite "go.kenn.io/docbank/pkg/sqlite" + "go.kenn.io/docbank/pkg/sqlite/mattn" + "go.kenn.io/docbank/pkg/sqlite/modernc" +) + +const contractOrigin = "contract-a1b2c3" + +func TestNewRefValidatesCanonicalReferences(t *testing.T) { + hash := strings.Repeat("a", 64) + tests := []struct { + name string + kind Kind + ref string + }{ + {name: "checkpoint", kind: KindCheckpoints, ref: "cp-0000000001.json"}, + {name: "manifest", kind: KindManifests, ref: hash + ".json"}, + {name: "segment", kind: KindSegments, ref: hash + ".ndjson"}, + {name: "metadata", kind: KindMeta, ref: "20260721T010203.000000000Z-0-" + hash + ".json"}, + {name: "raw", kind: KindRaw, ref: hash}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := NewRef(contractOrigin, tt.kind, tt.ref) + require.NoError(t, err) + assert.Equal(t, Ref{Origin: contractOrigin, Kind: tt.kind, Name: tt.ref}, got) + }) + } +} + +func TestNewRefRejectsNoncanonicalReferences(t *testing.T) { + hash := strings.Repeat("a", 64) + tests := []struct { + name string + origin string + kind Kind + ref string + }{ + {name: "missing origin", kind: KindRaw, ref: hash}, + {name: "unknown kind", origin: contractOrigin, kind: "future", ref: hash}, + {name: "checkpoint without extension", origin: contractOrigin, kind: KindCheckpoints, ref: "cp-0000000001"}, + {name: "manifest wire extension", origin: contractOrigin, kind: KindManifests, ref: hash + ".json.zst"}, + {name: "segment wire extension", origin: contractOrigin, kind: KindSegments, ref: hash + ".ndjson.zst"}, + {name: "metadata without extension", origin: contractOrigin, kind: KindMeta, ref: "clock-" + hash}, + {name: "uppercase hash", origin: contractOrigin, kind: KindRaw, ref: strings.Repeat("A", 64)}, + {name: "path separator", origin: contractOrigin, kind: KindRaw, ref: "../" + hash}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := NewRef(tt.origin, tt.kind, tt.ref) + assert.ErrorIs(t, err, ErrArtifactInvalid) + }) + } +} + +func TestNewIdentityValidatesCanonicalSHA256AndSize(t *testing.T) { + hash := strings.Repeat("a", 64) + identity, err := NewIdentity(hash, 0) + require.NoError(t, err) + assert.Equal(t, Identity{SHA256: hash, Size: 0}, identity) + + tests := []struct { + name string + hash string + size int64 + }{ + {name: "missing hash"}, + {name: "short hash", hash: "abcd"}, + {name: "uppercase hash", hash: strings.Repeat("A", 64)}, + {name: "negative size", hash: hash, size: -1}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := NewIdentity(tt.hash, tt.size) + assert.ErrorIs(t, err, ErrArtifactInvalid) + }) + } +} + +func TestArtifactOpErrorPreservesCause(t *testing.T) { + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + err := &ArtifactOpError{Op: "open", Ref: ref, Err: context.Canceled} + + assert.ErrorIs(t, err, context.Canceled) + assert.Contains(t, err.Error(), "open") + assert.Contains(t, err.Error(), ref.Name) + + transient := fmt.Errorf("writing artifact: %w", syscall.EAGAIN) + err = &ArtifactOpError{Op: "create", Ref: ref, Err: transient} + assert.ErrorIs(t, err, syscall.EAGAIN) +} + +type artifactStoreFactory func(t *testing.T) ArtifactStore + +func TestArtifactStoreContractDocbank(t *testing.T) { + for _, driver := range []docsqlite.Driver{mattn.Driver{}, modernc.Driver{}} { + t.Run(driver.Name(), func(t *testing.T) { + runArtifactStoreContract(t, func(t *testing.T) ArtifactStore { + vault, err := docbank.New(t.Context(), docbank.Config{ + Root: t.TempDir(), + SQLite: driver, + }) + require.NoError(t, err) + return newDocbankContent(vault) + }) + }) + } +} + +func TestArtifactStoreIteratorContractDocbank(t *testing.T) { + for _, driver := range []docsqlite.Driver{mattn.Driver{}, modernc.Driver{}} { + t.Run(driver.Name(), func(t *testing.T) { + store := newContractStore(t, func(t *testing.T) ArtifactStore { + vault, err := docbank.New(t.Context(), docbank.Config{ + Root: t.TempDir(), SQLite: driver, + }) + require.NoError(t, err) + return newDocbankContent(vault) + }) + iterable := store + + originalOrigins := []string{ + "alpha-a1b2c3", + "charlie-c3d4e5", + "echo-e5f6a7", + "foxtrot-f6a7b8", + } + for i, origin := range originalOrigins { + ref := requireContractRef(t, origin, KindCheckpoints, "cp-0000000001.json") + createContractArtifact(t, store, ref, []byte{byte(i + 1)}) + } + + origins, err := iterable.Origins(t.Context()) + require.NoError(t, err) + firstOrigins, err := origins.Next(t.Context(), 1) + require.NoError(t, err) + assert.Equal(t, originalOrigins[:1], firstOrigins) + + insertedOrigin := requireContractRef( + t, "bravo-b2c3d4", KindCheckpoints, "cp-0000000001.json", + ) + createContractArtifact(t, store, insertedOrigin, []byte("inserted origin")) + restOrigins, err := origins.Next(t.Context(), 3) + require.ErrorIs(t, err, io.EOF) + assert.Equal(t, originalOrigins[1:], restOrigins) + _, err = origins.Next(t.Context(), 1) + assert.ErrorIs(t, err, io.EOF) + require.NoError(t, origins.Close()) + + originalNames := []string{ + "cp-0000000001.json", + "cp-0000000003.json", + "cp-0000000005.json", + "cp-0000000007.json", + } + for _, name := range originalNames { + ref := requireContractRef(t, contractOrigin, KindCheckpoints, name) + createContractArtifact(t, store, ref, []byte(name)) + } + entries, err := iterable.Entries(t.Context(), contractOrigin, KindCheckpoints) + require.NoError(t, err) + firstEntries, err := entries.Next(t.Context(), 1) + require.NoError(t, err) + assert.Equal(t, originalNames[:1], entryNames(firstEntries)) + + insertedEntry := requireContractRef( + t, contractOrigin, KindCheckpoints, "cp-0000000002.json", + ) + createContractArtifact(t, store, insertedEntry, []byte(insertedEntry.Name)) + restEntries, err := entries.Next(t.Context(), 3) + require.ErrorIs(t, err, io.EOF) + assert.Equal(t, originalNames[1:], entryNames(restEntries)) + _, err = entries.Next(t.Context(), 1) + assert.ErrorIs(t, err, io.EOF) + require.NoError(t, entries.Close()) + + cancelEntries, err := iterable.Entries(t.Context(), contractOrigin, KindCheckpoints) + require.NoError(t, err) + _, err = cancelEntries.Next(t.Context(), 1) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + _, err = cancelEntries.Next(ctx, 1) + assert.ErrorIs(t, err, context.Canceled) + _, err = cancelEntries.Next(t.Context(), 1) + assert.ErrorIs(t, err, fs.ErrClosed) + + closedOrigins, err := iterable.Origins(t.Context()) + require.NoError(t, err) + require.NoError(t, closedOrigins.Close()) + require.NoError(t, closedOrigins.Close()) + _, err = closedOrigins.Next(t.Context(), 1) + assert.ErrorIs(t, err, fs.ErrClosed) + }) + } +} + +// runArtifactStoreContract exercises only the logical API. Backends register +// this helper from their own top-level tests and may not expose physical paths +// or implementation handles to it. +func runArtifactStoreContract(t *testing.T, factory artifactStoreFactory) { + t.Helper() + + t.Run("missing reads", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + + _, err := store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) + _, reader, err := store.Open(t.Context(), ref) + assert.Nil(t, reader) + assert.ErrorIs(t, err, ErrArtifactNotFound) + }) + + t.Run("identical retry is immutable and idempotent", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte(`{"origin":"contract-a1b2c3","sequence":1}`) + identity := identityForBytes(t, body) + + first, err := store.Create(t.Context(), ref, identity, "application/json", bytes.NewReader(body)) + require.NoError(t, err) + assert.True(t, first.Created) + assert.Equal(t, ref, first.Entry.Ref) + assert.Equal(t, identity, first.Entry.Identity) + assert.False(t, first.Entry.Modified.IsZero()) + + retry, err := store.Create(t.Context(), ref, identity, "application/json", bytes.NewReader(body)) + require.NoError(t, err) + assert.False(t, retry.Created) + assert.Equal(t, first.Entry, retry.Entry) + assert.Equal(t, body, readContractArtifact(t, store, ref)) + }) + + t.Run("different expected identity conflicts without mutation", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + original := []byte("original checkpoint") + replacement := []byte("replacement checkpoint") + createContractArtifact(t, store, ref, original) + + _, err := store.Create(t.Context(), ref, identityForBytes(t, replacement), + "application/json", bytes.NewReader(replacement)) + assert.ErrorIs(t, err, ErrArtifactConflict) + assert.Equal(t, original, readContractArtifact(t, store, ref)) + }) + + t.Run("duplicate path rejects stream mismatching expected identity", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + original := []byte("original checkpoint") + replacement := []byte("different bytes with the original expected identity") + first := createContractArtifact(t, store, ref, original) + + _, err := store.Create(t.Context(), ref, first.Entry.Identity, + "application/json", bytes.NewReader(replacement)) + assert.ErrorIs(t, err, ErrArtifactInvalid) + entry, err := store.Stat(t.Context(), ref) + require.NoError(t, err) + assert.Equal(t, first.Entry, entry) + assert.Equal(t, original, readContractArtifact(t, store, ref)) + }) + + t.Run("duplicate path rejects media type mismatch", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte("original checkpoint") + first := createContractArtifact(t, store, ref, body) + + _, err := store.Create(t.Context(), ref, first.Entry.Identity, + "application/octet-stream", bytes.NewReader(body)) + assert.ErrorIs(t, err, ErrArtifactConflict) + entry, err := store.Stat(t.Context(), ref) + require.NoError(t, err) + assert.Equal(t, first.Entry, entry) + retry, err := store.Create(t.Context(), ref, first.Entry.Identity, + "application/json", bytes.NewReader(body)) + require.NoError(t, err) + assert.False(t, retry.Created) + assert.Equal(t, first.Entry, retry.Entry) + assert.Equal(t, body, readContractArtifact(t, store, ref)) + }) + + t.Run("expected identity mismatch creates nothing", func(t *testing.T) { + tests := []struct { + name string + identity func(t *testing.T, body []byte) Identity + }{ + { + name: "hash", + identity: func(t *testing.T, body []byte) Identity { + different := []byte("malicious checkpoint") + require.Len(t, different, len(body)) + require.NotEqual(t, body, different) + return identityForBytes(t, different) + }, + }, + { + name: "size", + identity: func(t *testing.T, body []byte) Identity { + correct := identityForBytes(t, body) + identity, err := NewIdentity(correct.SHA256, correct.Size+1) + require.NoError(t, err) + return identity + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte("canonical checkpoint") + + _, err := store.Create(t.Context(), ref, tt.identity(t, body), + "application/json", bytes.NewReader(body)) + assert.ErrorIs(t, err, ErrArtifactInvalid) + _, err = store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) + }) + } + + t.Run("hash-bearing reference", func(t *testing.T) { + body := []byte("canonical artifact content") + identity := identityForBytes(t, body) + other := identityForBytes(t, []byte("different artifact content")) + tests := []struct { + name string + kind Kind + refName string + mediaType string + }{ + {name: "manifest", kind: KindManifests, refName: other.SHA256 + ".json", mediaType: "application/json"}, + {name: "segment", kind: KindSegments, refName: other.SHA256 + ".ndjson", mediaType: "application/x-ndjson"}, + { + name: "metadata", kind: KindMeta, + refName: "20260721T010203.000000000Z-0-" + other.SHA256 + ".json", mediaType: "application/json", + }, + {name: "raw", kind: KindRaw, refName: other.SHA256, mediaType: "application/octet-stream"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, tt.kind, tt.refName) + + _, err := store.Create(t.Context(), ref, identity, + tt.mediaType, bytes.NewReader(body)) + assert.ErrorIs(t, err, ErrArtifactInvalid) + _, err = store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) + }) + } + }) + }) + + t.Run("open verifies a streamed read", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte("streamed checkpoint content") + want := createContractArtifact(t, store, ref, body).Entry + + entry, reader, err := store.Open(t.Context(), ref) + require.NoError(t, err) + require.NotNil(t, reader) + assert.Equal(t, want, entry) + prefix := make([]byte, 8) + _, err = io.ReadFull(reader, prefix) + require.NoError(t, err) + assert.Equal(t, []byte("streamed"), prefix) + require.NoError(t, reader.Verify()) + assert.NoError(t, reader.Close()) + + entry, reader, err = store.Open(t.Context(), ref) + require.NoError(t, err) + assert.Equal(t, want, entry) + got, err := io.ReadAll(reader) + require.NoError(t, err) + assert.Equal(t, body, got) + assert.NoError(t, reader.Verify()) + assert.NoError(t, reader.Close()) + }) + + t.Run("early close does not drain or damage content", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte("content that must be explicitly verified") + createContractArtifact(t, store, ref, body) + + _, reader, err := store.Open(t.Context(), ref) + require.NoError(t, err) + one := make([]byte, 1) + _, err = reader.Read(one) + require.NoError(t, err) + assert.Error(t, reader.Close()) + assert.Equal(t, body, readContractArtifact(t, store, ref)) + }) + + t.Run("quarantine excludes content and permits recreation", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + createContractArtifact(t, store, ref, []byte("invalid current-format checkpoint")) + + require.NoError(t, store.Quarantine(t.Context(), ref, "semantic validation failed")) + _, err := store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) + assert.Empty(t, listAllContractEntries(t, store, contractOrigin, KindCheckpoints, 10)) + assert.Empty(t, listAllContractOrigins(t, store, 10)) + + replacement := []byte("trusted replacement checkpoint") + result := createContractArtifact(t, store, ref, replacement) + assert.True(t, result.Created) + assert.Equal(t, replacement, readContractArtifact(t, store, ref)) + }) + + t.Run("trash removes content from live reads and enumeration", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + createContractArtifact(t, store, ref, []byte("unreachable checkpoint")) + + require.NoError(t, store.Trash(t.Context(), ref)) + _, err := store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) + _, reader, err := store.Open(t.Context(), ref) + assert.Nil(t, reader) + assert.ErrorIs(t, err, ErrArtifactNotFound) + assert.Empty(t, listAllContractEntries(t, store, contractOrigin, KindCheckpoints, 10)) + assert.Empty(t, listAllContractOrigins(t, store, 10)) + }) + + t.Run("operations preserve cancellation", func(t *testing.T) { + store := newContractStore(t, factory) + existing := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + quarantineRef := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000002.json") + trashRef := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000003.json") + createContractArtifact(t, store, existing, []byte("existing")) + createContractArtifact(t, store, quarantineRef, []byte("quarantine")) + createContractArtifact(t, store, trashRef, []byte("trash")) + missing := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000004.json") + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err := store.Create(ctx, missing, identityForBytes(t, []byte("new")), + "application/json", strings.NewReader("new")) + assert.ErrorIs(t, err, context.Canceled) + _, err = store.Stat(ctx, existing) + assert.ErrorIs(t, err, context.Canceled) + _, reader, err := store.Open(ctx, existing) + if reader != nil { + _ = reader.Close() + } + assert.ErrorIs(t, err, context.Canceled) + _, err = store.Origins(ctx) + assert.ErrorIs(t, err, context.Canceled) + _, err = store.Entries(ctx, contractOrigin, KindCheckpoints) + assert.ErrorIs(t, err, context.Canceled) + assert.ErrorIs(t, store.Quarantine(ctx, quarantineRef, "cancelled"), context.Canceled) + assert.ErrorIs(t, store.Trash(ctx, trashRef), context.Canceled) + _, err = store.Stat(t.Context(), missing) + assert.ErrorIs(t, err, ErrArtifactNotFound) + _, err = store.Stat(t.Context(), quarantineRef) + assert.NoError(t, err) + _, err = store.Stat(t.Context(), trashRef) + assert.NoError(t, err) + }) + + t.Run("concurrent identical creates converge", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte("one immutable concurrent value") + identity := identityForBytes(t, body) + const writers = 8 + results := make([]CreateResult, writers) + errs := make([]error, writers) + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(writers) + for i := range writers { + go func() { + defer wg.Done() + <-start + results[i], errs[i] = store.Create(t.Context(), ref, identity, + "application/json", bytes.NewReader(body)) + }() + } + close(start) + wg.Wait() + + created := 0 + for i := range writers { + assert.NoError(t, errs[i]) + assert.Equal(t, ref, results[i].Entry.Ref) + assert.Equal(t, identity, results[i].Entry.Identity) + if results[i].Created { + created++ + } + } + assert.Equal(t, 1, created) + assert.Equal(t, body, readContractArtifact(t, store, ref)) + }) + + t.Run("concurrent distinct creates preserve one winner", func(t *testing.T) { + store := newContractStore(t, factory) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + bodies := [2][]byte{ + []byte("first immutable concurrent value"), + []byte("second immutable concurrent value"), + } + identities := [2]Identity{ + identityForBytes(t, bodies[0]), + identityForBytes(t, bodies[1]), + } + var results [2]CreateResult + var errs [2]error + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(len(bodies)) + for i := range bodies { + go func() { + defer wg.Done() + <-start + results[i], errs[i] = store.Create(t.Context(), ref, identities[i], + "application/json", bytes.NewReader(bodies[i])) + }() + } + close(start) + wg.Wait() + + winner := -1 + successes := 0 + conflicts := 0 + for i := range bodies { + if errs[i] == nil { + winner = i + successes++ + assert.True(t, results[i].Created) + assert.Equal(t, ref, results[i].Entry.Ref) + assert.Equal(t, identities[i], results[i].Entry.Identity) + continue + } + if assert.ErrorIs(t, errs[i], ErrArtifactConflict) { + conflicts++ + } + } + require.NotEqual(t, -1, winner) + assert.Equal(t, 1, successes) + assert.Equal(t, 1, conflicts) + assert.Equal(t, bodies[winner], readContractArtifact(t, store, ref)) + entry, err := store.Stat(t.Context(), ref) + require.NoError(t, err) + assert.Equal(t, identities[winner], entry.Identity) + }) +} + +func newContractStore(t *testing.T, factory artifactStoreFactory) ArtifactStore { + t.Helper() + store := factory(t) + require.NotNil(t, store) + t.Cleanup(func() { + closer, ok := any(store).(io.Closer) + require.True(t, ok) + require.NoError(t, closer.Close()) + }) + return store +} + +func requireContractRef(t *testing.T, origin string, kind Kind, name string) Ref { + t.Helper() + ref, err := NewRef(origin, kind, name) + require.NoError(t, err) + return ref +} + +func identityForBytes(t *testing.T, body []byte) Identity { + t.Helper() + sum := sha256.Sum256(body) + identity, err := NewIdentity(hex.EncodeToString(sum[:]), int64(len(body))) + require.NoError(t, err) + return identity +} + +func createContractArtifact(t *testing.T, store ArtifactStore, ref Ref, body []byte) CreateResult { + t.Helper() + result, err := store.Create(t.Context(), ref, identityForBytes(t, body), + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + return result +} + +func readContractArtifact(t *testing.T, store ArtifactStore, ref Ref) []byte { + t.Helper() + _, reader, err := store.Open(t.Context(), ref) + require.NoError(t, err) + require.NotNil(t, reader) + data, err := io.ReadAll(reader) + require.NoError(t, err) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + return data +} + +func listAllContractEntries( + t *testing.T, store ArtifactStore, origin string, kind Kind, limit int, +) []Entry { + t.Helper() + var entries []Entry + iterator, err := store.Entries(t.Context(), origin, kind) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, iterator.Close()) }) + for { + page, nextErr := iterator.Next(t.Context(), limit) + require.True(t, nextErr == nil || errors.Is(nextErr, io.EOF)) + assert.LessOrEqual(t, len(page), limit) + entries = append(entries, page...) + if errors.Is(nextErr, io.EOF) { + break + } + } + return entries +} + +func listAllContractOrigins(t *testing.T, store ArtifactStore, limit int) []string { + t.Helper() + var origins []string + iterator, err := store.Origins(t.Context()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, iterator.Close()) }) + for { + page, nextErr := iterator.Next(t.Context(), limit) + require.True(t, nextErr == nil || errors.Is(nextErr, io.EOF)) + assert.LessOrEqual(t, len(page), limit) + origins = append(origins, page...) + if errors.Is(nextErr, io.EOF) { + break + } + } + return origins +} + +func entryNames(entries []Entry) []string { + names := make([]string, 0, len(entries)) + for _, entry := range entries { + names = append(names, entry.Ref.Name) + } + return names +} diff --git a/internal/artifact/store_docbank.go b/internal/artifact/store_docbank.go new file mode 100644 index 000000000..76d9e780d --- /dev/null +++ b/internal/artifact/store_docbank.go @@ -0,0 +1,1031 @@ +package artifact + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "io/fs" + "strings" + "sync" + "time" + + "go.kenn.io/docbank" + "go.kenn.io/kit/pack" + "go.kenn.io/kit/packstore" +) + +const ( + docbankLiveRoot = "/v1" + docbankQuarantineRoot = "/.quarantine/v1" + docbankTrashCursorV1 = "agentsview-artifact-maintenance:v1:empty-trash" +) + +type docbankStore struct { + vault *docbank.Vault + quarantineMu sync.Mutex +} + +type docbankTraversal struct { + walker docbankWalker + page []docbank.WalkEntry + pageOffset int + lastOrigin string + pendingOrigin string + nextEntry *Entry + nextQuarantine *QuarantinedEntry +} + +func newDocbankContent(vault *docbank.Vault) *docbankStore { + store := &docbankStore{vault: vault} + return store +} + +func (s *docbankStore) Create( + ctx context.Context, + ref Ref, + identity Identity, + mediaType string, + body io.Reader, +) (CreateResult, error) { + if err := ctx.Err(); err != nil { + return CreateResult{}, artifactStoreError("create", ref, err) + } + if err := validateStoreRef(ref); err != nil { + return CreateResult{}, artifactStoreError("create", ref, err) + } + if err := validateStoreIdentity(identity); err != nil { + return CreateResult{}, artifactStoreError("create", ref, err) + } + if err := validateRefIdentity(ref, identity); err != nil { + return CreateResult{}, artifactStoreError("create", ref, err) + } + if body == nil { + return CreateResult{}, artifactStoreError("create", ref, + fmt.Errorf("%w: artifact body is required", ErrArtifactInvalid)) + } + if mediaType != canonicalArtifactMediaType(ref.Kind) { + if _, err := s.vault.Stat(ctx, docbankPath(ref)); err == nil { + return CreateResult{}, artifactStoreError("create", ref, + fmt.Errorf("%w: media type does not match existing artifact", ErrArtifactConflict)) + } else if !errors.Is(err, docbank.ErrNotFound) { + return CreateResult{}, artifactStoreError("create", ref, mapDocbankError(err)) + } + return CreateResult{}, artifactStoreError("create", ref, + fmt.Errorf("%w: unsupported media type %q", ErrArtifactInvalid, mediaType)) + } + receipt, err := s.vault.Create(ctx, docbankPath(ref), body, docbank.CreateOptions{ + MediaType: mediaType, + Expected: docbank.ContentIdentity{ + SHA256: identity.SHA256, + Size: identity.Size, + }, + }) + if err != nil { + return CreateResult{}, artifactStoreError("create", ref, mapDocbankError(err)) + } + entry, err := docbankEntry(ref, receipt.Node) + if err != nil { + return CreateResult{}, artifactStoreError("create", ref, err) + } + result := CreateResult{ + Entry: entry, + Created: receipt.Created, + } + if receipt.PhysicalCreated { + result.Physical = PhysicalWrite{ + Kind: receipt.Physical.Kind, + Encoding: receipt.Physical.Encoding, + LogicalBytes: receipt.Physical.LogicalBytes, + StoredBytes: receipt.Physical.StoredBytes, + PackEligible: receipt.Physical.PackEligible, + } + } + return result, nil +} + +func (s *docbankStore) Stat(ctx context.Context, ref Ref) (Entry, error) { + if err := ctx.Err(); err != nil { + return Entry{}, artifactStoreError("stat", ref, err) + } + if err := validateStoreRef(ref); err != nil { + return Entry{}, artifactStoreError("stat", ref, err) + } + node, err := s.vault.Stat(ctx, docbankPath(ref)) + if err != nil { + return Entry{}, artifactStoreError("stat", ref, mapDocbankError(err)) + } + entry, err := docbankEntry(ref, node) + if err != nil { + return Entry{}, artifactStoreError("stat", ref, err) + } + return entry, nil +} + +func (s *docbankStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + if err := ctx.Err(); err != nil { + return Entry{}, nil, artifactStoreError("open", ref, err) + } + if err := validateStoreRef(ref); err != nil { + return Entry{}, nil, artifactStoreError("open", ref, err) + } + content, err := s.vault.OpenContent(ctx, docbankPath(ref)) + if err != nil { + return Entry{}, nil, artifactStoreError("open", ref, mapDocbankError(err)) + } + entry, err := docbankEntry(ref, content.Node) + if err != nil { + closeErr := content.Reader.Close() + return Entry{}, nil, artifactStoreError("open", ref, errors.Join(err, closeErr)) + } + return entry, &docbankVerifiedReader{reader: content.Reader}, nil +} + +type docbankIterator struct { + mu sync.Mutex + + vault *docbank.Vault + root string + state docbankTraversal + + opened bool + done bool + closed bool +} + +type docbankOriginIterator struct { + iterator docbankIterator +} + +type docbankEntryIterator struct { + iterator docbankIterator + origin string + kind Kind +} + +type docbankQuarantineIterator struct { + iterator docbankIterator +} + +func (s *docbankStore) Origins(ctx context.Context) (OriginIterator, error) { + if err := ctx.Err(); err != nil { + return nil, artifactStoreError("iterate origins", Ref{}, err) + } + return &docbankOriginIterator{iterator: docbankIterator{ + vault: s.vault, + root: docbankLiveRoot, + }}, nil +} + +func (s *docbankStore) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + ref := Ref{Origin: origin, Kind: kind} + if err := ctx.Err(); err != nil { + return nil, artifactStoreError("iterate entries", ref, err) + } + if err := validateStoreCollection(origin, kind); err != nil { + return nil, artifactStoreError("iterate entries", ref, err) + } + return &docbankEntryIterator{ + iterator: docbankIterator{ + vault: s.vault, + root: docbankLiveRoot + "/" + origin + "/" + string(kind), + }, + origin: origin, + kind: kind, + }, nil +} + +func (s *docbankStore) Quarantined(ctx context.Context) (QuarantineIterator, error) { + if err := ctx.Err(); err != nil { + return nil, artifactStoreError("iterate quarantine", Ref{}, err) + } + return &docbankQuarantineIterator{iterator: docbankIterator{ + vault: s.vault, + root: docbankQuarantineRoot, + }}, nil +} + +func (i *docbankOriginIterator) Next( + ctx context.Context, limit int, +) ([]string, error) { + i.iterator.mu.Lock() + defer i.iterator.mu.Unlock() + if err := i.iterator.beforeNext(ctx, limit); err != nil { + return nil, artifactStoreError("iterate origins", Ref{}, err) + } + items, more, err := i.iterator.state.nextOrigins(ctx, limit) + if err != nil { + return nil, artifactStoreError( + "iterate origins", Ref{}, i.iterator.fail(err), + ) + } + if more { + return items, nil + } + closeErr := i.iterator.finish() + return items, artifactStoreError("iterate origins", Ref{}, errors.Join(io.EOF, closeErr)) +} + +func (i *docbankOriginIterator) Close() error { + i.iterator.mu.Lock() + defer i.iterator.mu.Unlock() + return artifactStoreErrorIfError("close origin iterator", Ref{}, i.iterator.close()) +} + +func (i *docbankEntryIterator) Next( + ctx context.Context, limit int, +) ([]Entry, error) { + i.iterator.mu.Lock() + defer i.iterator.mu.Unlock() + ref := Ref{Origin: i.origin, Kind: i.kind} + if err := i.iterator.beforeNext(ctx, limit); err != nil { + return nil, artifactStoreError("iterate entries", ref, err) + } + items, more, err := i.iterator.state.nextEntries(ctx, i.origin, i.kind, limit) + if err != nil { + return nil, artifactStoreError( + "iterate entries", ref, i.iterator.fail(err), + ) + } + if more { + return items, nil + } + closeErr := i.iterator.finish() + return items, artifactStoreError("iterate entries", ref, errors.Join(io.EOF, closeErr)) +} + +func (i *docbankEntryIterator) Close() error { + i.iterator.mu.Lock() + defer i.iterator.mu.Unlock() + ref := Ref{Origin: i.origin, Kind: i.kind} + return artifactStoreErrorIfError("close entry iterator", ref, i.iterator.close()) +} + +func (i *docbankQuarantineIterator) Next( + ctx context.Context, limit int, +) ([]QuarantinedEntry, error) { + i.iterator.mu.Lock() + defer i.iterator.mu.Unlock() + if err := i.iterator.beforeNext(ctx, limit); err != nil { + return nil, artifactStoreError("iterate quarantine", Ref{}, err) + } + items, more, err := i.iterator.state.nextQuarantinedEntries(ctx, limit) + if err != nil { + return nil, artifactStoreError( + "iterate quarantine", Ref{}, i.iterator.fail(err), + ) + } + if more { + return items, nil + } + closeErr := i.iterator.finish() + return items, artifactStoreError("iterate quarantine", Ref{}, errors.Join(io.EOF, closeErr)) +} + +func (i *docbankQuarantineIterator) Close() error { + i.iterator.mu.Lock() + defer i.iterator.mu.Unlock() + return artifactStoreErrorIfError("close quarantine iterator", Ref{}, i.iterator.close()) +} + +func (i *docbankIterator) beforeNext(ctx context.Context, limit int) error { + if i.closed { + return fs.ErrClosed + } + if err := ctx.Err(); err != nil { + return i.fail(err) + } + if limit <= 0 || limit > maxArtifactListPageSize { + return fmt.Errorf( + "%w: page limit must be between 1 and %d", + ErrArtifactInvalid, + maxArtifactListPageSize, + ) + } + if i.done { + return io.EOF + } + if i.opened { + return nil + } + i.opened = true + walker, err := i.vault.Walk(ctx, i.root, docbank.WalkOptions{ + PageSize: min(limit, docbank.MaxWalkPageSize), + }) + if errors.Is(err, docbank.ErrNotFound) { + i.done = true + return io.EOF + } + if err != nil { + i.closed = true + return mapDocbankError(err) + } + i.state.walker = walker + return nil +} + +func (i *docbankIterator) fail(err error) error { + i.closed = true + return errors.Join(mapDocbankError(err), i.closeWalker()) +} + +func (i *docbankIterator) finish() error { + i.done = true + return i.closeWalker() +} + +func (i *docbankIterator) close() error { + if i.closed { + return nil + } + i.closed = true + return i.closeWalker() +} + +func (i *docbankIterator) closeWalker() error { + if i.state.walker == nil { + return nil + } + err := i.state.walker.Close() + i.state.walker = nil + return mapDocbankError(err) +} + +func artifactStoreErrorIfError(op string, ref Ref, err error) error { + if err == nil { + return nil + } + return artifactStoreError(op, ref, err) +} + +func (s *docbankStore) TrashQuarantined(ctx context.Context, token string) error { + if err := ctx.Err(); err != nil { + return artifactStoreError("trash quarantine", Ref{}, err) + } + if _, valid := quarantinedRefFromDocbankPath(token); !valid { + return artifactStoreError("trash quarantine", Ref{}, + fmt.Errorf("%w: invalid quarantine token", ErrArtifactInvalid)) + } + node, err := s.vault.Stat(ctx, token) + if err != nil { + return artifactStoreError("trash quarantine", Ref{}, mapDocbankError(err)) + } + _, err = s.vault.TrashPath(ctx, token, docbank.RevisionOptions{ + IfRevision: node.Revision, + }) + if err != nil { + return artifactStoreError("trash quarantine", Ref{}, mapDocbankError(err)) + } + return nil +} + +func (s *docbankStore) Quarantine(ctx context.Context, ref Ref, _ string) error { + if err := ctx.Err(); err != nil { + return artifactStoreError("quarantine", ref, err) + } + if err := validateStoreRef(ref); err != nil { + return artifactStoreError("quarantine", ref, err) + } + s.quarantineMu.Lock() + defer s.quarantineMu.Unlock() + parent := docbankQuarantineParent(ref) + if err := s.ensureQuarantineParent(ctx, parent); err != nil { + return artifactStoreError("quarantine", ref, mapDocbankError(err)) + } + for range 4 { + id, err := newDocbankQuarantineID() + if err != nil { + return artifactStoreError("quarantine", ref, err) + } + destination := parent + "/" + id + "-" + ref.Name + _, err = s.vault.MovePath(ctx, docbankPath(ref), destination, docbank.RevisionOptions{}) + if err == nil { + return nil + } + if !errors.Is(err, docbank.ErrExists) { + return artifactStoreError("quarantine", ref, mapDocbankError(err)) + } + } + return artifactStoreError("quarantine", ref, + fmt.Errorf("%w: quarantine identifier collisions", ErrArtifactConflict)) +} + +func (s *docbankStore) Trash(ctx context.Context, ref Ref) error { + if err := ctx.Err(); err != nil { + return artifactStoreError("trash", ref, err) + } + if err := validateStoreRef(ref); err != nil { + return artifactStoreError("trash", ref, err) + } + _, err := s.vault.TrashPath(ctx, docbankPath(ref), docbank.RevisionOptions{}) + if err != nil { + return artifactStoreError("trash", ref, mapDocbankError(err)) + } + return nil +} + +func (s *docbankStore) Pack(ctx context.Context, maxBytes int64) (PackResult, error) { + if err := ctx.Err(); err != nil { + return PackResult{}, artifactStoreError("pack", Ref{}, err) + } + if maxBytes < 0 { + return PackResult{}, artifactStoreError("pack", Ref{}, + fmt.Errorf("%w: pack byte limit must not be negative", ErrArtifactInvalid)) + } + report, err := s.vault.Pack(ctx, docbank.PackOptions{MaxBytes: maxBytes}) + if err != nil { + return PackResult{}, artifactStoreError("pack", Ref{}, mapDocbankError(err)) + } + return PackResult{ + PackedObjects: report.BlobsPacked, + LogicalBytes: report.BytesPacked, + More: report.More, + }, nil +} + +func (s *docbankStore) LooseBacklog(ctx context.Context) (LooseBacklog, error) { + if err := ctx.Err(); err != nil { + return LooseBacklog{}, artifactStoreError("loose backlog", Ref{}, err) + } + backlog, err := s.vault.LooseBacklog(ctx) + if err != nil { + return LooseBacklog{}, artifactStoreError("loose backlog", Ref{}, mapDocbankError(err)) + } + return LooseBacklog{ + EligibleObjects: backlog.EligibleObjects, + EligibleBytes: backlog.EligibleBytes, + EligibleStoredBytes: backlog.EligibleStoredBytes, + }, nil +} + +func (s *docbankStore) Verify( + ctx context.Context, budget WorkBudget, +) (MaintenanceResult, error) { + if err := validateArtifactWorkBudget(budget); err != nil { + return MaintenanceResult{}, artifactStoreError("verify", Ref{}, err) + } + report, err := s.vault.Verify(ctx, docbank.VerifyOptions{ + Budget: docbankWorkBudget(budget), + }) + if err != nil { + return MaintenanceResult{}, artifactStoreError("verify", Ref{}, mapDocbankError(err)) + } + return MaintenanceResult{ + Processed: report.OK + len(report.Problems), + NextCursor: report.NextCursor, + More: report.More, + }, nil +} + +func (s *docbankStore) EmptyTrash( + ctx context.Context, olderThan time.Duration, budget WorkBudget, +) (MaintenanceResult, error) { + if err := validateArtifactWorkBudget(budget); err != nil { + return MaintenanceResult{}, artifactStoreError("empty trash", Ref{}, err) + } + if olderThan < 0 { + return MaintenanceResult{}, artifactStoreError("empty trash", Ref{}, + fmt.Errorf("%w: trash grace must not be negative", ErrArtifactInvalid)) + } + if budget.MaxBytes != 0 { + return MaintenanceResult{}, artifactStoreError("empty trash", Ref{}, + fmt.Errorf("%w: trash emptying supports only an object budget", ErrArtifactInvalid)) + } + if budget.Cursor != "" && budget.Cursor != docbankTrashCursorV1 { + return MaintenanceResult{}, artifactStoreError("empty trash", Ref{}, + fmt.Errorf("%w: invalid trash continuation cursor", ErrArtifactInvalid)) + } + report, err := s.vault.EmptyTrash(ctx, docbank.TrashEmptyOptions{ + OlderThan: olderThan, + MaxRoots: budget.MaxObjects, + }) + if err != nil { + return MaintenanceResult{}, artifactStoreError("empty trash", Ref{}, mapDocbankError(err)) + } + result := MaintenanceResult{ + Processed: int(report.Deleted), + More: report.More, + } + if report.More { + result.NextCursor = docbankTrashCursorV1 + } + return result, nil +} + +func (s *docbankStore) GarbageCollect( + ctx context.Context, budget WorkBudget, +) (MaintenanceResult, error) { + if err := validateArtifactWorkBudget(budget); err != nil { + return MaintenanceResult{}, artifactStoreError("garbage collect", Ref{}, err) + } + report, err := s.vault.GarbageCollect(ctx, docbank.GCOptions{ + Budget: docbankWorkBudget(budget), + }) + if err != nil { + return MaintenanceResult{}, artifactStoreError( + "garbage collect", Ref{}, mapDocbankError(err)) + } + return MaintenanceResult{ + Processed: report.CandidateBlobs + report.UntrackedFiles, + Bytes: report.ReclaimableBytes, + NextCursor: report.NextCursor, + More: report.More, + }, nil +} + +func (s *docbankStore) Repack( + ctx context.Context, budget WorkBudget, +) (MaintenanceResult, error) { + if err := validateArtifactWorkBudget(budget); err != nil { + return MaintenanceResult{}, artifactStoreError("repack", Ref{}, err) + } + report, err := s.vault.Repack(ctx, docbank.RepackOptions{ + Budget: docbankWorkBudget(budget), + }) + if err != nil { + return MaintenanceResult{}, artifactStoreError("repack", Ref{}, mapDocbankError(err)) + } + return MaintenanceResult{ + Processed: docbankRepackProcessed(report), + Bytes: report.BytesRepacked, + NextCursor: report.NextCursor, + More: report.More, + }, nil +} + +func docbankWorkBudget(budget WorkBudget) docbank.WorkBudget { + return docbank.WorkBudget{ + MaxObjects: budget.MaxObjects, + MaxBytes: budget.MaxBytes, + Cursor: budget.Cursor, + } +} + +func validateArtifactWorkBudget(budget WorkBudget) error { + if budget.MaxObjects < 0 || budget.MaxObjects > docbank.MaxMaintenanceObjects || + budget.MaxBytes < 0 { + return fmt.Errorf("%w: maintenance budget is outside the supported range", ErrArtifactInvalid) + } + return nil +} + +func docbankRepackProcessed(report docbank.RepackReport) int { + total := report.MappingsPruned + + int64(report.PacksSelected) + + int64(report.PacksRewritten) + + int64(report.PacksSealed) + + int64(report.PacksRemoved) + + int64(report.PacksDeferredOversized) + + int64(report.BlobsRepacked) + maxInt := int64(^uint(0) >> 1) + if total > maxInt { + return int(maxInt) + } + return int(total) +} + +// RepairContent replaces corrupt physical bytes for an existing canonical +// identity without changing any logical artifact reference. +func (s *docbankStore) RepairContent( + ctx context.Context, identity Identity, trusted io.Reader, +) error { + if err := validateStoreIdentity(identity); err != nil { + return artifactStoreError("repair content", Ref{}, err) + } + if trusted == nil { + return artifactStoreError("repair content", Ref{}, + fmt.Errorf("%w: trusted repair body is required", ErrArtifactInvalid)) + } + _, err := s.vault.RepairContent(ctx, docbank.ContentIdentity{ + SHA256: identity.SHA256, + Size: identity.Size, + }, trusted) + if err != nil { + return artifactStoreError("repair content", Ref{}, mapDocbankError(err)) + } + return nil +} + +func (s *docbankStore) Close() error { + if s == nil || s.vault == nil { + return nil + } + err := s.vault.Close() + s.vault = nil + return err +} + +// checkpointFloor traverses both Docbank namespaces through stable walkers, +// page-by-page, without materializing either checkpoint collection. +func (s *docbankStore) checkpointFloor( + ctx context.Context, origin string, +) (int, error) { + liveCollection := docbankLiveRoot + "/" + origin + "/" + string(KindCheckpoints) + live, err := s.walkCheckpointFloor(ctx, liveCollection, false) + if err != nil { + return 0, err + } + quarantineCollection := docbankQuarantineRoot + "/" + origin + "/" + string(KindCheckpoints) + quarantined, err := s.walkCheckpointFloor(ctx, quarantineCollection, true) + if err != nil { + return 0, err + } + return max(live, quarantined), nil +} + +func (s *docbankStore) walkCheckpointFloor( + ctx context.Context, collection string, quarantined bool, +) (_ int, retErr error) { + return walkCheckpointFloor(ctx, collection, quarantined, func( + ctx context.Context, + ) (docbankWalker, error) { + return s.vault.Walk(ctx, collection, docbank.WalkOptions{ + PageSize: checkpointFloorPageSize, + }) + }) +} + +func walkCheckpointFloor( + ctx context.Context, + collection string, + quarantined bool, + open func(context.Context) (docbankWalker, error), +) (_ int, retErr error) { + if err := ctx.Err(); err != nil { + return 0, err + } + walker, err := open(ctx) + if errors.Is(err, docbank.ErrNotFound) { + return 0, nil + } + if err != nil { + return 0, mapDocbankError(err) + } + defer func() { retErr = errors.Join(retErr, walker.Close()) }() + prefix := collection + "/" + floor := 0 + for { + if err := ctx.Err(); err != nil { + return 0, err + } + page, nextErr := walker.Next(ctx) + if err := ctx.Err(); err != nil { + return 0, err + } + if errors.Is(nextErr, io.EOF) { + return floor, nil + } + if nextErr != nil { + return 0, mapDocbankError(nextErr) + } + for _, item := range page { + if item.Node.BlobHash == "" { + continue + } + name := strings.TrimPrefix(item.Path, prefix) + if name == item.Path || strings.Contains(name, "/") { + continue + } + if quarantined { + if len(name) <= 33 || name[32] != '-' { + continue + } + quarantineID, err := hex.DecodeString(name[:32]) + if err != nil || len(quarantineID) != 16 { + continue + } + name = name[33:] + } + sequence, err := checkpointSequence(name) + if err == nil { + floor = max(floor, sequence) + } + } + } +} + +type docbankWalker interface { + Next(context.Context) ([]docbank.WalkEntry, error) + Close() error +} + +func collectDocbankWalk( + ctx context.Context, + open func(context.Context) (docbankWalker, error), +) (entries []docbank.WalkEntry, retErr error) { + walker, err := open(ctx) + if errors.Is(err, docbank.ErrNotFound) { + return []docbank.WalkEntry{}, nil + } + if err != nil { + return nil, err + } + defer func() { retErr = errors.Join(retErr, walker.Close()) }() + for { + page, nextErr := walker.Next(ctx) + if errors.Is(nextErr, io.EOF) { + return entries, nil + } + if nextErr != nil { + return nil, nextErr + } + entries = append(entries, page...) + } +} + +func (s *docbankStore) ensureQuarantineParent(ctx context.Context, parent string) error { + if _, err := s.vault.Stat(ctx, parent); err == nil { + return nil + } else if !errors.Is(err, docbank.ErrNotFound) { + return err + } + emptyHash := sha256.Sum256(nil) + anchor := parent + "/.anchor" + receipt, err := s.vault.Create(ctx, anchor, strings.NewReader(""), docbank.CreateOptions{ + MediaType: "application/octet-stream", + Expected: docbank.ContentIdentity{ + SHA256: hex.EncodeToString(emptyHash[:]), + Size: 0, + }, + }) + if err != nil { + return err + } + _, err = s.vault.TrashPath(ctx, anchor, docbank.RevisionOptions{ + IfRevision: receipt.Node.Revision, + }) + return err +} + +func (s *docbankTraversal) nextWalk(ctx context.Context) (docbank.WalkEntry, bool, error) { + for { + if s.pageOffset < len(s.page) { + item := s.page[s.pageOffset] + s.pageOffset++ + return item, true, nil + } + if s.walker == nil { + return docbank.WalkEntry{}, false, nil + } + page, err := s.walker.Next(ctx) + if errors.Is(err, io.EOF) { + return docbank.WalkEntry{}, false, nil + } + if err != nil { + return docbank.WalkEntry{}, false, err + } + s.page = page + s.pageOffset = 0 + } +} + +func (s *docbankTraversal) nextOrigin(ctx context.Context) (string, bool, error) { + if s.pendingOrigin != "" { + origin := s.pendingOrigin + s.pendingOrigin = "" + return origin, true, nil + } + for { + item, ok, err := s.nextWalk(ctx) + if err != nil || !ok { + return "", ok, err + } + ref, valid := refFromDocbankPath(item.Path) + if !valid || item.Node.BlobHash == "" || ref.Origin == s.lastOrigin { + continue + } + s.lastOrigin = ref.Origin + return ref.Origin, true, nil + } +} + +func (s *docbankTraversal) nextOrigins( + ctx context.Context, limit int, +) ([]string, bool, error) { + items := make([]string, 0, limit) + for len(items) < limit { + origin, ok, err := s.nextOrigin(ctx) + if err != nil || !ok { + return items, false, err + } + items = append(items, origin) + } + next, ok, err := s.nextOrigin(ctx) + if err != nil || !ok { + return items, false, err + } + s.pendingOrigin = next + return items, true, nil +} + +func (s *docbankTraversal) nextLogicalEntry( + ctx context.Context, origin string, kind Kind, +) (Entry, bool, error) { + if s.nextEntry != nil { + entry := *s.nextEntry + s.nextEntry = nil + return entry, true, nil + } + for { + item, ok, err := s.nextWalk(ctx) + if err != nil || !ok { + return Entry{}, ok, err + } + ref, valid := refFromDocbankPath(item.Path) + if !valid || item.Node.BlobHash == "" || ref.Origin != origin || ref.Kind != kind { + continue + } + entry, err := docbankEntry(ref, item.Node) + return entry, true, err + } +} + +func (s *docbankTraversal) nextEntries( + ctx context.Context, origin string, kind Kind, limit int, +) ([]Entry, bool, error) { + items := make([]Entry, 0, limit) + for len(items) < limit { + entry, ok, err := s.nextLogicalEntry(ctx, origin, kind) + if err != nil || !ok { + return items, false, err + } + items = append(items, entry) + } + next, ok, err := s.nextLogicalEntry(ctx, origin, kind) + if err != nil || !ok { + return items, false, err + } + s.nextEntry = &next + return items, true, nil +} + +func (s *docbankTraversal) nextQuarantinedEntry( + ctx context.Context, +) (QuarantinedEntry, bool, error) { + if s.nextQuarantine != nil { + entry := *s.nextQuarantine + s.nextQuarantine = nil + return entry, true, nil + } + for { + item, ok, err := s.nextWalk(ctx) + if err != nil || !ok { + return QuarantinedEntry{}, ok, err + } + ref, valid := quarantinedRefFromDocbankPath(item.Path) + if !valid || item.Node.BlobHash == "" { + continue + } + entry, err := docbankEntry(ref, item.Node) + if err != nil { + return QuarantinedEntry{}, false, err + } + return QuarantinedEntry{ + Token: item.Path, Ref: ref, Identity: entry.Identity, Modified: entry.Modified, + }, true, nil + } +} + +func (s *docbankTraversal) nextQuarantinedEntries( + ctx context.Context, limit int, +) ([]QuarantinedEntry, bool, error) { + items := make([]QuarantinedEntry, 0, limit) + for len(items) < limit { + entry, ok, err := s.nextQuarantinedEntry(ctx) + if err != nil || !ok { + return items, false, err + } + items = append(items, entry) + } + next, ok, err := s.nextQuarantinedEntry(ctx) + if err != nil || !ok { + return items, false, err + } + s.nextQuarantine = &next + return items, true, nil +} + +func docbankPath(ref Ref) string { + return docbankLiveRoot + "/" + ref.Origin + "/" + string(ref.Kind) + "/" + ref.Name +} + +func docbankQuarantineParent(ref Ref) string { + return docbankQuarantineRoot + "/" + ref.Origin + "/" + string(ref.Kind) +} + +func refFromDocbankPath(value string) (Ref, bool) { + if !strings.HasPrefix(value, docbankLiveRoot+"/") { + return Ref{}, false + } + parts := strings.Split(strings.TrimPrefix(value, docbankLiveRoot+"/"), "/") + if len(parts) != 3 || value != docbankLiveRoot+"/"+strings.Join(parts, "/") { + return Ref{}, false + } + ref, err := NewRef(parts[0], Kind(parts[1]), parts[2]) + return ref, err == nil +} + +func quarantinedRefFromDocbankPath(value string) (Ref, bool) { + if !strings.HasPrefix(value, docbankQuarantineRoot+"/") { + return Ref{}, false + } + parts := strings.Split(strings.TrimPrefix(value, docbankQuarantineRoot+"/"), "/") + if len(parts) != 3 || value != docbankQuarantineRoot+"/"+strings.Join(parts, "/") { + return Ref{}, false + } + quarantinedName := parts[2] + if len(quarantinedName) <= 33 || quarantinedName[32] != '-' { + return Ref{}, false + } + id, err := hex.DecodeString(quarantinedName[:32]) + if err != nil || len(id) != 16 { + return Ref{}, false + } + ref, err := NewRef(parts[0], Kind(parts[1]), quarantinedName[33:]) + return ref, err == nil +} + +func docbankEntry(ref Ref, node docbank.Node) (Entry, error) { + identity, err := NewIdentity(node.BlobHash, node.Size) + if err != nil { + return Entry{}, fmt.Errorf("%w: invalid Docbank node identity: %w", ErrArtifactCorrupt, err) + } + if err := validateRefIdentity(ref, identity); err != nil { + return Entry{}, fmt.Errorf("%w: reference and Docbank blob identity differ: %w", + ErrArtifactCorrupt, err) + } + modified, err := time.Parse(time.RFC3339Nano, node.ModifiedAt) + if err != nil { + return Entry{}, fmt.Errorf("%w: invalid Docbank modification time: %w", ErrArtifactCorrupt, err) + } + return Entry{Ref: ref, Identity: identity, Modified: modified}, nil +} + +func newDocbankQuarantineID() (string, error) { + var raw [16]byte + if _, err := rand.Read(raw[:]); err != nil { + return "", err + } + return hex.EncodeToString(raw[:]), nil +} + +func mapDocbankError(err error) error { + switch { + case err == nil: + return nil + case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded): + return err + case errors.Is(err, docbank.ErrNotFound): + return fmt.Errorf("%w: %w", ErrArtifactNotFound, err) + case errors.Is(err, docbank.ErrContentConflict), errors.Is(err, docbank.ErrExists): + return fmt.Errorf("%w: %w", ErrArtifactConflict, err) + case errors.Is(err, docbank.ErrStaleRevision): + return fmt.Errorf("%w: %w", ErrArtifactConflict, err) + case errors.Is(err, docbank.ErrDigestMismatch), errors.Is(err, docbank.ErrSizeMismatch), + errors.Is(err, docbank.ErrInvalidMaintenanceCursor): + return fmt.Errorf("%w: %w", ErrArtifactInvalid, err) + case errors.Is(err, docbank.ErrContentUnavailable), errors.Is(err, packstore.ErrContentMismatch): + return fmt.Errorf("%w: %w", ErrArtifactCorrupt, err) + default: + return err + } +} + +type docbankVerifiedReader struct { + reader docbank.VerifiedReadCloser +} + +func (r *docbankVerifiedReader) Read(p []byte) (int, error) { + n, err := r.reader.Read(p) + return n, mapDocbankReadError(err) +} + +func (r *docbankVerifiedReader) Verify() error { + return mapDocbankReadError(r.reader.Verify()) +} + +func (r *docbankVerifiedReader) Close() error { + return mapDocbankReadError(r.reader.Close()) +} + +func mapDocbankReadError(err error) error { + switch { + case err == nil, errors.Is(err, io.EOF): + return err + case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded): + return err + case errors.Is(err, pack.ErrVerificationIncomplete): + return fmt.Errorf("%w: %w", errIncompleteArtifact, err) + default: + return fmt.Errorf("%w: %w", ErrArtifactCorrupt, err) + } +} + +var _ ArtifactStore = (*docbankStore)(nil) diff --git a/internal/artifact/store_docbank_test.go b/internal/artifact/store_docbank_test.go new file mode 100644 index 000000000..99c8f72d0 --- /dev/null +++ b/internal/artifact/store_docbank_test.go @@ -0,0 +1,485 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/sha256" + "errors" + "fmt" + "io" + "io/fs" + "regexp" + "strconv" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/docbank" +) + +func TestDocbankStoreUsesCanonicalNamespaceAndMovesQuarantineNode(t *testing.T) { + vault, store := newTestDocbankStore(t, docbank.Config{}) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000042.json") + body := []byte(`{"origin":"contract-a1b2c3","sequence":42}`) + created := createContractArtifact(t, store, ref, body) + + livePath := "/v1/contract-a1b2c3/checkpoints/cp-0000000042.json" + live, err := vault.Stat(t.Context(), livePath) + require.NoError(t, err) + assert.Equal(t, created.Entry.Identity.SHA256, live.BlobHash) + + require.NoError(t, store.Quarantine(t.Context(), ref, "semantic validation failed")) + _, err = vault.Stat(t.Context(), livePath) + assert.ErrorIs(t, err, docbank.ErrNotFound) + + quarantined := walkDocbankTestEntries(t, vault, docbankQuarantineRoot) + require.Len(t, quarantined, 1) + assert.Equal(t, live.ID, quarantined[0].Node.ID, "quarantine must move the stable node") + assert.Equal(t, live.BlobHash, quarantined[0].Node.BlobHash, "quarantine must not copy content") + assert.Regexp(t, + regexp.MustCompile(`^/\.quarantine/v1/contract-a1b2c3/checkpoints/[0-9a-f]{32}-cp-0000000042\.json$`), + quarantined[0].Path, + ) + + replacement := []byte("trusted replacement") + recreated := createContractArtifact(t, store, ref, replacement) + assert.True(t, recreated.Created) + assert.NotEqual(t, live.ID, mustDocbankNode(t, vault, livePath).ID) + assert.Equal(t, replacement, readContractArtifact(t, store, ref)) +} + +func TestDocbankStoreRejectsReferenceAndNodeIdentityMismatchBeforeRead(t *testing.T) { + vault, store := newTestDocbankStore(t, docbank.Config{}) + body := []byte("catalog-authorized but incorrectly named content") + identity := identityForBytes(t, body) + wrongHash := strings.Repeat("a", 64) + require.NotEqual(t, identity.SHA256, wrongHash) + ref := requireContractRef(t, contractOrigin, KindRaw, wrongHash) + _, err := vault.Create(t.Context(), docbankPath(ref), bytes.NewReader(body), docbank.CreateOptions{ + MediaType: "application/octet-stream", + Expected: docbank.ContentIdentity{SHA256: identity.SHA256, Size: identity.Size}, + }) + require.NoError(t, err) + + _, err = store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactCorrupt) + _, reader, err := store.Open(t.Context(), ref) + assert.Nil(t, reader) + assert.ErrorIs(t, err, ErrArtifactCorrupt) +} + +func TestDocbankStorePreservesTypedDocbankCauses(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + + _, err := store.Stat(t.Context(), ref) + assert.ErrorIs(t, err, ErrArtifactNotFound) + assert.ErrorIs(t, err, docbank.ErrNotFound) + + body := []byte("actual bytes") + expected := identityForBytes(t, bytes.Repeat([]byte("x"), len(body))) + require.Equal(t, int64(len(body)), expected.Size) + _, err = store.Create(t.Context(), ref, expected, "application/json", bytes.NewReader(body)) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.ErrorIs(t, err, docbank.ErrDigestMismatch) + + created := createContractArtifact(t, store, ref, body) + _, err = store.Create(t.Context(), ref, created.Entry.Identity, + "application/octet-stream", bytes.NewReader(body)) + assert.ErrorIs(t, err, ErrArtifactConflict) +} + +func TestDocbankStoreIdempotentRetryDoesNotReportPhysicalWrite(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + body := []byte(`{"origin":"contract-a1b2c3","sequence":1}`) + identity := identityForBytes(t, body) + + first, err := store.Create(t.Context(), ref, identity, "application/json", bytes.NewReader(body)) + require.NoError(t, err) + assert.True(t, first.Created) + assert.NotEqual(t, PhysicalWrite{}, first.Physical) + + retry, err := store.Create(t.Context(), ref, identity, "application/json", bytes.NewReader(body)) + require.NoError(t, err) + assert.False(t, retry.Created) + assert.Equal(t, first.Entry, retry.Entry) + assert.Equal(t, PhysicalWrite{}, retry.Physical) +} + +func TestDocbankStoreDistinctReferencesCountOnePhysicalWrite(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + + body := []byte(`{"origin":"contract-a1b2c3","shared":"physical-content"}`) + first := createCheckpointBody(t, store, 1, body) + second := createCheckpointBody(t, store, 2, body) + assert.True(t, first.Created) + assert.True(t, second.Created, "the second logical reference is new") + assert.NotEqual(t, PhysicalWrite{}, first.Physical) + assert.Equal(t, PhysicalWrite{}, second.Physical, + "the second logical reference must not claim a duplicate physical publication") +} + +func TestDocbankStoreReportsConfiguredLooseCompressionAndBacklog(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{LooseCompression: docbank.LooseCompressionOptions{ + Enabled: true, + MinBytes: 4 << 10, + MinSavingsPercent: 10, + }}) + + compressible := bytes.Repeat([]byte("canonical manifest payload\n"), 200) + compressible = compressible[:4<<10] + compressed := createCheckpointBody(t, store, 1, compressible) + assert.Equal(t, "loose", compressed.Physical.Kind) + assert.Equal(t, "zstd", compressed.Physical.Encoding) + assert.Equal(t, int64(len(compressible)), compressed.Physical.LogicalBytes) + assert.Less(t, compressed.Physical.StoredBytes, int64(len(compressible))*9/10) + + belowThreshold := bytes.Repeat([]byte("x"), (4<<10)-1) + rawSmall := createCheckpointBody(t, store, 2, belowThreshold) + assert.Equal(t, "raw", rawSmall.Physical.Encoding) + + incompressible := deterministicDocbankBytes(4 << 10) + rawSavings := createCheckpointBody(t, store, 3, incompressible) + assert.Equal(t, "raw", rawSavings.Physical.Encoding) + + assertCanonical := func(result CreateResult, want []byte) { + t.Helper() + entry, reader, err := store.Open(t.Context(), result.Entry.Ref) + require.NoError(t, err) + got, err := io.ReadAll(reader) + require.NoError(t, err) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + assert.Equal(t, want, got) + assert.Equal(t, identityForBytes(t, want), entry.Identity, + "physical encoding must not change the canonical SHA-256 or size") + } + artifacts := []struct { + result CreateResult + body []byte + }{ + {result: compressed, body: compressible}, + {result: rawSmall, body: belowThreshold}, + {result: rawSavings, body: incompressible}, + } + for _, artifact := range artifacts { + assertCanonical(artifact.result, artifact.body) + } + + backlog, err := store.LooseBacklog(t.Context()) + require.NoError(t, err) + assert.Equal(t, int64(3), backlog.EligibleObjects) + assert.Equal(t, int64(len(compressible)+len(belowThreshold)+len(incompressible)), + backlog.EligibleBytes) + assert.Equal(t, + compressed.Physical.StoredBytes+rawSmall.Physical.StoredBytes+rawSavings.Physical.StoredBytes, + backlog.EligibleStoredBytes, + "the indexed backlog must report physical loose bytes after compression", + ) + + packed, err := store.Pack(t.Context(), 1<<20) + require.NoError(t, err) + assert.Equal(t, 3, packed.PackedObjects) + assert.Equal(t, backlog.EligibleBytes, packed.LogicalBytes) + assert.False(t, packed.More) + for _, artifact := range artifacts { + assertCanonical(artifact.result, artifact.body) + } + backlog, err = store.LooseBacklog(t.Context()) + require.NoError(t, err) + assert.Zero(t, backlog.EligibleObjects) + assert.Zero(t, backlog.EligibleBytes) + assert.Zero(t, backlog.EligibleStoredBytes) +} + +func TestCollectDocbankWalkJoinsCleanupErrors(t *testing.T) { + nextErr := errors.New("walk page failed") + closeErr := errors.New("walk cleanup failed") + entry := docbank.WalkEntry{Path: "/v1"} + for _, test := range []struct { + name string + walker *docbankWalkerStub + want []docbank.WalkEntry + wantErrors []error + }{ + { + name: "EOF remains successful", + walker: &docbankWalkerStub{pages: [][]docbank.WalkEntry{{entry}}}, + want: []docbank.WalkEntry{entry}, + }, + { + name: "EOF exposes cleanup failure", + walker: &docbankWalkerStub{pages: [][]docbank.WalkEntry{{entry}}, closeErr: closeErr}, + want: []docbank.WalkEntry{entry}, + wantErrors: []error{closeErr}, + }, + { + name: "page and cleanup failures are joined", + walker: &docbankWalkerStub{nextErr: nextErr, closeErr: closeErr}, + wantErrors: []error{nextErr, closeErr}, + }, + } { + t.Run(test.name, func(t *testing.T) { + entries, err := collectDocbankWalk(t.Context(), func( + context.Context, + ) (docbankWalker, error) { + return test.walker, nil + }) + assert.Equal(t, test.want, entries) + if len(test.wantErrors) == 0 { + assert.NoError(t, err) + } + for _, wantErr := range test.wantErrors { + assert.ErrorIs(t, err, wantErr) + } + assert.Equal(t, 1, test.walker.closeCalls) + }) + } +} + +func TestWalkCheckpointFloorStreamsAllQuarantinePages(t *testing.T) { + collection := docbankQuarantineRoot + "/" + contractOrigin + "/" + string(KindCheckpoints) + prefix := collection + "/" + strings.Repeat("a", 32) + "-" + firstPage := make([]docbank.WalkEntry, checkpointFloorPageSize) + for i := range firstPage { + firstPage[i] = docbank.WalkEntry{ + Path: prefix + fmt.Sprintf("cp-%010d.json", i+1), + Node: docbank.Node{BlobHash: "blob"}, + } + } + walker := &docbankWalkerStub{pages: [][]docbank.WalkEntry{ + firstPage, + { + {Path: collection + "/malformed", Node: docbank.Node{BlobHash: "blob"}}, + {Path: prefix + "not-a-checkpoint", Node: docbank.Node{BlobHash: "blob"}}, + {Path: prefix + "cp-0000000900.json", Node: docbank.Node{BlobHash: "blob"}}, + }, + }} + + floor, err := walkCheckpointFloor(t.Context(), collection, true, func( + context.Context, + ) (docbankWalker, error) { + return walker, nil + }) + require.NoError(t, err) + assert.Equal(t, 900, floor) + assert.Equal(t, 1, walker.closeCalls) +} + +func TestWalkCheckpointFloorSkipsMalformedLiveNamesAcrossPages(t *testing.T) { + collection := docbankLiveRoot + "/" + contractOrigin + "/" + string(KindCheckpoints) + walker := &docbankWalkerStub{pages: [][]docbank.WalkEntry{ + { + {Path: collection + "/cp-0000000007.json", Node: docbank.Node{BlobHash: "blob"}}, + {Path: collection + "/cp-malformed.json", Node: docbank.Node{BlobHash: "blob"}}, + }, + { + {Path: collection + "/nested/cp-0000000999.json", Node: docbank.Node{BlobHash: "blob"}}, + {Path: collection + "/cp-0000000042.json", Node: docbank.Node{BlobHash: "blob"}}, + }, + }} + + floor, err := walkCheckpointFloor(t.Context(), collection, false, func( + context.Context, + ) (docbankWalker, error) { + return walker, nil + }) + require.NoError(t, err) + assert.Equal(t, 42, floor) +} + +func TestWalkCheckpointFloorPropagatesCancellationAndCloseError(t *testing.T) { + closeErr := errors.New("close walker") + ctx, cancel := context.WithCancel(t.Context()) + walker := &docbankWalkerStub{ + pages: [][]docbank.WalkEntry{{}}, + closeErr: closeErr, + onNext: cancel, + } + + _, err := walkCheckpointFloor(ctx, "/v1/origin/checkpoints", false, func( + context.Context, + ) (docbankWalker, error) { + return walker, nil + }) + assert.ErrorIs(t, err, context.Canceled) + assert.ErrorIs(t, err, closeErr) +} + +func TestDocbankIteratorClosesWalkerExactlyOnce(t *testing.T) { + t.Run("explicit close", func(t *testing.T) { + walker := &docbankWalkerStub{} + iterator := &docbankOriginIterator{ + iterator: docbankIterator{ + opened: true, + state: docbankTraversal{walker: walker}, + }, + } + + require.NoError(t, iterator.Close()) + require.NoError(t, iterator.Close()) + assert.Equal(t, 1, walker.closeCalls) + _, err := iterator.Next(t.Context(), 1) + assert.ErrorIs(t, err, fs.ErrClosed) + }) + + t.Run("cancellation", func(t *testing.T) { + walker := &docbankWalkerStub{} + iterator := &docbankOriginIterator{ + iterator: docbankIterator{ + opened: true, + state: docbankTraversal{walker: walker}, + }, + } + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err := iterator.Next(ctx, 1) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, 1, walker.closeCalls) + _, err = iterator.Next(t.Context(), 1) + assert.ErrorIs(t, err, fs.ErrClosed) + }) + + t.Run("end of iteration", func(t *testing.T) { + walker := &docbankWalkerStub{} + iterator := &docbankOriginIterator{ + iterator: docbankIterator{ + opened: true, + state: docbankTraversal{walker: walker}, + }, + } + + _, err := iterator.Next(t.Context(), 1) + assert.ErrorIs(t, err, io.EOF) + assert.Equal(t, 1, walker.closeCalls) + _, err = iterator.Next(t.Context(), 1) + assert.ErrorIs(t, err, io.EOF) + assert.Equal(t, 1, walker.closeCalls) + }) +} + +type docbankWalkerStub struct { + pages [][]docbank.WalkEntry + nextErr error + closeErr error + next int + closeCalls int + onNext func() +} + +func (w *docbankWalkerStub) Next(context.Context) ([]docbank.WalkEntry, error) { + if w.onNext != nil { + w.onNext() + w.onNext = nil + } + if w.next < len(w.pages) { + page := w.pages[w.next] + w.next++ + return page, nil + } + if w.nextErr != nil { + return nil, w.nextErr + } + return nil, io.EOF +} + +func (w *docbankWalkerStub) Close() error { + w.closeCalls++ + return w.closeErr +} + +func newTestDocbankStore(t *testing.T, config docbank.Config) (*docbank.Vault, *docbankStore) { + t.Helper() + if config.Root == "" { + config.Root = t.TempDir() + } + vault, err := docbank.New(t.Context(), config) + require.NoError(t, err) + store := newDocbankContent(vault) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + return vault, store +} + +func newTestArtifactStore(t *testing.T) *docbankStore { + t.Helper() + _, store := newTestDocbankStore(t, docbank.Config{}) + return store +} + +func createMetadataArtifactInStore( + t *testing.T, store ArtifactStore, event metadataEvent, +) Ref { + t.Helper() + stamp, err := ParseHLCTimestamp(event.HLC) + require.NoError(t, err) + data, err := canonicalJSON(event) + require.NoError(t, err) + hash := hashHex(data) + ref, err := NewRef( + event.Origin, KindMeta, stamp.OrderingKey(hash)+metadataEventExtension, + ) + require.NoError(t, err) + createContractArtifact(t, store, ref, data) + return ref +} + +// newProtocolTestStore opens a real isolated Docbank store at root without +// registering test cleanup. Callers use it for explicit close/reopen protocol +// and transport scenarios. +func newProtocolTestStore(root string) (*docbankStore, error) { + vault, err := docbank.New(context.Background(), docbank.Config{Root: root}) + if err != nil { + return nil, err + } + return newDocbankContent(vault), nil +} + +func walkDocbankTestEntries( + t *testing.T, vault *docbank.Vault, root string, +) []docbank.WalkEntry { + t.Helper() + walker, err := vault.Walk(t.Context(), root, docbank.WalkOptions{PageSize: 2}) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, walker.Close()) }) + var entries []docbank.WalkEntry + for { + page, err := walker.Next(t.Context()) + if errors.Is(err, io.EOF) { + return entries + } + require.NoError(t, err) + for _, entry := range page { + if entry.Node.BlobHash != "" { + entries = append(entries, entry) + } + } + } +} + +func mustDocbankNode(t *testing.T, vault *docbank.Vault, path string) docbank.Node { + t.Helper() + node, err := vault.Stat(t.Context(), path) + require.NoError(t, err) + return node +} + +func createCheckpointBody( + t *testing.T, store ArtifactStore, sequence int, body []byte, +) CreateResult { + t.Helper() + ref := requireContractRef(t, contractOrigin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequence)) + return createContractArtifact(t, store, ref, body) +} + +func deterministicDocbankBytes(size int) []byte { + data := make([]byte, 0, size) + for counter := 0; len(data) < size; counter++ { + sum := sha256.Sum256([]byte(strconv.Itoa(counter))) + data = append(data, sum[:]...) + } + return data[:size] +} diff --git a/internal/artifact/store_iterator_test.go b/internal/artifact/store_iterator_test.go new file mode 100644 index 000000000..cfde5ee7a --- /dev/null +++ b/internal/artifact/store_iterator_test.go @@ -0,0 +1,84 @@ +package artifact + +import ( + "context" + "errors" + "io" +) + +type testOriginIterator struct { + next func(context.Context, int) ([]string, error) + close func() error +} + +func (i *testOriginIterator) Next(ctx context.Context, limit int) ([]string, error) { + return i.next(ctx, limit) +} + +func (i *testOriginIterator) Close() error { + if i.close == nil { + return nil + } + return i.close() +} + +type testEntryIterator struct { + next func(context.Context, int) ([]Entry, error) + close func() error +} + +func (i *testEntryIterator) Next(ctx context.Context, limit int) ([]Entry, error) { + return i.next(ctx, limit) +} + +func (i *testEntryIterator) Close() error { + if i.close == nil { + return nil + } + return i.close() +} + +func firstStoreEntryPage( + ctx context.Context, store ArtifactStore, origin string, kind Kind, limit int, +) (_ Page, retErr error) { + iterator, err := store.Entries(ctx, origin, kind) + if err != nil { + return Page{}, err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + items, err := iterator.Next(ctx, limit) + if errors.Is(err, io.EOF) { + err = nil + } + return Page{Items: items}, err +} + +func firstStoreOriginPage( + ctx context.Context, store ArtifactStore, limit int, +) (_ []string, retErr error) { + iterator, err := store.Origins(ctx) + if err != nil { + return nil, err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + items, err := iterator.Next(ctx, limit) + if errors.Is(err, io.EOF) { + err = nil + } + return items, err +} + +func firstStoreQuarantinePage( + ctx context.Context, store ArtifactQuarantineStore, limit int, +) (_ []QuarantinedEntry, retErr error) { + iterator, err := store.Quarantined(ctx) + if err != nil { + return nil, err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + items, err := iterator.Next(ctx, limit) + if errors.Is(err, io.EOF) { + err = nil + } + return items, err +} diff --git a/internal/artifact/stream_bench_test.go b/internal/artifact/stream_bench_test.go new file mode 100644 index 000000000..6adecb090 --- /dev/null +++ b/internal/artifact/stream_bench_test.go @@ -0,0 +1,223 @@ +package artifact + +import ( + "bytes" + "context" + "fmt" + "io" + "os" + "runtime" + "testing" + "time" + + "go.kenn.io/docbank" +) + +var artifactBenchmarkSizes = []struct { + name string + size int +}{ + {name: "32KiB", size: 32 << 10}, + {name: "1MiB", size: 1 << 20}, + {name: "16MiB", size: 16 << 20}, +} + +func BenchmarkWireEncode(b *testing.B) { + for _, codec := range []struct { + name string + kind Kind + }{ + {name: "identity", kind: KindRaw}, + {name: "zstd", kind: KindSegments}, + } { + for _, size := range artifactBenchmarkSizes { + b.Run(codec.name+"/"+size.name, func(b *testing.B) { + body := deterministicDocbankBytes(size.size) + ref := benchmarkArtifactRef(b, codec.kind, body) + b.ReportAllocs() + b.SetBytes(int64(len(body))) + b.ResetTimer() + for b.Loop() { + if err := EncodeWire(context.Background(), ref, bytes.NewReader(body), io.Discard); err != nil { + b.Fatal(err) + } + } + }) + } + } +} + +func BenchmarkWireDecode(b *testing.B) { + for _, codec := range []struct { + name string + kind Kind + }{ + {name: "identity", kind: KindRaw}, + {name: "zstd", kind: KindSegments}, + } { + for _, size := range artifactBenchmarkSizes { + b.Run(codec.name+"/"+size.name, func(b *testing.B) { + body := deterministicDocbankBytes(size.size) + ref := benchmarkArtifactRef(b, codec.kind, body) + wire, err := ToWireRef(ref) + if err != nil { + b.Fatal(err) + } + var encoded bytes.Buffer + if err := EncodeWire(context.Background(), ref, bytes.NewReader(body), &encoded); err != nil { + b.Fatal(err) + } + limits := WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body)), + } + b.ReportAllocs() + b.SetBytes(int64(len(body))) + b.ResetTimer() + for b.Loop() { + if err := DecodeWire(context.Background(), wire, + bytes.NewReader(encoded.Bytes()), io.Discard, limits); err != nil { + b.Fatal(err) + } + } + }) + } + } +} + +func BenchmarkDocbankVerifiedRead(b *testing.B) { + for _, size := range artifactBenchmarkSizes { + b.Run(size.name, func(b *testing.B) { + ctx := context.Background() + vault, err := docbank.New(ctx, docbank.Config{ + Root: b.TempDir(), + LooseCompression: docbank.LooseCompressionOptions{ + Enabled: true, + MinBytes: 1, + MinSavingsPercent: 1, + }, + }) + if err != nil { + b.Fatal(err) + } + store := newDocbankContent(vault) + b.Cleanup(func() { + if err := store.Close(); err != nil { + b.Error(err) + } + }) + pattern := []byte("agent session message with repeated text\n") + body := bytes.Repeat(pattern, (size.size+len(pattern)-1)/len(pattern))[:size.size] + ref := benchmarkArtifactRef(b, KindRaw, body) + identity, err := NewIdentity(hashHex(body), int64(len(body))) + if err != nil { + b.Fatal(err) + } + created, err := store.Create(ctx, ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + if err != nil { + b.Fatal(err) + } + if created.Physical.Encoding != "zstd" { + b.Fatalf("expected compressed loose object, got %q", created.Physical.Encoding) + } + packed, err := store.Pack(ctx, int64(len(body))+1) + if err != nil { + b.Fatal(err) + } + if packed.PackedObjects != 1 { + b.Fatalf("expected one packed object, got %d", packed.PackedObjects) + } + + b.ReportAllocs() + b.SetBytes(int64(len(body))) + stopHeapSamples := startArtifactHeapSamples(b) + b.ResetTimer() + for b.Loop() { + entry, reader, err := store.Open(ctx, ref) + if err != nil { + b.Fatal(err) + } + if _, err := io.Copy(io.Discard, reader); err != nil { + b.Fatal(err) + } + if err := reader.Verify(); err != nil { + b.Fatal(err) + } + if err := reader.Close(); err != nil { + b.Fatal(err) + } + if entry.Identity != identity { + b.Fatalf("verified read identity changed: got %+v want %+v", entry.Identity, identity) + } + } + b.StopTimer() + stopHeapSamples() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + runtime.GC() + runtime.ReadMemStats(&after) + if os.Getenv("AGENTSVIEW_ARTIFACT_PROFILE_SAMPLES") != "" { + b.Logf("forced-gc heap_alloc_before=%d heap_alloc_after=%d heap_objects_before=%d heap_objects_after=%d", + before.HeapAlloc, after.HeapAlloc, before.HeapObjects, after.HeapObjects) + } + b.ReportMetric(float64(before.HeapAlloc), "live-heap-before-final-gc-B") + b.ReportMetric(float64(after.HeapAlloc), "live-heap-after-final-gc-B") + runtime.KeepAlive(body) + }) + } +} + +func startArtifactHeapSamples(b *testing.B) func() { + b.Helper() + if os.Getenv("AGENTSVIEW_ARTIFACT_PROFILE_SAMPLES") == "" { + return func() {} + } + started := time.Now() + done := make(chan struct{}) + stopped := make(chan struct{}) + sample := func(label string) { + var stats runtime.MemStats + runtime.ReadMemStats(&stats) + b.Logf("retention-sample label=%s elapsed=%s heap_alloc=%d heap_objects=%d heap_sys=%d", + label, time.Since(started).Round(time.Second), stats.HeapAlloc, stats.HeapObjects, stats.HeapSys) + } + sample("start") + go func() { + defer close(stopped) + ticker := time.NewTicker(30 * time.Second) + defer ticker.Stop() + for { + select { + case <-ticker.C: + sample("periodic") + case <-done: + return + } + } + }() + return func() { + close(done) + <-stopped + sample("end") + } +} + +func benchmarkArtifactRef(b *testing.B, kind Kind, body []byte) Ref { + b.Helper() + name := hashHex(body) + switch kind { + case KindSegments: + name += ".ndjson" + case KindManifests: + name += ".json" + case KindRaw: + default: + b.Fatalf("unsupported benchmark artifact kind %q", kind) + } + ref, err := NewRef(contractOrigin, kind, name) + if err != nil { + b.Fatal(fmt.Errorf("creating benchmark ref: %w", err)) + } + return ref +} diff --git a/internal/artifact/sync.go b/internal/artifact/sync.go new file mode 100644 index 000000000..b849425e2 --- /dev/null +++ b/internal/artifact/sync.go @@ -0,0 +1,3430 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "hash" + "io" + "io/fs" + "net" + "net/url" + "os" + "path/filepath" + "reflect" + "slices" + "sort" + "strconv" + "strings" + "sync" + "time" + + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" +) + +const ( + formatVersion = 1 + originStateKey = "artifact_origin_id" + importStatePrefix = "artifact_import:" + exportStatePrefix = "artifact_export:" + exportSourceStatePrefix = "artifact_export_source:" + tempFilePrefix = ".tmp-" + segmentTargetSize = int64(32 << 20) + + // Cardinality caps complement the byte caps: 4,096 records keeps one + // segment's decoded object graph bounded, while 32,768 records and 256 MiB + // leave ample room for unusually long sessions without letting many valid + // chunks amplify during aggregation. Sixteen references accommodate uneven + // 32 MiB chunks; the aggregate byte cap remains the final session bound. + maxManifestSegments = 16 + maxManifestUsageEvents = 32_768 + maxSegmentMessages = 4_096 + maxSessionMessages = 32_768 + maxSessionDecodedBytes = int64(256 << 20) + + // Nested collections need independent caps because compact empty objects can + // amplify far beyond the decoded byte budget when unmarshaled. A message may + // still describe unusually wide tool fan-out, and one tool may retain a long + // result history. Segment totals keep one decoded chunk modest; session totals + // allow eight full nested-budget segments, matching the message-count ratio. + maxMessageToolCalls = 256 + maxToolResultEvents = 1_024 + maxSegmentToolCalls = 8_192 + maxSegmentResultEvents = 32_768 + maxSessionToolCalls = 65_536 + maxSessionResultEvents = 262_144 +) + +const ( + checkpointFloorPageSize = 128 + checkpointDecodedLimit = int64(64 << 20) + artifactImportPageSize = 128 +) + +type artifactCheckpointSequenceDB interface { + GetArtifactCheckpointFloor(context.Context, string) (int, bool, error) + ReserveArtifactCheckpointSequence(context.Context, string, int) (int, error) +} + +type checkpointFloorStore interface { + checkpointFloor(context.Context, string) (int, error) +} + +// artifactLimits bounds decoded collection cardinality in addition to raw +// bytes. The production values are intentionally generous for real sessions +// while preventing small JSON records from amplifying into unbounded Go +// object graphs. +type artifactLimits struct { + manifestSegments int + manifestUsageEvents int + segmentMessages int + sessionMessages int + sessionDecodedBytes int64 + messageToolCalls int + toolResultEvents int + segmentToolCalls int + segmentResultEvents int + sessionToolCalls int + sessionResultEvents int +} + +func productionArtifactLimits() artifactLimits { + return artifactLimits{ + manifestSegments: maxManifestSegments, + manifestUsageEvents: maxManifestUsageEvents, + segmentMessages: maxSegmentMessages, + sessionMessages: maxSessionMessages, + sessionDecodedBytes: maxSessionDecodedBytes, + messageToolCalls: maxMessageToolCalls, + toolResultEvents: maxToolResultEvents, + segmentToolCalls: maxSegmentToolCalls, + segmentResultEvents: maxSegmentResultEvents, + sessionToolCalls: maxSessionToolCalls, + sessionResultEvents: maxSessionResultEvents, + } +} + +type nestedCollectionCounts struct { + toolCalls int + resultEvents int +} + +type segmentPreflight struct { + records [][]byte + nested nestedCollectionCounts +} + +func exceedsCollectionLimit(current, additional, limit int) bool { + return current > limit || additional > limit-current +} + +var errIncompleteArtifact = errors.New("incomplete artifact") + +var errFutureArtifactVersion = errors.New("future artifact version") + +// SyncOptions configures a local-first artifact folder sync. +type SyncOptions struct { + DataDir string + Target string + Origin string + // Now is the wall-clock source for advancing the metadata HLC past + // observed remote events. When nil, time.Now is used. Sharing it with the + // local metadata recorder keeps import and local edits on one time base. + Now func() time.Time + // Token is the Bearer token for an HTTP peer target. It is ignored by + // folder and object-store targets. + Token string + // AllowInsecure permits plaintext HTTP to a non-loopback peer. Loopback + // HTTP remains allowed without this override. + AllowInsecure bool + // BaselineMetadata writes metadata events for existing local curation before + // exchanging artifacts. It is intended for first-time initialization. + BaselineMetadata bool + // OnDataChanged is called after a foreign import writes local rows. + OnDataChanged func() +} + +// Sync runs one artifact sync, selecting the transport from the target shape: +// an http(s):// URL uses the HTTP peer transport, anything else is treated as a +// local folder target. +func Sync(ctx context.Context, database *db.DB, opts SyncOptions) (SyncResult, error) { + if IsFolderTarget(opts.Target) { + return syncFolderWithOwnedStore(ctx, database, opts) + } + tr, err := syncTransport(nil, opts) + if err != nil { + return SyncResult{}, err + } + return syncWithTransport(ctx, database, opts, tr) +} + +// SyncWithStore runs one exchange through a caller-owned artifact store. The +// store remains open when the exchange returns. +func SyncWithStore( + ctx context.Context, database *db.DB, store ArtifactStore, opts SyncOptions, +) (SyncResult, error) { + if store == nil { + return SyncResult{}, errors.New("artifact store is required") + } + tr, err := syncTransport(nil, opts) + if err != nil { + return SyncResult{}, err + } + return syncContentWithTransport(ctx, database, nil, store, nil, opts, tr) +} + +// SyncWithRepository runs one exchange through a caller-owned local +// repository, including folder-target identity validation and batch packing. +// The repository remains open when the exchange returns. +func SyncWithRepository( + ctx context.Context, database *db.DB, repository *Repository, opts SyncOptions, +) (SyncResult, error) { + if repository == nil || repository.Closed() { + return SyncResult{}, errors.New("artifact repository is required") + } + tr, err := syncTransport(repository, opts) + if err != nil { + return SyncResult{}, err + } + return syncRepositoryWithTransport(ctx, database, repository, opts, tr) +} + +func syncTransport(repository *Repository, opts SyncOptions) (Transport, error) { + if err := ValidateSyncTarget(opts.Target); err != nil { + return nil, err + } + if IsHTTPTarget(opts.Target) { + tr, err := newHTTPTransport(opts.Target, opts.Token, opts.AllowInsecure) + if err != nil { + return nil, err + } + return tr, nil + } + if IsObjectTarget(opts.Target) { + tr, err := newObjectTransport(opts.Target, ObjectStoreOptionsFromEnv()) + if err != nil { + return nil, err + } + return tr, nil + } + if repository == nil { + return nil, fmt.Errorf( + "%w: folder exchange requires retained repository identity", + ErrArtifactUnsupported, + ) + } + transport, err := repository.NewFolderTransport(opts.Target) + if err != nil { + return nil, err + } + return transport, nil +} + +// ValidateSyncTarget rejects URL components that can carry credentials or be +// confused with transport-owned API paths. Folder targets are left untouched. +func ValidateSyncTarget(target string) error { + if target == "" { + return fmt.Errorf("%w: artifact sync target is required", ErrArtifactInvalid) + } + if !IsHTTPTarget(target) && !IsObjectTarget(target) { + return nil + } + u, err := url.Parse(target) + if err != nil || u == nil || u.Host == "" { + return fmt.Errorf("%w: artifact sync URL is invalid", ErrArtifactInvalid) + } + if u.User != nil || u.RawQuery != "" || u.Fragment != "" { + return fmt.Errorf("%w: artifact sync URL must not contain credentials, query, or fragment", + ErrArtifactInvalid) + } + return nil +} + +// SyncResult summarizes a folder artifact sync run. +type SyncResult struct { + Origin string + ExportedSessions int + ImportedSessions int + ImportedMessages int + ImportedMetadata int +} + +// ImportResult summarizes local rows changed by artifact import. +type ImportResult struct { + Sessions int + Messages int + Metadata int + Deferred int +} + +// Changed reports whether the import wrote user-visible local data. +func (r ImportResult) Changed() bool { + return r.Sessions > 0 || r.Messages > 0 || r.Metadata > 0 +} + +// SyncFolder exports local sessions to the local artifact store, exchanges the +// store with target, and imports foreign origins from the exchanged artifacts. +func SyncFolder(ctx context.Context, database *db.DB, opts SyncOptions) (SyncResult, error) { + return Sync(ctx, database, opts) +} + +func syncFolderWithOwnedStore( + ctx context.Context, database *db.DB, opts SyncOptions, +) (_ SyncResult, retErr error) { + if opts.DataDir == "" { + return SyncResult{}, errors.New("artifact sync data dir is required") + } + repository, err := OpenRepository(ctx, opts.DataDir) + if err != nil { + return SyncResult{}, fmt.Errorf("opening artifact repository: %w", err) + } + defer func() { retErr = errors.Join(retErr, repository.Close()) }() + transport, err := syncTransport(repository, opts) + if err != nil { + return SyncResult{}, err + } + return syncRepositoryWithTransport(ctx, database, repository, opts, transport) +} + +// syncWithTransport runs one artifact sync over any transport: export local +// sessions, exchange the store with the remote via set-union, then import +// foreign origins. Folder, HTTP peer, and object-store targets differ only in +// the transport's Prepare and Exchange. +func syncWithTransport( + ctx context.Context, + database *db.DB, + opts SyncOptions, + tr Transport, +) (_ SyncResult, retErr error) { + if opts.DataDir == "" { + return SyncResult{}, errors.New("artifact sync data dir is required") + } + repository, err := OpenRepository(ctx, opts.DataDir) + if err != nil { + return SyncResult{}, fmt.Errorf("opening artifact repository: %w", err) + } + defer func() { retErr = errors.Join(retErr, repository.Close()) }() + return syncRepositoryWithTransport(ctx, database, repository, opts, tr) +} + +func syncRepositoryWithTransport( + ctx context.Context, + database *db.DB, + repository *Repository, + opts SyncOptions, + tr Transport, +) (SyncResult, error) { + return syncContentWithTransport( + ctx, database, repository, repository.Content(), repository.NotifyBatch, opts, tr, + ) +} + +func syncContentWithTransport( + ctx context.Context, + database *db.DB, + repository *Repository, + localStore ArtifactStore, + notifyBatch func(context.Context), + opts SyncOptions, + tr Transport, +) (_ SyncResult, retErr error) { + notificationCtx := ctx + defer func() { + if retErr == nil && notifyBatch != nil { + notifyBatch(notificationCtx) + } + }() + if closer, ok := tr.(io.Closer); ok { + defer func() { retErr = errors.Join(retErr, closer.Close()) }() + } + ctx = SuppressArtifactMaintenance(ctx) + if err := tr.Prepare(ctx, localStore); err != nil { + return SyncResult{}, err + } + origin := opts.Origin + if origin == "" { + var err error + origin, err = EnsureOrigin(database) + if err != nil { + return SyncResult{}, err + } + } else if err := validateOriginID(origin); err != nil { + return SyncResult{}, err + } + coordinator := NewStoreImportCoordinator(database, localStore, origin) + coordinator.now = opts.Now + transportStore := newCoordinatedTransportStore(database, localStore, coordinator) + + var imported ImportResult + var baselineSnapshot db.MetadataBaselineSnapshot + if opts.BaselineMetadata { + var err error + baselineSnapshot, err = database.MetadataBaselineSnapshot(ctx) + if err != nil { + return SyncResult{}, err + } + if err := tr.Exchange(ctx, transportStore); err != nil { + return SyncResult{}, err + } + preBaselineImported, err := coordinator.Finalize(ctx) + if err != nil { + return SyncResult{}, err + } + imported.Sessions += preBaselineImported.Sessions + imported.Messages += preBaselineImported.Messages + imported.Metadata += preBaselineImported.Metadata + + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, + Store: localStore, + Now: opts.Now, + }) + if _, err := recorder.AppendBaselineSnapshot(ctx, baselineSnapshot); err != nil { + return SyncResult{}, err + } + } + var ( + exported ExportResult + err error + ) + if repository != nil { + exported, err = PublishRepositoryArtifacts(ctx, database, repository, ExportOptions{ + Origin: origin, + }) + } else { + exported, err = ExportToStore(ctx, database, localStore, ExportOptions{ + Origin: origin, + }) + } + if err != nil { + return SyncResult{}, err + } + if err := tr.Exchange(ctx, transportStore); err != nil { + return SyncResult{}, err + } + postExportImported, err := coordinator.Finalize(ctx) + if err != nil { + return SyncResult{}, err + } + imported.Sessions += postExportImported.Sessions + imported.Messages += postExportImported.Messages + imported.Metadata += postExportImported.Metadata + if imported.Changed() && opts.OnDataChanged != nil { + opts.OnDataChanged() + } + return SyncResult{ + Origin: origin, + ExportedSessions: exported.ExportedSessions, + ImportedSessions: imported.Sessions, + ImportedMessages: imported.Messages, + ImportedMetadata: imported.Metadata, + }, nil +} + +// EnsureOrigin returns the persisted origin ID, creating one when absent. +func EnsureOrigin(database *db.DB) (string, error) { + origin, err := StoredOrigin(database) + if err != nil { + return "", err + } + if origin != "" { + return origin, nil + } + origin, err = newOriginID() + if err != nil { + return "", err + } + if err := validateOriginID(origin); err != nil { + return "", fmt.Errorf("generated artifact origin: %w", err) + } + if err := database.SetSyncState(originStateKey, origin); err != nil { + return "", fmt.Errorf("persisting artifact origin: %w", err) + } + return origin, nil +} + +// AdoptOrigin persists origin as this machine's artifact origin in the database +// sync state so DB-derived lookups (EnsureOrigin and its callers) agree with the +// authoritative config origin. It validates the input and is idempotent: it only +// writes when the stored value differs. The config origin always wins, so a +// previously stored value is overwritten to converge on a single origin. +func AdoptOrigin(database *db.DB, origin string) error { + if err := validateOriginID(origin); err != nil { + return fmt.Errorf("adopting artifact origin: %w", err) + } + existing, err := StoredOrigin(database) + if err != nil { + return err + } + if existing == origin { + return nil + } + if err := database.SetSyncState(originStateKey, origin); err != nil { + return fmt.Errorf("persisting artifact origin: %w", err) + } + return nil +} + +// StoredOrigin returns the persisted origin ID without creating one. +func StoredOrigin(database *db.DB) (string, error) { + origin, err := database.GetSyncState(originStateKey) + if err != nil { + return "", fmt.Errorf("reading artifact origin: %w", err) + } + if origin != "" { + if err := validateOriginID(origin); err != nil { + return "", fmt.Errorf("stored artifact origin: %w", err) + } + return origin, nil + } + return "", nil +} + +type syncStateValueReader interface { + SyncStateValues(keys []string) (map[string]string, error) +} + +// ImportedSessionIDs returns the candidate session IDs with durable artifact +// import provenance. A foreign machine~id shape is shared by other import +// mechanisms, so callers must query the exact provenance keys rather than +// infer artifact ownership from the session row or scan all historical imports. +func ImportedSessionIDs( + database syncStateValueReader, candidateIDs []string, +) (map[string]struct{}, error) { + ids := make(map[string]struct{}) + if len(candidateIDs) == 0 { + return ids, nil + } + keys := make([]string, 0, len(candidateIDs)) + keyToID := make(map[string]string, len(candidateIDs)) + for _, gid := range candidateIDs { + origin, nativeID, ok := strings.Cut(gid, "~") + if !ok || origin == "" || nativeID == "" { + continue + } + key := importStateKey(origin, gid) + keys = append(keys, key) + keyToID[key] = gid + } + if len(keys) == 0 { + return ids, nil + } + states, err := database.SyncStateValues(keys) + if err != nil { + return nil, fmt.Errorf("reading artifact import provenance: %w", err) + } + for key := range states { + if gid, ok := keyToID[key]; ok { + ids[gid] = struct{}{} + } + } + return ids, nil +} + +func newOriginID() (string, error) { + host, err := os.Hostname() + if err != nil || strings.TrimSpace(host) == "" { + host = "machine" + } + host = sanitizeOriginPart(host) + if host == "" || host == "local" { + host = "machine" + } + var suffix [3]byte + if _, err := rand.Read(suffix[:]); err != nil { + return "", fmt.Errorf("generating artifact origin suffix: %w", err) + } + return fmt.Sprintf("%s-%s", host, hex.EncodeToString(suffix[:])), nil +} + +func validateOriginID(origin string) error { + return config.ValidateArtifactOriginID(origin) +} + +func validateDisjointRoots(localRoot, target string) error { + localAbs, err := filepath.Abs(localRoot) + if err != nil { + return fmt.Errorf("resolving local artifact store: %w", err) + } + targetAbs, err := filepath.Abs(target) + if err != nil { + return fmt.Errorf("resolving artifact sync target: %w", err) + } + localAbs = filepath.Clean(localAbs) + targetAbs = filepath.Clean(targetAbs) + localCanonical, err := canonicalArtifactPath(localAbs) + if err != nil { + return fmt.Errorf("resolving local artifact store symlinks: %w", err) + } + targetCanonical, err := canonicalArtifactPath(targetAbs) + if err != nil { + return fmt.Errorf("resolving artifact sync target symlinks: %w", err) + } + if rootsOverlap(localAbs, targetAbs) || rootsOverlap(localCanonical, targetCanonical) { + return fmt.Errorf( + "artifact sync target %s must not overlap local artifact store %s", + targetCanonical, localCanonical, + ) + } + return nil +} + +func canonicalArtifactPath(path string) (string, error) { + missing := make([]string, 0, 2) + current := path + for { + resolved, err := filepath.EvalSymlinks(current) + if err == nil { + for _, part := range slices.Backward(missing) { + resolved = filepath.Join(resolved, part) + } + return filepath.Clean(resolved), nil + } + if !errors.Is(err, fs.ErrNotExist) { + return "", err + } + parent := filepath.Dir(current) + if parent == current { + return "", err + } + missing = append(missing, filepath.Base(current)) + current = parent + } +} + +func rootsOverlap(a, b string) bool { + return a == b || pathContains(a, b) || pathContains(b, a) +} + +func pathContains(parent, child string) bool { + rel, err := filepath.Rel(parent, child) + if err != nil { + return false + } + return rel != "." && rel != ".." && !strings.HasPrefix(rel, ".."+string(filepath.Separator)) && !filepath.IsAbs(rel) +} + +func sanitizeOriginPart(s string) string { + s = strings.ToLower(strings.TrimSpace(s)) + var b strings.Builder + lastDash := false + for _, r := range s { + ok := r >= 'a' && r <= 'z' || r >= '0' && r <= '9' + if ok { + b.WriteRune(r) + lastDash = false + continue + } + if !lastDash { + b.WriteByte('-') + lastDash = true + } + } + return strings.Trim(b.String(), "-") +} + +type checkpoint struct { + Version int `json:"v"` + Origin string `json:"origin"` + Sequence int `json:"seq"` + Sessions map[string]string `json:"sessions"` +} + +type manifest struct { + Version int `json:"v"` + Origin string `json:"origin"` + NativeSessionID string `json:"native_session_id"` + Session manifestSession `json:"session"` + SessionName *string `json:"session_name,omitempty"` + Segments []string `json:"segments"` + UsageEvents []artifactUsageEvent `json:"usage_events,omitempty"` + RawSource *rawSourceRef `json:"raw_source,omitempty"` + DataVersion int `json:"data_version"` + Generation int `json:"generation"` + // Signal state persisted on the session row but absent from the wire + // Session above, which mirrors only db.Session's JSON-visible fields. + // Carried explicitly so an imported session keeps its tool-call, context, + // and quality signal state instead of resetting to false/zero. Secret-scan + // state is deliberately not carried: findings live outside the manifest, + // so imported sessions are treated as unscanned (see rewriteForImport). + SessionHasToolCalls bool `json:"session_has_tool_calls,omitempty"` + SessionHasContextData bool `json:"session_has_context_data,omitempty"` + SessionQualitySignals *manifestQualitySignals `json:"session_quality_signals,omitempty"` +} + +type artifactUsageEvent struct { + MessageOrdinal *int `json:"message_ordinal,omitempty"` + Source string `json:"source"` + Model string `json:"model"` + InputTokens int `json:"input_tokens,omitempty"` + OutputTokens int `json:"output_tokens,omitempty"` + CacheCreationInputTokens int `json:"cache_creation_input_tokens,omitempty"` + CacheReadInputTokens int `json:"cache_read_input_tokens,omitempty"` + ReasoningTokens int `json:"reasoning_tokens,omitempty"` + CostUSD *float64 `json:"cost_usd,omitempty"` + CostStatus string `json:"cost_status,omitempty"` + CostSource string `json:"cost_source,omitempty"` + OccurredAt string `json:"occurred_at,omitempty"` + DedupKey string `json:"dedup_key,omitempty"` +} + +type rawSourceRef struct { + Hash string `json:"hash"` + Size int64 `json:"size"` + MediaType string `json:"media_type,omitempty"` + Path string `json:"path,omitempty"` +} + +type metadataEvent struct { + Version int `json:"v"` + HLC string `json:"hlc"` + Origin string `json:"origin"` + SessionGID string `json:"session_gid"` + Op string `json:"op"` + Value json.RawMessage `json:"value,omitempty"` + Pin *MetadataPin `json:"pin,omitempty"` +} + +type segmentMessage struct { + Version int `json:"v"` + Ordinal int `json:"ordinal"` + Role string `json:"role"` + Content string `json:"content"` + ThinkingText string `json:"thinking_text,omitempty"` + Timestamp string `json:"timestamp,omitempty"` + HasThinking bool `json:"has_thinking,omitempty"` + HasToolUse bool `json:"has_tool_use,omitempty"` + ContentLength int `json:"content_length,omitempty"` + Model string `json:"model,omitempty"` + TokenUsage json.RawMessage `json:"token_usage,omitempty"` + ContextTokens int `json:"context_tokens,omitempty"` + OutputTokens int `json:"output_tokens,omitempty"` + HasContextTokens bool `json:"has_context_tokens,omitempty"` + HasOutputTokens bool `json:"has_output_tokens,omitempty"` + ClaudeMessageID string `json:"claude_message_id,omitempty"` + ClaudeRequestID string `json:"claude_request_id,omitempty"` + ToolCalls []segmentToolCall `json:"tool_calls,omitempty"` + IsSystem bool `json:"is_system,omitempty"` + SourceType string `json:"source_type,omitempty"` + SourceSubtype string `json:"source_subtype,omitempty"` + SourceUUID string `json:"source_uuid,omitempty"` + SourceParentUUID string `json:"source_parent_uuid,omitempty"` + IsSidechain bool `json:"is_sidechain,omitempty"` + IsCompactBoundary bool `json:"is_compact_boundary,omitempty"` +} + +type segmentToolCall struct { + CallIndex int `json:"call_index"` + ToolName string `json:"tool_name"` + Category string `json:"category,omitempty"` + ToolUseID string `json:"tool_use_id,omitempty"` + InputJSON string `json:"input_json,omitempty"` + FilePath string `json:"file_path,omitempty"` + SkillName string `json:"skill_name,omitempty"` + ResultContentLength int `json:"result_content_length,omitempty"` + ResultContent string `json:"result_content,omitempty"` + SubagentSessionID string `json:"subagent_session_id,omitempty"` + ResultEvents []segmentResultEvent `json:"result_events,omitempty"` +} + +type segmentResultEvent struct { + ToolUseID string `json:"tool_use_id,omitempty"` + AgentID string `json:"agent_id,omitempty"` + SubagentSessionID string `json:"subagent_session_id,omitempty"` + Source string `json:"source"` + Status string `json:"status"` + Content string `json:"content"` + ContentLength int `json:"content_length,omitempty"` + Timestamp string `json:"timestamp,omitempty"` + EventIndex int `json:"event_index"` +} + +type artifactExportStore interface { + ListOwnedSessionIDsForExport(context.Context) ([]string, error) + PendingArtifactExports(context.Context, int) ([]db.ArtifactExportQueueItem, error) + ArtifactExportClaims(context.Context, []string) ([]db.ArtifactExportQueueItem, error) + GetSessionFull(context.Context, string) (*db.Session, error) + GetAllMessages(context.Context, string) ([]db.Message, error) + GetUsageEvents(context.Context, string) ([]db.UsageEvent, error) + ApplyArtifactPublicationChanges(context.Context, string, []db.ArtifactPublicationChange) (int64, bool, error) + AcknowledgeArtifactExports(context.Context, []db.ArtifactExportQueueItem) error + GetArtifactCheckpointHead(context.Context, string) (db.ArtifactCheckpointHead, bool, error) + StreamArtifactPublications(context.Context, string, func(db.ArtifactPublication) error) (int64, error) + RecordArtifactCheckpointHead(context.Context, db.ArtifactCheckpointHead, []db.ArtifactExportQueueItem) error + GetArtifactCheckpointFloor(context.Context, string) (int, bool, error) + ReserveArtifactCheckpointSequence(context.Context, string, int) (int, error) +} + +const artifactExportBatchSize = 128 + +// ExportOptions selects artifact publication work. The default and explicit-ID +// modes process one bounded batch. Full drains bounded claim pages before +// recreating every currently-owned session's immutable dependencies one body +// at a time; publication authority still changes only through guarded claims. +type ExportOptions struct { + Origin string + SessionIDs []string + Full bool +} + +// ExportResult summarizes one canonical store export. +type ExportResult struct { + ExportedSessions int + CheckpointCreated bool + CheckpointSequence int +} + +type queuedArtifactExportStore interface { + PendingArtifactExports(context.Context, int) ([]db.ArtifactExportQueueItem, error) + GetSessionFull(context.Context, string) (*db.Session, error) + GetAllMessages(context.Context, string) ([]db.Message, error) + GetUsageEvents(context.Context, string) ([]db.UsageEvent, error) +} + +type queuedArtifactExport struct { + Item db.ArtifactExportQueueItem + Session *db.Session + Messages []db.Message + UsageEvents []db.UsageEvent +} + +// forEachQueuedArtifactExport loads only the bounded dirty batch and at most +// one complete session body at a time. A missing session represents a pending +// publication deletion and deliberately performs no message or usage reads. +func forEachQueuedArtifactExport( + ctx context.Context, + store queuedArtifactExportStore, + limit int, + visit func(queuedArtifactExport) error, +) error { + if visit == nil { + return errors.New("queued artifact export visitor is required") + } + items, err := store.PendingArtifactExports(ctx, limit) + if err != nil { + return fmt.Errorf("reading queued artifact exports: %w", err) + } + for _, item := range items { + if err := ctx.Err(); err != nil { + return err + } + work := queuedArtifactExport{Item: item} + work.Session, err = store.GetSessionFull(ctx, item.SessionID) + if err != nil { + return fmt.Errorf("loading queued artifact session %s: %w", item.SessionID, err) + } + if work.Session != nil && + (work.Session.Machine != "local" || work.Session.DeletedAt != nil) { + work.Session = nil + } + if work.Session != nil { + work.Messages, err = store.GetAllMessages(ctx, item.SessionID) + if err != nil { + return fmt.Errorf("loading queued artifact messages %s: %w", item.SessionID, err) + } + work.UsageEvents, err = store.GetUsageEvents(ctx, item.SessionID) + if err != nil { + return fmt.Errorf("loading queued artifact usage %s: %w", item.SessionID, err) + } + } + if err := visit(work); err != nil { + return err + } + } + return nil +} + +// ExportToStore publishes generation-guarded work into the canonical artifact +// store. Immutable dependencies are created before their manifest, and each +// bounded page's checkpoint is created last. Full mode may publish several +// bounded pages before its dependency-recovery pass completes. +func ExportToStore( + ctx context.Context, + database artifactExportStore, + store ArtifactStore, + opts ExportOptions, +) (_ ExportResult, retErr error) { + if database == nil { + return ExportResult{}, errors.New("artifact export database is required") + } + if store == nil { + return ExportResult{}, errors.New("artifact export store is required") + } + if err := validateOriginID(opts.Origin); err != nil { + return ExportResult{}, err + } + if len(opts.SessionIDs) > 1024 { + return ExportResult{}, errors.New("artifact export session batch exceeds 1024 rows") + } + if opts.Full && len(opts.SessionIDs) > 0 { + return ExportResult{}, errors.New("full artifact export cannot select session IDs") + } + if opts.Full { + return exportFullToStore(ctx, database, store, opts.Origin) + } + + var claims []db.ArtifactExportQueueItem + var err error + if len(opts.SessionIDs) > 0 { + claims, err = database.ArtifactExportClaims(ctx, opts.SessionIDs) + } else { + queueLimit := artifactExportBatchSize + if opts.Full { + queueLimit = 1024 + } + claims, err = database.PendingArtifactExports(ctx, queueLimit) + } + if err != nil { + return ExportResult{}, fmt.Errorf("reading artifact export queue: %w", err) + } + claimByID := make(map[string]db.ArtifactExportQueueItem, len(claims)) + for _, claim := range claims { + claimByID[claim.SessionID] = claim + } + + selected, err := selectArtifactExportSessionIDs(ctx, database, opts, claims, claimByID) + if err != nil { + return ExportResult{}, err + } + changes := make([]db.ArtifactPublicationChange, 0, len(claims)) + acknowledged := make([]db.ArtifactExportQueueItem, 0, len(claims)) + result := ExportResult{} + for _, sessionID := range selected { + if err := ctx.Err(); err != nil { + return result, err + } + claim, claimed := claimByID[sessionID] + sess, err := database.GetSessionFull(ctx, sessionID) + if err != nil { + return result, fmt.Errorf("loading artifact export session %s: %w", sessionID, err) + } + if sess == nil || sess.Machine != "local" || sess.DeletedAt != nil { + if claimed { + changes = append(changes, db.ArtifactPublicationChange{ + SessionID: sessionID, Generation: claim.Generation, Delete: true, + }) + acknowledged = append(acknowledged, claim) + } + continue + } + messages, err := database.GetAllMessages(ctx, sessionID) + if err != nil { + return result, fmt.Errorf("loading artifact export messages %s: %w", sessionID, err) + } + usageEvents, err := database.GetUsageEvents(ctx, sessionID) + if err != nil { + return result, fmt.Errorf("loading artifact export usage %s: %w", sessionID, err) + } + manifestHash, created, err := exportLoadedSessionToStore( + ctx, store, opts.Origin, sess, messages, usageEvents, + productionArtifactLimits(), + ) + if err != nil { + return result, err + } + if created && claimed { + result.ExportedSessions++ + } + if claimed { + changes = append(changes, db.ArtifactPublicationChange{ + SessionID: sessionID, Generation: claim.Generation, + ManifestHash: manifestHash, SourceFingerprint: manifestHash, + }) + acknowledged = append(acknowledged, claim) + } + } + + head, hadHead, err := database.GetArtifactCheckpointHead(ctx, opts.Origin) + if err != nil { + return result, fmt.Errorf("reading artifact checkpoint head: %w", err) + } + publicationRevision, changed, err := database.ApplyArtifactPublicationChanges( + ctx, opts.Origin, changes, + ) + if err != nil { + return result, err + } + if !changed && hadHead && head.PublicationRevision == publicationRevision { + verified, err := statRecordedCheckpoint(ctx, store, head) + if err != nil { + return result, err + } + if verified { + if err := database.AcknowledgeArtifactExports(ctx, acknowledged); err != nil { + return result, err + } + return result, nil + } + } + + comparableHead := !changed && hadHead + if !hadHead { + head, comparableHead, err = latestValidCheckpointHead(ctx, store, opts.Origin) + if err != nil { + return result, err + } + } + mapSpool, mapDigest, mapRevision, err := spoolArtifactPublicationMap(ctx, database, opts.Origin) + if err != nil { + return result, err + } + defer func() { + if mapSpool != nil { + retErr = errors.Join(retErr, closeAndRemoveExportSpool(mapSpool)) + } + }() + if comparableHead && head.SessionMapSHA256 == mapDigest { + checkpointSpool, checkpointIdentity, err := spoolArtifactCheckpoint( + ctx, mapSpool, opts.Origin, head.Sequence, + ) + if err != nil { + return result, err + } + defer func() { + if checkpointSpool != nil { + retErr = errors.Join(retErr, closeAndRemoveExportSpool(checkpointSpool)) + } + }() + if checkpointIdentity.SHA256 != head.CheckpointSHA256 { + return result, fmt.Errorf( + "%w: recorded checkpoint %d hash differs from canonical publications", + ErrArtifactCorrupt, head.Sequence, + ) + } + head.PublicationRevision = mapRevision + head.CheckpointSize = checkpointIdentity.Size + checkpointRef, err := NewRef(opts.Origin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", head.Sequence)) + if err != nil { + return result, err + } + if err := closeAndRemoveExportSpool(mapSpool); err != nil { + return result, fmt.Errorf("cleaning artifact session map spool: %w", err) + } + mapSpool = nil + create, err := store.Create(ctx, checkpointRef, checkpointIdentity, + canonicalArtifactMediaType(KindCheckpoints), checkpointSpool) + if err != nil { + return result, fmt.Errorf("recreating artifact checkpoint: %w", err) + } + if err := closeAndRemoveExportSpool(checkpointSpool); err != nil { + return result, fmt.Errorf("cleaning artifact checkpoint spool: %w", err) + } + checkpointSpool = nil + if err := database.RecordArtifactCheckpointHead(ctx, head, acknowledged); err != nil { + return result, err + } + result.CheckpointCreated = create.Created + result.CheckpointSequence = head.Sequence + return result, nil + } + + sequence, err := reserveCheckpointSequenceFromStore(ctx, database, store, opts.Origin) + if err != nil { + return result, err + } + checkpointSpool, checkpointIdentity, err := spoolArtifactCheckpoint( + ctx, mapSpool, opts.Origin, sequence, + ) + if err != nil { + return result, err + } + defer func() { + if checkpointSpool != nil { + retErr = errors.Join(retErr, closeAndRemoveExportSpool(checkpointSpool)) + } + }() + checkpointRef, err := NewRef(opts.Origin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequence)) + if err != nil { + return result, err + } + if err := closeAndRemoveExportSpool(mapSpool); err != nil { + return result, fmt.Errorf("cleaning artifact session map spool: %w", err) + } + mapSpool = nil + if _, err := store.Create(ctx, checkpointRef, checkpointIdentity, + canonicalArtifactMediaType(KindCheckpoints), checkpointSpool, + ); err != nil { + return result, fmt.Errorf("creating artifact checkpoint: %w", err) + } + if err := closeAndRemoveExportSpool(checkpointSpool); err != nil { + return result, fmt.Errorf("cleaning artifact checkpoint spool: %w", err) + } + checkpointSpool = nil + if err := database.RecordArtifactCheckpointHead(ctx, db.ArtifactCheckpointHead{ + Origin: opts.Origin, Sequence: sequence, PublicationRevision: mapRevision, + SessionMapSHA256: mapDigest, CheckpointSHA256: checkpointIdentity.SHA256, + CheckpointSize: checkpointIdentity.Size, + }, acknowledged); err != nil { + return result, err + } + result.CheckpointCreated = true + result.CheckpointSequence = sequence + return result, nil +} + +func exportFullToStore( + ctx context.Context, + database artifactExportStore, + store ArtifactStore, + origin string, +) (ExportResult, error) { + result := ExportResult{} + processed := make(map[string]struct{}) + drain := func() error { + for { + if err := ctx.Err(); err != nil { + return err + } + claims, err := database.PendingArtifactExports(ctx, maxArtifactExportBatchSize) + if err != nil { + return fmt.Errorf("reading full artifact export queue: %w", err) + } + if len(claims) == 0 { + return nil + } + ids := make([]string, len(claims)) + for i, claim := range claims { + ids[i] = claim.SessionID + processed[claim.SessionID] = struct{}{} + } + page, err := ExportToStore(ctx, database, store, ExportOptions{ + Origin: origin, SessionIDs: ids, + }) + if err != nil { + return err + } + mergeArtifactExportResult(&result, page) + } + } + if err := drain(); err != nil { + return result, err + } + ids, err := database.ListOwnedSessionIDsForExport(ctx) + if err != nil { + return result, fmt.Errorf("listing sessions for full artifact export: %w", err) + } + for _, sessionID := range ids { + if err := ctx.Err(); err != nil { + return result, err + } + if _, ok := processed[sessionID]; ok { + continue + } + sess, err := database.GetSessionFull(ctx, sessionID) + if err != nil { + return result, fmt.Errorf("loading full artifact export session %s: %w", sessionID, err) + } + if sess == nil || sess.Machine != "local" || sess.DeletedAt != nil { + continue + } + messages, err := database.GetAllMessages(ctx, sessionID) + if err != nil { + return result, fmt.Errorf("loading full artifact export messages %s: %w", sessionID, err) + } + usageEvents, err := database.GetUsageEvents(ctx, sessionID) + if err != nil { + return result, fmt.Errorf("loading full artifact export usage %s: %w", sessionID, err) + } + if _, _, err := exportLoadedSessionToStore( + ctx, store, origin, sess, messages, usageEvents, productionArtifactLimits(), + ); err != nil { + return result, err + } + } + if err := drain(); err != nil { + return result, err + } + for { + final, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin}) + if err != nil { + return result, err + } + mergeArtifactExportResult(&result, final) + pending, err := database.PendingArtifactExports(ctx, 1) + if err != nil { + return result, fmt.Errorf("checking concurrent full artifact work: %w", err) + } + if len(pending) == 0 { + return result, nil + } + if err := drain(); err != nil { + return result, err + } + } +} + +const maxArtifactExportBatchSize = 1024 + +func mergeArtifactExportResult(total *ExportResult, page ExportResult) { + total.ExportedSessions += page.ExportedSessions + if page.CheckpointCreated { + total.CheckpointCreated = true + } + if page.CheckpointSequence > total.CheckpointSequence { + total.CheckpointSequence = page.CheckpointSequence + } +} + +func selectArtifactExportSessionIDs( + ctx context.Context, + database artifactExportStore, + opts ExportOptions, + claims []db.ArtifactExportQueueItem, + claimByID map[string]db.ArtifactExportQueueItem, +) ([]string, error) { + selected := make(map[string]struct{}) + switch { + case opts.Full: + ids, err := database.ListOwnedSessionIDsForExport(ctx) + if err != nil { + return nil, fmt.Errorf("listing sessions for full artifact export: %w", err) + } + for _, id := range ids { + selected[id] = struct{}{} + } + for _, claim := range claims { + selected[claim.SessionID] = struct{}{} + } + case len(opts.SessionIDs) > 0: + for _, id := range opts.SessionIDs { + if _, ok := claimByID[id]; ok { + selected[id] = struct{}{} + } + } + default: + for _, claim := range claims { + selected[claim.SessionID] = struct{}{} + } + } + ids := make([]string, 0, len(selected)) + for id := range selected { + ids = append(ids, id) + } + sort.Strings(ids) + return ids, nil +} + +func exportLoadedSessionToStore( + ctx context.Context, + store ArtifactStore, + origin string, + sess *db.Session, + messages []db.Message, + usageEvents []db.UsageEvent, + limits artifactLimits, +) (string, bool, error) { + if len(messages) > limits.sessionMessages { + return "", false, fmt.Errorf( + "session message limit exceeded for %s: got %d, limit %d", + sess.ID, len(messages), limits.sessionMessages, + ) + } + if len(usageEvents) > limits.manifestUsageEvents { + return "", false, fmt.Errorf( + "manifest usage event limit exceeded for %s: got %d, limit %d", + sess.ID, len(usageEvents), limits.manifestUsageEvents, + ) + } + if err := validateExportNestedCollections(messages, limits); err != nil { + return "", false, fmt.Errorf("validating nested collections for %s: %w", sess.ID, err) + } + + segmentHashes, err := exportMessageSegmentsToStore( + ctx, store, origin, sess.ID, messages, limits, + ) + if err != nil { + return "", false, err + } + + wireSession := manifestSessionFromDB(*sess) + wireSession.Machine = origin + normalizeManifestSessionLocalState(&wireSession) + m := manifest{ + Version: formatVersion, Origin: origin, NativeSessionID: sess.ID, + Session: wireSession, SessionName: sess.SessionName, + Segments: segmentHashes, UsageEvents: canonicalUsageEvents(usageEvents), + DataVersion: sess.DataVersion, Generation: 1, + SessionHasToolCalls: sess.HasToolCalls, + SessionHasContextData: sess.HasContextData, + SessionQualitySignals: manifestQualitySignalsFromDB(sess.StoredQualitySignals()), + } + data, err := canonicalJSON(m) + if err != nil { + return "", false, err + } + if int64(len(data)) > manifestDecodedLimit { + return "", false, fmt.Errorf( + "generated manifest exceeds %d-byte readable limit: got %d bytes", + manifestDecodedLimit, len(data), + ) + } + hash := hashHex(data) + identity, err := NewIdentity(hash, int64(len(data))) + if err != nil { + return "", false, err + } + ref, err := NewRef(origin, KindManifests, hash+".json") + if err != nil { + return "", false, err + } + created, err := store.Create(ctx, ref, identity, + canonicalArtifactMediaType(KindManifests), bytes.NewReader(data)) + if err != nil { + return "", false, fmt.Errorf("creating manifest for %s: %w", sess.ID, err) + } + return hash, created.Created, nil +} + +func exportMessageSegmentsToStore( + ctx context.Context, + store ArtifactStore, + origin string, + sessionID string, + messages []db.Message, + limits artifactLimits, +) (_ []string, retErr error) { + segmentHashes := make([]string, 0, 1) + seen := make(map[string]struct{}) + var decodedBytes int64 + var spool *os.File + var hasher hash.Hash + var segmentBytes int64 + segmentMessages := 0 + segmentNested := nestedCollectionCounts{} + defer func() { + if spool != nil { + retErr = errors.Join(retErr, closeAndRemoveExportSpool(spool)) + } + }() + + start := func() error { + var err error + spool, err = os.CreateTemp("", "agentsview-artifact-segment-*") + if err != nil { + return fmt.Errorf("creating artifact segment spool: %w", err) + } + if err := spool.Chmod(0o600); err != nil { + return fmt.Errorf("securing artifact segment spool: %w", err) + } + hasher = sha256.New() + segmentBytes = 0 + segmentMessages = 0 + segmentNested = nestedCollectionCounts{} + return nil + } + flush := func() error { + if len(segmentHashes) >= limits.manifestSegments { + return fmt.Errorf("manifest segment reference limit exceeded for %s: limit %d", + sessionID, limits.manifestSegments) + } + digest := hex.EncodeToString(hasher.Sum(nil)) + if _, duplicate := seen[digest]; duplicate { + return fmt.Errorf("generated duplicate segment reference %s", digest) + } + identity, err := NewIdentity(digest, segmentBytes) + if err != nil { + return err + } + ref, err := NewRef(origin, KindSegments, digest+".ndjson") + if err != nil { + return err + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return fmt.Errorf("rewinding artifact segment spool: %w", err) + } + if _, err := store.Create(ctx, ref, identity, + canonicalArtifactMediaType(KindSegments), spool); err != nil { + return fmt.Errorf("creating segment for %s: %w", sessionID, err) + } + cleanupErr := closeAndRemoveExportSpool(spool) + spool = nil + if cleanupErr != nil { + return fmt.Errorf("cleaning artifact segment spool: %w", cleanupErr) + } + seen[digest] = struct{}{} + segmentHashes = append(segmentHashes, digest) + return nil + } + + if err := start(); err != nil { + return nil, err + } + for _, message := range messages { + if err := ctx.Err(); err != nil { + return nil, err + } + messageNested, err := dbMessageNestedCounts(message, limits) + if err != nil { + return nil, err + } + if err := validateMessageFitsSegment(message.Ordinal, messageNested, limits); err != nil { + return nil, err + } + data, err := canonicalJSON(segmentMessageFromDB(message)) + if err != nil { + return nil, fmt.Errorf("encoding message segment: %w", err) + } + if int64(len(data)) > segmentDecodedLimit { + return nil, fmt.Errorf( + "encoded message record at ordinal %d exceeds %d-byte readable limit", + message.Ordinal, segmentDecodedLimit, + ) + } + if segmentBytes > 0 && (segmentBytes+int64(len(data)) > segmentTargetSize || + segmentMessages >= limits.segmentMessages || + exceedsCollectionLimit(segmentNested.toolCalls, + messageNested.toolCalls, limits.segmentToolCalls) || + exceedsCollectionLimit(segmentNested.resultEvents, + messageNested.resultEvents, limits.segmentResultEvents)) { + if err := flush(); err != nil { + return nil, err + } + if err := start(); err != nil { + return nil, err + } + } + if int64(len(data)) > limits.sessionDecodedBytes-decodedBytes { + return nil, fmt.Errorf("session decoded byte limit exceeded for %s: limit %d", + sessionID, limits.sessionDecodedBytes) + } + if _, err := io.MultiWriter(spool, hasher).Write(data); err != nil { + return nil, fmt.Errorf("writing artifact segment spool: %w", err) + } + segmentBytes += int64(len(data)) + decodedBytes += int64(len(data)) + segmentMessages++ + segmentNested.toolCalls += messageNested.toolCalls + segmentNested.resultEvents += messageNested.resultEvents + } + if err := flush(); err != nil { + return nil, err + } + return segmentHashes, nil +} + +func spoolArtifactPublicationMap( + ctx context.Context, + database artifactExportStore, + origin string, +) (_ *os.File, _ string, _ int64, retErr error) { + spool, err := os.CreateTemp("", "agentsview-artifact-map-*") + if err != nil { + return nil, "", 0, fmt.Errorf("creating artifact session map spool: %w", err) + } + failed := true + defer func() { + if failed { + retErr = errors.Join(retErr, exportSpoolCleanup(spool)) + } + }() + if err := exportSpoolChmod(spool); err != nil { + return nil, "", 0, fmt.Errorf("securing artifact session map spool: %w", err) + } + hasher := sha256.New() + writer := io.MultiWriter(spool, hasher) + if _, err := io.WriteString(writer, "{"); err != nil { + return nil, "", 0, err + } + first := true + revision, err := database.StreamArtifactPublications(ctx, origin, func(publication db.ArtifactPublication) error { + if err := ctx.Err(); err != nil { + return err + } + if publication.Origin != origin { + return fmt.Errorf("artifact publication origin mismatch: got %q", publication.Origin) + } + gid, err := json.Marshal(origin + "~" + publication.SessionID) + if err != nil { + return err + } + hash, err := json.Marshal(publication.ManifestHash) + if err != nil { + return err + } + if !first { + if _, err := io.WriteString(writer, ","); err != nil { + return err + } + } + first = false + _, err = writer.Write(append(append(gid, ':'), hash...)) + return err + }) + if err != nil { + return nil, "", 0, fmt.Errorf("streaming artifact session map: %w", err) + } + if _, err := io.WriteString(writer, "}"); err != nil { + return nil, "", 0, err + } + if _, err := hasher.Write([]byte{'\n'}); err != nil { + return nil, "", 0, err + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return nil, "", 0, fmt.Errorf("rewinding artifact session map spool: %w", err) + } + failed = false + return spool, hex.EncodeToString(hasher.Sum(nil)), revision, nil +} + +func spoolArtifactCheckpoint( + ctx context.Context, + mapSpool *os.File, + origin string, + sequence int, +) (_ *os.File, _ Identity, retErr error) { + if _, err := mapSpool.Seek(0, io.SeekStart); err != nil { + return nil, Identity{}, fmt.Errorf("rewinding artifact session map: %w", err) + } + spool, err := os.CreateTemp("", "agentsview-artifact-checkpoint-*") + if err != nil { + return nil, Identity{}, fmt.Errorf("creating artifact checkpoint spool: %w", err) + } + failed := true + defer func() { + if failed { + retErr = errors.Join(retErr, exportSpoolCleanup(spool)) + } + }() + if err := exportSpoolChmod(spool); err != nil { + return nil, Identity{}, fmt.Errorf("securing artifact checkpoint spool: %w", err) + } + hasher := sha256.New() + writer := io.MultiWriter(spool, hasher) + originJSON, err := json.Marshal(origin) + if err != nil { + return nil, Identity{}, err + } + if _, err := fmt.Fprintf(writer, `{"origin":%s,"seq":%d,"sessions":`, originJSON, sequence); err != nil { + return nil, Identity{}, err + } + if _, err := io.Copy(writer, &contextArtifactReader{ctx: ctx, reader: mapSpool}); err != nil { + return nil, Identity{}, fmt.Errorf("copying artifact session map: %w", err) + } + if _, err := io.WriteString(writer, `,"v":1}`+"\n"); err != nil { + return nil, Identity{}, err + } + info, err := spool.Stat() + if err != nil { + return nil, Identity{}, fmt.Errorf("stating artifact checkpoint spool: %w", err) + } + identity, err := NewIdentity(hex.EncodeToString(hasher.Sum(nil)), info.Size()) + if err != nil { + return nil, Identity{}, err + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return nil, Identity{}, fmt.Errorf("rewinding artifact checkpoint spool: %w", err) + } + failed = false + return spool, identity, nil +} + +func closeAndRemoveExportSpool(file *os.File) error { + if file == nil { + return nil + } + name := file.Name() + closeErr := file.Close() + removeErr := os.Remove(name) + if errors.Is(removeErr, fs.ErrNotExist) { + removeErr = nil + } + return errors.Join(closeErr, removeErr) +} + +var ( + exportSpoolChmod = func(file *os.File) error { return file.Chmod(0o600) } + exportSpoolCleanup = closeAndRemoveExportSpool +) + +// statRecordedCheckpoint trusts the store's catalog identity, which is +// established by verified immutable creation and checked again on normal +// reads. Periodic unchanged export must remain constant work; full physical +// verification belongs bootstrap and maintenance. +func statRecordedCheckpoint( + ctx context.Context, + store ArtifactStore, + head db.ArtifactCheckpointHead, +) (bool, error) { + ref, err := NewRef(head.Origin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", head.Sequence)) + if err != nil { + return false, err + } + entry, err := store.Stat(ctx, ref) + if errors.Is(err, ErrArtifactNotFound) { + return false, nil + } + if err != nil { + return false, fmt.Errorf("stating recorded artifact checkpoint: %w", err) + } + if entry.Identity.SHA256 != head.CheckpointSHA256 || entry.Identity.Size != head.CheckpointSize { + quarantineErr := store.Quarantine(ctx, ref, "recorded checkpoint identity mismatch") + return false, quarantineErr + } + return true, nil +} + +func latestValidCheckpointHead( + ctx context.Context, + store ArtifactStore, + origin string, +) (_ db.ArtifactCheckpointHead, _ bool, retErr error) { + var head db.ArtifactCheckpointHead + iterator, err := openStoreEntryIterator(ctx, store, origin, KindCheckpoints) + if err != nil { + return db.ArtifactCheckpointHead{}, false, fmt.Errorf("listing artifact checkpoints: %w", err) + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + for { + entries, nextErr := iterator.Next(ctx, checkpointFloorPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return db.ArtifactCheckpointHead{}, false, fmt.Errorf("listing artifact checkpoints: %w", nextErr) + } + for _, entry := range entries { + sequence, err := checkpointSequence(entry.Ref.Name) + if err != nil || sequence <= head.Sequence { + continue + } + if entry.Identity.Size > checkpointDecodedLimit { + continue + } + _, reader, err := store.Open(ctx, entry.Ref) + if errors.Is(err, ErrArtifactNotFound) || errors.Is(err, ErrArtifactCorrupt) { + continue + } + if err != nil { + return db.ArtifactCheckpointHead{}, false, + fmt.Errorf("opening artifact checkpoint: %w", err) + } + candidate, decodeErr := decodeCanonicalCheckpointHead( + reader, origin, entry.Ref.Name, entry.Identity, + ) + verifyErr := reader.Verify() + closeErr := reader.Close() + if closeErr != nil && !errors.Is(closeErr, ErrArtifactCorrupt) { + return db.ArtifactCheckpointHead{}, false, + fmt.Errorf("closing artifact checkpoint: %w", closeErr) + } + if verifyErr != nil && !errors.Is(verifyErr, ErrArtifactCorrupt) { + return db.ArtifactCheckpointHead{}, false, + fmt.Errorf("verifying artifact checkpoint: %w", verifyErr) + } + if errors.Is(decodeErr, errFutureArtifactVersion) { + return db.ArtifactCheckpointHead{}, false, decodeErr + } + if decodeErr != nil || verifyErr != nil || closeErr != nil { + continue + } + head = candidate + } + if errors.Is(nextErr, io.EOF) { + break + } + } + return head, head.Sequence > 0, nil +} + +func decodeCanonicalCheckpointHead( + reader io.Reader, + origin string, + name string, + identity Identity, +) (db.ArtifactCheckpointHead, error) { + decoder := json.NewDecoder(reader) + decoder.UseNumber() + token, err := decoder.Token() + if err != nil || token != json.Delim('{') { + return db.ArtifactCheckpointHead{}, errors.New("checkpoint is not a JSON object") + } + expectedFields := []string{"origin", "seq", "sessions", "v"} + var sequence int + var mapDigest string + for _, expected := range expectedFields { + token, err := decoder.Token() + if err != nil { + return db.ArtifactCheckpointHead{}, err + } + field, ok := token.(string) + if !ok || field != expected { + return db.ArtifactCheckpointHead{}, fmt.Errorf( + "checkpoint is not canonical: expected field %q", expected, + ) + } + switch field { + case "origin": + var got string + if err := decoder.Decode(&got); err != nil { + return db.ArtifactCheckpointHead{}, err + } + if got != origin { + return db.ArtifactCheckpointHead{}, fmt.Errorf( + "checkpoint origin mismatch for %s: got %q", origin, got, + ) + } + case "seq": + var number json.Number + if err := decoder.Decode(&number); err != nil { + return db.ArtifactCheckpointHead{}, err + } + value, err := strconv.ParseInt(number.String(), 10, 32) + if err != nil || value < 1 { + return db.ArtifactCheckpointHead{}, errors.New("checkpoint sequence is invalid") + } + sequence = int(value) + case "sessions": + mapDigest, err = decodeCanonicalCheckpointSessionMap(decoder, origin) + if err != nil { + return db.ArtifactCheckpointHead{}, err + } + case "v": + var number json.Number + if err := decoder.Decode(&number); err != nil { + return db.ArtifactCheckpointHead{}, err + } + version, err := strconv.Atoi(number.String()) + if err != nil || version < 1 { + return db.ArtifactCheckpointHead{}, errors.New("checkpoint version is unsupported") + } + if version > formatVersion { + return db.ArtifactCheckpointHead{}, fmt.Errorf( + "%w: checkpoint version %d", errFutureArtifactVersion, version, + ) + } + if version != formatVersion { + return db.ArtifactCheckpointHead{}, errors.New("checkpoint version is unsupported") + } + } + } + token, err = decoder.Token() + if err != nil || token != json.Delim('}') { + return db.ArtifactCheckpointHead{}, errors.New("checkpoint object is incomplete") + } + var trailing any + if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) { + if err == nil { + return db.ArtifactCheckpointHead{}, errors.New("checkpoint has trailing JSON") + } + return db.ArtifactCheckpointHead{}, err + } + if fmt.Sprintf("cp-%010d.json", sequence) != name { + return db.ArtifactCheckpointHead{}, fmt.Errorf( + "checkpoint sequence identity mismatch: got %s", name, + ) + } + return db.ArtifactCheckpointHead{ + Origin: origin, Sequence: sequence, + SessionMapSHA256: mapDigest, CheckpointSHA256: identity.SHA256, + CheckpointSize: identity.Size, + }, nil +} + +func decodeCanonicalCheckpointSessionMap( + decoder *json.Decoder, + origin string, +) (string, error) { + token, err := decoder.Token() + if err != nil || token != json.Delim('{') { + return "", errors.New("checkpoint sessions is not an object") + } + hasher := sha256.New() + _, _ = io.WriteString(hasher, "{") + first := true + previous := "" + for decoder.More() { + token, err := decoder.Token() + if err != nil { + return "", err + } + gid, ok := token.(string) + if !ok || gid == "" || !strings.HasPrefix(gid, origin+"~") { + return "", errors.New("checkpoint session identity is invalid") + } + if !first && gid <= previous { + return "", errors.New("checkpoint sessions are not in canonical order") + } + var manifestHash string + if err := decoder.Decode(&manifestHash); err != nil { + return "", err + } + if err := validateHashHex(manifestHash); err != nil { + return "", fmt.Errorf("checkpoint manifest hash is invalid: %w", err) + } + if !first { + _, _ = io.WriteString(hasher, ",") + } + gidJSON, _ := json.Marshal(gid) + hashJSON, _ := json.Marshal(manifestHash) + _, _ = hasher.Write(gidJSON) + _, _ = io.WriteString(hasher, ":") + _, _ = hasher.Write(hashJSON) + first = false + previous = gid + } + token, err = decoder.Token() + if err != nil || token != json.Delim('}') { + return "", errors.New("checkpoint sessions object is incomplete") + } + _, _ = io.WriteString(hasher, "}\n") + return hex.EncodeToString(hasher.Sum(nil)), nil +} + +// Export temporarily preserves the root-based API while canonical publication +// migrates to ArtifactStore. The reference filesystem store is isolated from +// the legacy wire tree, then encoded into that tree for existing transports. +func reserveCheckpointSequenceFromStore( + ctx context.Context, + database artifactCheckpointSequenceDB, + store ArtifactStore, + origin string, +) (_ int, retErr error) { + _, bootstrapped, err := database.GetArtifactCheckpointFloor(ctx, origin) + if err != nil { + return 0, fmt.Errorf("reading checkpoint floor for %s: %w", origin, err) + } + if bootstrapped { + sequence, err := database.ReserveArtifactCheckpointSequence(ctx, origin, 0) + if err != nil { + return 0, fmt.Errorf("reserving checkpoint sequence for %s: %w", origin, err) + } + return sequence, nil + } + observedFloor := 0 + if observer, ok := store.(checkpointFloorStore); ok { + floor, err := observer.checkpointFloor(ctx, origin) + if err != nil { + return 0, fmt.Errorf("listing checkpoint floor for %s: %w", origin, err) + } + observedFloor = floor + } else { + iterator, err := openStoreEntryIterator(ctx, store, origin, KindCheckpoints) + if err != nil { + return 0, fmt.Errorf("listing checkpoint floor for %s: %w", origin, err) + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + for { + entries, nextErr := iterator.Next(ctx, checkpointFloorPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return 0, fmt.Errorf("listing checkpoint floor for %s: %w", origin, nextErr) + } + for _, entry := range entries { + sequence, err := checkpointSequence(entry.Ref.Name) + if err != nil { + continue + } + observedFloor = max(observedFloor, sequence) + } + if errors.Is(nextErr, io.EOF) { + break + } + } + } + sequence, err := database.ReserveArtifactCheckpointSequence(ctx, origin, observedFloor) + if err != nil { + return 0, fmt.Errorf("reserving checkpoint sequence for %s: %w", origin, err) + } + return sequence, nil +} + +func normalizeManifestSessionLocalState(sess *manifestSession) { + // Keep non-content, machine-local state out of the canonical manifest so a + // source-only change to it does not alter the content hash and trigger a + // re-import that clears the importer's local findings. secret_leak_count is + // import-discarded secret state (see rewriteForImport); local_modified_at is + // the local sync watermark, which import ignores (the importer stamps its + // own) -- and a secret rescan bumps both even when no exported message + // content changed. The file_* fields are source-file bookkeeping that + // import clears (see clearImportedSessionSourceState); a touch, move, or + // re-download of the source file changes them without changing any + // exported content. + sess.SecretLeakCount = 0 + sess.LocalModifiedAt = nil + sess.FilePath = nil + sess.FileSize = nil + sess.FileMtime = nil + sess.FileInode = nil + sess.FileDevice = nil + sess.FileHash = nil +} + +type boundedCursorCycleGuard struct { + anchor Cursor + power uint64 + length uint64 +} + +// Observe implements Brent's cycle detector over a deterministic cursor chain +// while retaining constant state regardless of traversal cardinality. +func (g *boundedCursorCycleGuard) Observe(cursor Cursor) bool { + if cursor == "" { + return false + } + if g.anchor == "" { + g.anchor = cursor + g.power = 1 + return false + } + g.length++ + if cursor == g.anchor { + return true + } + if g.length == g.power { + g.anchor = cursor + if g.power <= ^uint64(0)/2 { + g.power *= 2 + } + g.length = 0 + } + return false +} + +func readVerifiedStoreArtifact( + ctx context.Context, + database *db.DB, + store ArtifactStore, + listed Entry, + limit int64, +) ([]byte, error) { + if listed.Identity.Size > limit { + return nil, fmt.Errorf("%w: artifact exceeds %d-byte decoded limit", + ErrArtifactInvalid, limit) + } + if err := validateRefIdentity(listed.Ref, listed.Identity); err != nil { + return nil, fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + entry, reader, err := store.Open(ctx, listed.Ref) + if err != nil { + if errors.Is(err, ErrArtifactNotFound) { + return nil, fmt.Errorf("%w: %s", errIncompleteArtifact, listed.Ref.Name) + } + if errors.Is(err, ErrArtifactCorrupt) { + if qerr := enqueueArtifactRepair(ctx, database, listed); qerr != nil { + return nil, errors.Join(err, qerr) + } + } + return nil, err + } + if entry.Ref != listed.Ref || entry.Identity != listed.Identity { + closeErr := reader.Close() + repairErr := enqueueArtifactRepair(ctx, database, listed) + return nil, errors.Join( + fmt.Errorf("%w: artifact catalog identity changed", ErrArtifactCorrupt), + closeErr, repairErr, + ) + } + data, readErr := io.ReadAll(io.LimitReader(reader, limit+1)) + verifyErr := reader.Verify() + closeErr := reader.Close() + readErr = errors.Join(readErr, verifyErr, closeErr) + if readErr != nil { + if errors.Is(readErr, context.Canceled) || errors.Is(readErr, context.DeadlineExceeded) { + return nil, readErr + } + if qerr := enqueueArtifactRepair(ctx, database, entry); qerr != nil { + return nil, errors.Join(fmt.Errorf("%w: %v", ErrArtifactCorrupt, readErr), qerr) + } + return nil, fmt.Errorf("%w: %v", ErrArtifactCorrupt, readErr) + } + if int64(len(data)) != entry.Identity.Size { + if qerr := enqueueArtifactRepair(ctx, database, entry); qerr != nil { + return nil, errors.Join(fmt.Errorf("%w: artifact size mismatch", ErrArtifactCorrupt), qerr) + } + return nil, fmt.Errorf("%w: artifact size mismatch", ErrArtifactCorrupt) + } + return data, nil +} + +func enqueueArtifactRepair(ctx context.Context, database *db.DB, entry Entry) error { + return database.EnqueueArtifactRepair(ctx, db.ArtifactRepair{ + Origin: entry.Ref.Origin, + Kind: string(entry.Ref.Kind), + Name: entry.Ref.Name, + SHA256: entry.Identity.SHA256, + Size: entry.Identity.Size, + }) +} + +type artifactImportRetryScheduler interface { + RecordChanged(context.Context, Entry) error +} + +// StoreImportCoordinator coalesces dependency arrivals and repairs into one +// explicit store import at the end of a transfer batch. +type StoreImportCoordinator struct { + database *db.DB + store ArtifactStore + localOrigin string + now func() time.Time + + runMu sync.Mutex + mu sync.Mutex + + generation uint64 + completed uint64 +} + +type coordinatedTransportStore struct { + ArtifactStore + database *db.DB + coordinator *StoreImportCoordinator +} + +func newCoordinatedTransportStore( + database *db.DB, + store ArtifactStore, + coordinator *StoreImportCoordinator, +) *coordinatedTransportStore { + return &coordinatedTransportStore{ + ArtifactStore: store, + database: database, + coordinator: coordinator, + } +} + +func (s *coordinatedTransportStore) RecordTransportChanged( + ctx context.Context, entry Entry, +) error { + return s.coordinator.RecordChanged(ctx, entry) +} + +func (s *coordinatedTransportStore) PendingTransportRepair( + ctx context.Context, ref Ref, +) (Entry, bool, error) { + repair, found, err := s.database.ArtifactRepairForRef( + ctx, ref.Origin, string(ref.Kind), ref.Name, + ) + if err != nil || !found { + return Entry{}, found, err + } + identity, err := NewIdentity(repair.SHA256, repair.Size) + if err != nil { + return Entry{}, false, err + } + return Entry{Ref: ref, Identity: identity}, true, nil +} + +func (s *coordinatedTransportStore) RepairTransportArtifact( + ctx context.Context, entry Entry, trusted io.Reader, +) error { + return s.RepairContent(ctx, entry.Identity, trusted) +} + +func (s *coordinatedTransportStore) AcknowledgeTransportRepair( + ctx context.Context, entry Entry, +) error { + return s.database.AcknowledgeArtifactRepair(ctx, db.ArtifactRepair{ + Origin: entry.Ref.Origin, + Kind: string(entry.Ref.Kind), + Name: entry.Ref.Name, + SHA256: entry.Identity.SHA256, + Size: entry.Identity.Size, + }) +} + +func NewStoreImportCoordinator( + database *db.DB, store ArtifactStore, localOrigin string, +) *StoreImportCoordinator { + return &StoreImportCoordinator{ + database: database, store: store, localOrigin: localOrigin, + generation: 1, + } +} + +// requestDrain marks the current transfer batch for another import. +func (c *StoreImportCoordinator) requestDrain() error { + if c == nil { + return errors.New("artifact import coordinator is required") + } + c.mu.Lock() + c.generation++ + c.mu.Unlock() + return nil +} + +// Finalize consumes one coalesced retry signal. A transient import failure +// retains the signal for a later finalize attempt. +func (c *StoreImportCoordinator) Finalize(ctx context.Context) (ImportResult, error) { + if c == nil { + return ImportResult{}, errors.New("artifact import coordinator is required") + } + c.runMu.Lock() + defer c.runMu.Unlock() + + c.mu.Lock() + generation := c.generation + if c.completed >= generation { + c.mu.Unlock() + return ImportResult{}, nil + } + c.mu.Unlock() + + result, err := c.drainQueuedImports(ctx) + if err == nil { + c.mu.Lock() + c.completed = generation + c.mu.Unlock() + } + return result, err +} + +type checkpointClosureOutcome uint8 + +const ( + checkpointClosureComplete checkpointClosureOutcome = iota + checkpointClosureDeferred + checkpointClosureCurrentInvalid +) + +func inspectCheckpointClosureFromStore( + ctx context.Context, + database *db.DB, + store ArtifactStore, + origin string, + cp checkpoint, +) (checkpointClosureOutcome, error) { + keys := make([]string, 0, len(cp.Sessions)) + for gid := range cp.Sessions { + keys = append(keys, gid) + } + sort.Strings(keys) + importStates, err := loadImportStates(database, origin, keys) + if err != nil { + return checkpointClosureComplete, fmt.Errorf("reading import state for %s: %w", origin, err) + } + for _, gid := range keys { + if err := ctx.Err(); err != nil { + return checkpointClosureComplete, err + } + manifestHash := cp.Sessions[gid] + if importStates[importStateKey(origin, gid)] == manifestHash { + continue + } + _, _, outcome, err := readStoreSession( + ctx, database, store, origin, gid, manifestHash, + ) + if err != nil { + return checkpointClosureComplete, err + } + if outcome != checkpointClosureComplete { + return outcome, nil + } + } + return checkpointClosureComplete, nil +} + +// RepairArtifactFromTrustedPeer repairs one queued physical identity from a +// trusted peer stream and acknowledges the exact SQLite claim only afterward. +func RepairArtifactFromTrustedPeer( + ctx context.Context, + database *db.DB, + store ArtifactStore, + repair db.ArtifactRepair, + trusted io.Reader, + retry artifactImportRetryScheduler, +) error { + if database == nil || store == nil || trusted == nil || retry == nil || isTypedNil(retry) { + return errors.New("artifact repair requires database, store, trusted content, and retry coordinator") + } + ref, err := NewRef(repair.Origin, Kind(repair.Kind), repair.Name) + if err != nil { + return err + } + identity, err := NewIdentity(repair.SHA256, repair.Size) + if err != nil { + return err + } + if err := validateRefIdentity(ref, identity); err != nil { + return err + } + if err := store.RepairContent(ctx, identity, trusted); err != nil { + return fmt.Errorf("repairing artifact content: %w", err) + } + if err := retry.RecordChanged(ctx, Entry{Ref: ref, Identity: identity}); err != nil { + return fmt.Errorf("scheduling artifact import retry: %w", err) + } + if err := database.AcknowledgeArtifactRepair(ctx, repair); err != nil { + return fmt.Errorf("acknowledging artifact repair: %w", err) + } + return nil +} + +func isTypedNil(value any) bool { + v := reflect.ValueOf(value) + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: + return v.IsNil() + default: + return false + } +} + +func importCheckpointFromStore( + ctx context.Context, + database *db.DB, + store ArtifactStore, + origin string, + cp checkpoint, +) (ImportResult, error) { + keys := make([]string, 0, len(cp.Sessions)) + for gid := range cp.Sessions { + keys = append(keys, gid) + } + sort.Strings(keys) + importStates, err := loadImportStates(database, origin, keys) + if err != nil { + return ImportResult{}, fmt.Errorf("reading import state for %s: %w", origin, err) + } + result := ImportResult{} + complete := true + for _, gid := range keys { + if err := ctx.Err(); err != nil { + return result, err + } + manifestHash := cp.Sessions[gid] + stateKey := importStateKey(origin, gid) + if importStates[stateKey] == manifestHash { + continue + } + m, msgs, outcome, err := readStoreSession( + ctx, database, store, origin, gid, manifestHash, + ) + if err != nil { + return result, err + } + if outcome != checkpointClosureComplete { + result.Deferred++ + complete = false + continue + } + write := rewriteForImport(m, msgs) + writeResult, err := database.WriteSessionBatchAtomic([]db.SessionBatchWrite{write}) + if err != nil { + if errors.Is(err, db.ErrSessionTrashed) { + continue + } + if errors.Is(err, db.ErrSessionExcluded) { + complete = false + continue + } + return result, fmt.Errorf("importing artifact session %s: %w", gid, err) + } + result.Sessions += writeResult.WrittenSessions + result.Messages += writeResult.WrittenMessages + if _, err := database.ReapplyMetadataReplayState(ctx, gid, write.Session.ID); err != nil { + return result, fmt.Errorf("reapplying metadata after importing artifact session %s: %w", gid, err) + } + if err := database.SetSyncState(stateKey, manifestHash); err != nil { + return result, fmt.Errorf("writing import state for %s: %w", gid, err) + } + } + if complete { + err := database.RecordArtifactCheckpointLanding(ctx, db.ArtifactCheckpointLanding{ + Origin: origin, Sequence: cp.Sequence, + }, cp.Sessions) + if err != nil { + return result, err + } + } + return result, nil +} + +func readStoreSession( + ctx context.Context, + database *db.DB, + store ArtifactStore, + origin, gid, manifestHash string, +) (manifest, []db.Message, checkpointClosureOutcome, error) { + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + if err != nil { + return manifest{}, nil, checkpointClosureComplete, err + } + manifestEntry, err := store.Stat(ctx, manifestRef) + if errors.Is(err, ErrArtifactNotFound) { + return manifest{}, nil, checkpointClosureDeferred, nil + } + if err != nil { + return manifest{}, nil, checkpointClosureComplete, err + } + data, err := readVerifiedStoreArtifact(ctx, database, store, manifestEntry, manifestDecodedLimit) + if errors.Is(err, errIncompleteArtifact) { + return manifest{}, nil, checkpointClosureDeferred, nil + } + if err != nil { + if errors.Is(err, ErrArtifactInvalid) { + if qerr := store.Quarantine(ctx, manifestRef, err.Error()); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + return manifest{}, nil, checkpointClosureComplete, err + } + m, err := decodeManifestWithLimits(data, productionArtifactLimits()) + if err != nil { + if qerr := store.Quarantine(ctx, manifestRef, "manifest JSON is invalid"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + if m.Version > formatVersion { + return manifest{}, nil, checkpointClosureDeferred, nil + } + canonical, canonicalErr := canonicalJSON(m) + if canonicalErr != nil || !bytes.Equal(canonical, data) { + if qerr := store.Quarantine(ctx, manifestRef, "manifest JSON is not canonical"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(canonicalErr, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + if err := validateManifest(m, origin, gid); err != nil { + if errors.Is(err, errFutureArtifactVersion) { + return manifest{}, nil, checkpointClosureDeferred, nil + } + if qerr := store.Quarantine(ctx, manifestRef, err.Error()); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + var messages []db.Message + var decodedBytes int64 + totalNested := nestedCollectionCounts{} + for _, segmentHash := range m.Segments { + segmentRef, err := NewRef(origin, KindSegments, segmentHash+".ndjson") + if err != nil { + if errors.Is(err, ErrArtifactInvalid) { + if qerr := store.Quarantine(ctx, segmentRef, err.Error()); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + return manifest{}, nil, checkpointClosureComplete, err + } + segmentEntry, err := store.Stat(ctx, segmentRef) + if errors.Is(err, ErrArtifactNotFound) { + return manifest{}, nil, checkpointClosureDeferred, nil + } + if err != nil { + return manifest{}, nil, checkpointClosureComplete, err + } + segmentData, err := readVerifiedStoreArtifact( + ctx, database, store, segmentEntry, segmentDecodedLimit, + ) + if errors.Is(err, errIncompleteArtifact) { + return manifest{}, nil, checkpointClosureDeferred, nil + } + if err != nil { + if errors.Is(err, ErrArtifactInvalid) { + if qerr := store.Quarantine(ctx, segmentRef, err.Error()); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + return manifest{}, nil, checkpointClosureComplete, err + } + segmentMessages, err := decodeSegment(segmentData) + if errors.Is(err, errFutureArtifactVersion) { + return manifest{}, nil, checkpointClosureDeferred, nil + } + if err != nil { + if qerr := store.Quarantine(ctx, segmentRef, "segment NDJSON is invalid"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + canonical, canonicalErr := encodeSegment(segmentMessages) + if canonicalErr != nil || !bytes.Equal(canonical, segmentData) { + if qerr := store.Quarantine(ctx, segmentRef, "segment NDJSON is not canonical"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(canonicalErr, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + if int64(len(segmentData)) > maxSessionDecodedBytes-decodedBytes { + if qerr := store.Quarantine(ctx, manifestRef, "session decoded byte limit exceeded"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, qerr + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + for _, message := range segmentMessages { + counts, err := dbMessageNestedCounts(message, productionArtifactLimits()) + if err != nil { + if qerr := store.Quarantine(ctx, manifestRef, err.Error()); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, errors.Join(err, qerr) + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + if exceedsCollectionLimit(totalNested.toolCalls, counts.toolCalls, maxSessionToolCalls) || + exceedsCollectionLimit(totalNested.resultEvents, counts.resultEvents, maxSessionResultEvents) { + if qerr := store.Quarantine(ctx, manifestRef, "session nested collection limit exceeded"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, qerr + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + totalNested.toolCalls += counts.toolCalls + totalNested.resultEvents += counts.resultEvents + } + decodedBytes += int64(len(segmentData)) + messages = append(messages, segmentMessages...) + } + if len(messages) > maxSessionMessages { + if qerr := store.Quarantine(ctx, manifestRef, "session message limit exceeded"); qerr != nil { + return manifest{}, nil, checkpointClosureComplete, qerr + } + return manifest{}, nil, checkpointClosureCurrentInvalid, nil + } + return m, messages, checkpointClosureComplete, nil +} + +func loadImportStates( + database syncStateValueReader, origin string, gids []string, +) (map[string]string, error) { + if len(gids) == 0 { + return map[string]string{}, nil + } + keys := make([]string, len(gids)) + for i, gid := range gids { + keys[i] = importStateKey(origin, gid) + } + return database.SyncStateValues(keys) +} + +func importStateKey(origin, gid string) string { + return importStatePrefix + origin + ":" + gid +} + +func validateCheckpoint(cp *checkpoint, origin string) error { + if cp.Version > formatVersion { + return fmt.Errorf( + "%w: checkpoint for %s has artifact version %d", + errFutureArtifactVersion, origin, cp.Version, + ) + } + if cp.Version != formatVersion { + return fmt.Errorf( + "checkpoint for %s has unsupported artifact version %d", + origin, cp.Version, + ) + } + if cp.Origin != origin { + return fmt.Errorf( + "checkpoint origin mismatch for %s: got %q", + origin, cp.Origin, + ) + } + return validateCheckpointReferences(cp, origin) +} + +func validateCheckpointReferences(cp *checkpoint, origin string) error { + for gid, manifestHash := range cp.Sessions { + if gid == "" { + return fmt.Errorf("checkpoint for %s contains empty session id", origin) + } + if !strings.HasPrefix(gid, origin+"~") { + return fmt.Errorf( + "checkpoint session %s does not belong to origin %s", + gid, origin, + ) + } + if strings.TrimSpace(manifestHash) == "" { + return fmt.Errorf("checkpoint session %s has empty manifest hash", gid) + } + if err := validateHashHex(manifestHash); err != nil { + return fmt.Errorf("checkpoint session %s has invalid manifest hash: %w", gid, err) + } + } + return nil +} + +func validateManifest(m manifest, origin, gid string) error { + if m.Version > formatVersion { + return fmt.Errorf( + "%w: manifest %s has artifact version %d", + errFutureArtifactVersion, gid, m.Version, + ) + } + if m.Version != formatVersion { + return fmt.Errorf( + "manifest %s has unsupported artifact version %d", + gid, m.Version, + ) + } + if m.Origin != origin { + return fmt.Errorf( + "manifest origin mismatch for %s: got %q", + gid, m.Origin, + ) + } + if m.NativeSessionID == "" { + return fmt.Errorf("manifest %s has empty native session id", gid) + } + expectedGID := origin + "~" + m.NativeSessionID + if gid != expectedGID { + return fmt.Errorf( + "manifest session id mismatch: checkpoint has %s, manifest has %s", + gid, expectedGID, + ) + } + if m.Session.ID != m.NativeSessionID { + return fmt.Errorf( + "manifest %s session row id mismatch: got %q", + gid, m.Session.ID, + ) + } + if m.Session.Machine != origin { + return fmt.Errorf( + "manifest %s session row machine mismatch: got %q", + gid, m.Session.Machine, + ) + } + if len(m.Segments) == 0 { + return fmt.Errorf("manifest %s has no message segments", gid) + } + if err := validateManifestReferences(m); err != nil { + return err + } + return nil +} + +func validateManifestReferences(m manifest) error { + return validateManifestReferencesWithLimits(m, productionArtifactLimits()) +} + +func validateManifestReferencesWithLimits(m manifest, limits artifactLimits) error { + if len(m.Segments) > limits.manifestSegments { + return fmt.Errorf( + "manifest segment reference limit exceeded: got %d, limit %d", + len(m.Segments), limits.manifestSegments, + ) + } + seen := make(map[string]struct{}, len(m.Segments)) + for _, segmentHash := range m.Segments { + if err := validateHashHex(segmentHash); err != nil { + return fmt.Errorf("manifest segment has invalid hash: %w", err) + } + if _, ok := seen[segmentHash]; ok { + return fmt.Errorf("manifest has duplicate segment reference %s", segmentHash) + } + seen[segmentHash] = struct{}{} + } + if len(m.UsageEvents) > limits.manifestUsageEvents { + return fmt.Errorf( + "manifest usage event limit exceeded: got %d, limit %d", + len(m.UsageEvents), limits.manifestUsageEvents, + ) + } + if m.RawSource != nil && m.RawSource.Hash != "" { + if err := validateHashHex(m.RawSource.Hash); err != nil { + return fmt.Errorf("manifest raw source has invalid hash: %w", err) + } + } + return nil +} + +func rewriteForImport(m manifest, msgs []db.Message) db.SessionBatchWrite { + importedID := m.Origin + "~" + m.NativeSessionID + sess := m.Session.dbSession() + sess.ID = importedID + sess.Machine = m.Origin + sess.SessionName = m.SessionName + clearImportedSessionSourceState(&sess) + // Restore signal state dropped from the Session JSON; signalsFromSession + // reads these fields below to persist the imported session's signal columns. + sess.HasToolCalls = m.SessionHasToolCalls + sess.HasContextData = m.SessionHasContextData + sess.ApplyQualitySignals(m.SessionQualitySignals.dbQualitySignals()) + // Secret findings are not carried in the manifest, so an imported session has + // no finding rows. Treat it as unscanned rather than trusting the source scan: + // clear the rules version (json:"-", so already absent) and the leak count + // (carried in the Session JSON) so the count stays consistent with the zero + // findings and `secrets scan --backfill` rescans it with local rules. Stamping + // it scanned-at-source-version would make backfill (secrets_rules_version != + // current) skip a secret-bearing session, leaving no revealable findings. + sess.SecretsRulesVersion = "" + sess.SecretLeakCount = 0 + sess.SourceSessionID = prefixImportedSessionID(m.Origin, sess.SourceSessionID) + if sess.ParentSessionID != nil { + prefixed := prefixImportedSessionID(m.Origin, *sess.ParentSessionID) + sess.ParentSessionID = &prefixed + } + for i := range msgs { + msgs[i].ID = 0 + msgs[i].SessionID = importedID + for j := range msgs[i].ToolCalls { + msgs[i].ToolCalls[j].MessageID = 0 + msgs[i].ToolCalls[j].SessionID = importedID + msgs[i].ToolCalls[j].SubagentSessionID = prefixImportedSessionID( + m.Origin, + msgs[i].ToolCalls[j].SubagentSessionID, + ) + for k := range msgs[i].ToolCalls[j].ResultEvents { + ev := &msgs[i].ToolCalls[j].ResultEvents[k] + ev.SubagentSessionID = prefixImportedSessionID(m.Origin, ev.SubagentSessionID) + } + } + } + usageEvents := dbUsageEvents(m.UsageEvents, importedID) + return db.SessionBatchWrite{ + Session: sess, + Messages: msgs, + UsageEvents: usageEvents, + Signals: signalsFromSession(sess), + DataVersion: m.DataVersion, + ReplaceMessages: true, + } +} + +func clearImportedSessionSourceState(sess *db.Session) { + sess.FilePath = nil + sess.FileSize = nil + sess.FileMtime = nil + sess.NextOrdinal = 0 + sess.LastEntryUUID = nil + sess.FileInode = nil + sess.FileDevice = nil + sess.FileHash = nil +} + +func prefixImportedSessionID(origin, id string) string { + if id == "" || strings.Contains(id, "~") { + return id + } + return origin + "~" + id +} + +func signalsFromSession(s db.Session) db.SessionSignalUpdate { + update := db.SessionSignalUpdate{ + ToolFailureSignalCount: s.ToolFailureSignalCount, + ToolRetryCount: s.ToolRetryCount, + EditChurnCount: s.EditChurnCount, + ConsecutiveFailureMax: s.ConsecutiveFailureMax, + Outcome: s.Outcome, + OutcomeConfidence: s.OutcomeConfidence, + EndedWithRole: s.EndedWithRole, + FinalFailureStreak: s.FinalFailureStreak, + SignalsPendingSince: s.SignalsPendingSince, + CompactionCount: s.CompactionCount, + MidTaskCompactionCount: s.MidTaskCompactionCount, + ContextPressureMax: s.ContextPressureMax, + HealthScore: s.HealthScore, + HealthGrade: s.HealthGrade, + HasToolCalls: s.HasToolCalls, + HasContextData: s.HasContextData, + SecretLeakCount: s.SecretLeakCount, + SecretsRulesVersion: s.SecretsRulesVersion, + } + if qs := s.StoredQualitySignals(); qs != nil { + update.QualitySignals = *qs + } + return update +} + +func decodeManifestWithLimits(data []byte, limits artifactLimits) (manifest, error) { + var envelope struct { + Version int `json:"v"` + Origin string `json:"origin"` + Segments json.RawMessage `json:"segments"` + UsageEvents json.RawMessage `json:"usage_events"` + } + if err := json.Unmarshal(data, &envelope); err != nil { + return manifest{}, err + } + // Future manifests are retained for forward compatibility. Reading only + // their scalar header avoids allocating collections whose schema this + // version does not understand. + if envelope.Version > formatVersion { + return manifest{Version: envelope.Version, Origin: envelope.Origin}, nil + } + if err := preflightManifestCollections( + envelope.Segments, envelope.UsageEvents, limits, + ); err != nil { + return manifest{}, err + } + var m manifest + if err := json.Unmarshal(data, &m); err != nil { + return manifest{}, err + } + return m, nil +} + +func preflightManifestCollections( + segments, usageEvents json.RawMessage, + limits artifactLimits, +) error { + if err := preflightSegmentReferences(segments, limits.manifestSegments); err != nil { + return err + } + return preflightJSONArrayCount( + usageEvents, "manifest usage event", limits.manifestUsageEvents, + ) +} + +func preflightSegmentReferences(data json.RawMessage, limit int) error { + trimmed := bytes.TrimSpace(data) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil + } + dec := json.NewDecoder(bytes.NewReader(trimmed)) + token, err := dec.Token() + if err != nil { + return err + } + if token != json.Delim('[') { + return errors.New("manifest segments must be an array") + } + seen := make(map[string]struct{}, min(limit, 16)) + count := 0 + for dec.More() { + if count >= limit { + return fmt.Errorf("manifest segment reference limit exceeded: limit %d", limit) + } + var hash string + if err := dec.Decode(&hash); err != nil { + return fmt.Errorf("decoding manifest segment reference: %w", err) + } + if _, ok := seen[hash]; ok { + return fmt.Errorf("manifest has duplicate segment reference %s", hash) + } + seen[hash] = struct{}{} + count++ + } + _, err = dec.Token() + return err +} + +func preflightJSONArrayCount(data json.RawMessage, name string, limit int) error { + _, err := countJSONArrayElements(data, name, limit) + return err +} + +func countJSONArrayElements( + data json.RawMessage, + name string, + limit int, +) (int, error) { + trimmed := bytes.TrimSpace(data) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return 0, nil + } + dec := json.NewDecoder(bytes.NewReader(trimmed)) + token, err := dec.Token() + if err != nil { + return 0, err + } + if token != json.Delim('[') { + return 0, fmt.Errorf("%ss must be an array", name) + } + count := 0 + for dec.More() { + if count >= limit { + return 0, fmt.Errorf("%s limit exceeded: limit %d", name, limit) + } + var value json.RawMessage + if err := dec.Decode(&value); err != nil { + return 0, fmt.Errorf("decoding %s: %w", name, err) + } + count++ + } + if _, err := dec.Token(); err != nil { + return 0, err + } + return count, nil +} + +func canonicalMessages(msgs []db.Message) []db.Message { + out := make([]db.Message, len(msgs)) + for i, msg := range msgs { + msg.ID = 0 + msg.SessionID = "" + if len(msg.ToolCalls) > 0 { + calls := make([]db.ToolCall, len(msg.ToolCalls)) + copy(calls, msg.ToolCalls) + for j := range calls { + calls[j].MessageID = 0 + calls[j].SessionID = "" + } + msg.ToolCalls = calls + } + out[i] = msg + } + return out +} + +func canonicalUsageEvents(events []db.UsageEvent) []artifactUsageEvent { + out := make([]artifactUsageEvent, len(events)) + for i, ev := range events { + out[i] = artifactUsageEvent{ + MessageOrdinal: ev.MessageOrdinal, + Source: ev.Source, + Model: ev.Model, + InputTokens: ev.InputTokens, + OutputTokens: ev.OutputTokens, + CacheCreationInputTokens: ev.CacheCreationInputTokens, + CacheReadInputTokens: ev.CacheReadInputTokens, + ReasoningTokens: ev.ReasoningTokens, + CostUSD: ev.CostUSD, + CostStatus: ev.CostStatus, + CostSource: ev.CostSource, + OccurredAt: ev.OccurredAt, + DedupKey: ev.DedupKey, + } + } + return out +} + +func validateExportNestedCollections(msgs []db.Message, limits artifactLimits) error { + total := nestedCollectionCounts{} + for _, msg := range msgs { + messageNested, err := dbMessageNestedCounts(msg, limits) + if err != nil { + return err + } + if err := validateMessageFitsSegment(msg.Ordinal, messageNested, limits); err != nil { + return err + } + if exceedsCollectionLimit( + total.toolCalls, messageNested.toolCalls, limits.sessionToolCalls, + ) { + return fmt.Errorf( + "session tool call limit exceeded at message ordinal %d: limit %d", + msg.Ordinal, limits.sessionToolCalls, + ) + } + if exceedsCollectionLimit( + total.resultEvents, messageNested.resultEvents, limits.sessionResultEvents, + ) { + return fmt.Errorf( + "session result event limit exceeded at message ordinal %d: limit %d", + msg.Ordinal, limits.sessionResultEvents, + ) + } + total.toolCalls += messageNested.toolCalls + total.resultEvents += messageNested.resultEvents + } + return nil +} + +func dbMessageNestedCounts( + msg db.Message, + limits artifactLimits, +) (nestedCollectionCounts, error) { + if len(msg.ToolCalls) > limits.messageToolCalls { + return nestedCollectionCounts{}, fmt.Errorf( + "tool call limit exceeded for message ordinal %d: got %d, limit %d", + msg.Ordinal, len(msg.ToolCalls), limits.messageToolCalls, + ) + } + counts := nestedCollectionCounts{toolCalls: len(msg.ToolCalls)} + for toolIndex, call := range msg.ToolCalls { + if len(call.ResultEvents) > limits.toolResultEvents { + return nestedCollectionCounts{}, fmt.Errorf( + "result event limit exceeded for tool call %d in message ordinal %d: got %d, limit %d", + toolIndex, msg.Ordinal, len(call.ResultEvents), limits.toolResultEvents, + ) + } + counts.resultEvents += len(call.ResultEvents) + } + return counts, nil +} + +func validateMessageFitsSegment( + ordinal int, + counts nestedCollectionCounts, + limits artifactLimits, +) error { + if counts.toolCalls > limits.segmentToolCalls { + return fmt.Errorf( + "message ordinal %d cannot fit in one segment: got %d tool calls, segment limit %d", + ordinal, counts.toolCalls, limits.segmentToolCalls, + ) + } + if counts.resultEvents > limits.segmentResultEvents { + return fmt.Errorf( + "message ordinal %d cannot fit in one segment: got %d result events, segment limit %d", + ordinal, counts.resultEvents, limits.segmentResultEvents, + ) + } + return nil +} + +func encodeSegment(msgs []db.Message) ([]byte, error) { + var buf bytes.Buffer + for _, msg := range msgs { + data, err := canonicalJSON(segmentMessageFromDB(msg)) + if err != nil { + return nil, fmt.Errorf("encoding message segment: %w", err) + } + buf.Write(data) + } + return buf.Bytes(), nil +} + +func decodeSegment(data []byte) ([]db.Message, error) { + return decodeSegmentWithLimits(data, productionArtifactLimits()) +} + +func decodeSegmentWithLimits(data []byte, limits artifactLimits) ([]db.Message, error) { + preflight, err := preflightSegmentData(data, limits) + if err != nil { + return nil, err + } + return decodePreflightedSegment(preflight) +} + +func decodePreflightedSegment(preflight segmentPreflight) ([]db.Message, error) { + msgs := make([]db.Message, 0, len(preflight.records)) + for _, line := range preflight.records { + var record segmentMessage + if err := json.Unmarshal(line, &record); err != nil { + return nil, fmt.Errorf("decoding message segment: %w", err) + } + msgs = append(msgs, record.dbMessage()) + } + return msgs, nil +} + +func segmentRecords(data []byte, limit int) ([][]byte, error) { + capacity := min(max(limit, 0), 64) + records := make([][]byte, 0, capacity) + remaining := data + lineNumber := 0 + for len(remaining) > 0 { + lineNumber++ + newline := bytes.IndexByte(remaining, '\n') + line := remaining + if newline >= 0 { + line = remaining[:newline] + remaining = remaining[newline+1:] + } else { + remaining = nil + } + if len(records) >= limit { + return nil, fmt.Errorf( + "message record limit exceeded: limit %d per segment", limit, + ) + } + if len(bytes.TrimSpace(line)) == 0 { + return nil, fmt.Errorf("blank message record at line %d", lineNumber) + } + records = append(records, line) + } + return records, nil +} + +func preflightSegmentData(data []byte, limits artifactLimits) (segmentPreflight, error) { + records, err := segmentRecords(data, limits.segmentMessages) + if err != nil { + return segmentPreflight{}, err + } + preflight := segmentPreflight{records: records} + for _, line := range records { + var header struct { + Version int `json:"v"` + } + if err := json.Unmarshal(line, &header); err != nil { + return segmentPreflight{}, fmt.Errorf("decoding message segment header: %w", err) + } + if header.Version > formatVersion { + return segmentPreflight{}, fmt.Errorf( + "%w: message segment has artifact version %d", + errFutureArtifactVersion, header.Version, + ) + } + if header.Version != formatVersion { + return segmentPreflight{}, fmt.Errorf( + "message segment has unsupported artifact version %d", + header.Version, + ) + } + messageNested, err := preflightMessageNestedCollections(line, limits) + if err != nil { + return segmentPreflight{}, err + } + if exceedsCollectionLimit( + preflight.nested.toolCalls, + messageNested.toolCalls, + limits.segmentToolCalls, + ) { + return segmentPreflight{}, fmt.Errorf( + "segment tool call limit exceeded: limit %d", limits.segmentToolCalls, + ) + } + if exceedsCollectionLimit( + preflight.nested.resultEvents, + messageNested.resultEvents, + limits.segmentResultEvents, + ) { + return segmentPreflight{}, fmt.Errorf( + "segment result event limit exceeded: limit %d", + limits.segmentResultEvents, + ) + } + preflight.nested.toolCalls += messageNested.toolCalls + preflight.nested.resultEvents += messageNested.resultEvents + } + return preflight, nil +} + +func preflightMessageNestedCollections( + line []byte, + limits artifactLimits, +) (nestedCollectionCounts, error) { + var envelope struct { + Ordinal int `json:"ordinal"` + ToolCalls json.RawMessage `json:"tool_calls"` + } + if err := json.Unmarshal(line, &envelope); err != nil { + return nestedCollectionCounts{}, fmt.Errorf( + "decoding message segment collections: %w", err, + ) + } + return preflightToolCallCollections(envelope.ToolCalls, envelope.Ordinal, limits) +} + +func preflightToolCallCollections( + data json.RawMessage, + ordinal int, + limits artifactLimits, +) (nestedCollectionCounts, error) { + trimmed := bytes.TrimSpace(data) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nestedCollectionCounts{}, nil + } + dec := json.NewDecoder(bytes.NewReader(trimmed)) + token, err := dec.Token() + if err != nil { + return nestedCollectionCounts{}, err + } + if token != json.Delim('[') { + return nestedCollectionCounts{}, errors.New("message tool_calls must be an array") + } + counts := nestedCollectionCounts{} + for dec.More() { + if counts.toolCalls >= limits.messageToolCalls { + return nestedCollectionCounts{}, fmt.Errorf( + "tool call limit exceeded for message ordinal %d: limit %d per message", + ordinal, limits.messageToolCalls, + ) + } + var toolEnvelope struct { + ResultEvents json.RawMessage `json:"result_events"` + } + if err := dec.Decode(&toolEnvelope); err != nil { + return nestedCollectionCounts{}, fmt.Errorf( + "decoding tool call %d in message ordinal %d: %w", + counts.toolCalls, ordinal, err, + ) + } + resultEvents, err := countJSONArrayElements( + toolEnvelope.ResultEvents, "result event", limits.toolResultEvents, + ) + if err != nil { + return nestedCollectionCounts{}, fmt.Errorf( + "preflighting tool call %d in message ordinal %d: %w", + counts.toolCalls, ordinal, err, + ) + } + counts.toolCalls++ + counts.resultEvents += resultEvents + } + if _, err := dec.Token(); err != nil { + return nestedCollectionCounts{}, err + } + return counts, nil +} + +func dbUsageEvents(events []artifactUsageEvent, sessionID string) []db.UsageEvent { + out := make([]db.UsageEvent, len(events)) + for i, ev := range events { + out[i] = db.UsageEvent{ + SessionID: sessionID, + MessageOrdinal: ev.MessageOrdinal, + Source: ev.Source, + Model: ev.Model, + InputTokens: ev.InputTokens, + OutputTokens: ev.OutputTokens, + CacheCreationInputTokens: ev.CacheCreationInputTokens, + CacheReadInputTokens: ev.CacheReadInputTokens, + ReasoningTokens: ev.ReasoningTokens, + CostUSD: ev.CostUSD, + CostStatus: ev.CostStatus, + CostSource: ev.CostSource, + OccurredAt: ev.OccurredAt, + DedupKey: ev.DedupKey, + } + } + return out +} + +func segmentMessageFromDB(msg db.Message) segmentMessage { + record := segmentMessage{ + Version: formatVersion, + Ordinal: msg.Ordinal, + Role: msg.Role, + Content: msg.Content, + ThinkingText: msg.ThinkingText, + Timestamp: msg.Timestamp, + HasThinking: msg.HasThinking, + HasToolUse: msg.HasToolUse, + ContentLength: msg.ContentLength, + Model: msg.Model, + TokenUsage: msg.TokenUsage, + ContextTokens: msg.ContextTokens, + OutputTokens: msg.OutputTokens, + HasContextTokens: msg.HasContextTokens, + HasOutputTokens: msg.HasOutputTokens, + ClaudeMessageID: msg.ClaudeMessageID, + ClaudeRequestID: msg.ClaudeRequestID, + IsSystem: msg.IsSystem, + SourceType: msg.SourceType, + SourceSubtype: msg.SourceSubtype, + SourceUUID: msg.SourceUUID, + SourceParentUUID: msg.SourceParentUUID, + IsSidechain: msg.IsSidechain, + IsCompactBoundary: msg.IsCompactBoundary, + } + if len(msg.ToolCalls) > 0 { + record.ToolCalls = make([]segmentToolCall, len(msg.ToolCalls)) + for i, call := range msg.ToolCalls { + record.ToolCalls[i] = segmentToolCall{ + CallIndex: i, + ToolName: call.ToolName, + Category: call.Category, + ToolUseID: call.ToolUseID, + InputJSON: call.InputJSON, + FilePath: call.FilePath, + SkillName: call.SkillName, + ResultContentLength: call.ResultContentLength, + ResultContent: call.ResultContent, + SubagentSessionID: call.SubagentSessionID, + } + if len(call.ResultEvents) > 0 { + record.ToolCalls[i].ResultEvents = make([]segmentResultEvent, len(call.ResultEvents)) + for j, ev := range call.ResultEvents { + record.ToolCalls[i].ResultEvents[j] = segmentResultEvent{ + ToolUseID: ev.ToolUseID, + AgentID: ev.AgentID, + SubagentSessionID: ev.SubagentSessionID, + Source: ev.Source, + Status: ev.Status, + Content: ev.Content, + ContentLength: ev.ContentLength, + Timestamp: ev.Timestamp, + EventIndex: ev.EventIndex, + } + } + } + } + } + return record +} + +func (m segmentMessage) dbMessage() db.Message { + msg := db.Message{ + Ordinal: m.Ordinal, + Role: m.Role, + Content: m.Content, + ThinkingText: m.ThinkingText, + Timestamp: m.Timestamp, + HasThinking: m.HasThinking, + HasToolUse: m.HasToolUse, + ContentLength: m.ContentLength, + Model: m.Model, + TokenUsage: m.TokenUsage, + ContextTokens: m.ContextTokens, + OutputTokens: m.OutputTokens, + HasContextTokens: m.HasContextTokens, + HasOutputTokens: m.HasOutputTokens, + ClaudeMessageID: m.ClaudeMessageID, + ClaudeRequestID: m.ClaudeRequestID, + IsSystem: m.IsSystem, + SourceType: m.SourceType, + SourceSubtype: m.SourceSubtype, + SourceUUID: m.SourceUUID, + SourceParentUUID: m.SourceParentUUID, + IsSidechain: m.IsSidechain, + IsCompactBoundary: m.IsCompactBoundary, + } + if len(m.ToolCalls) > 0 { + msg.ToolCalls = make([]db.ToolCall, len(m.ToolCalls)) + for i, call := range m.ToolCalls { + msg.ToolCalls[i] = db.ToolCall{ + ToolName: call.ToolName, + Category: call.Category, + ToolUseID: call.ToolUseID, + InputJSON: call.InputJSON, + FilePath: call.FilePath, + SkillName: call.SkillName, + ResultContentLength: call.ResultContentLength, + ResultContent: call.ResultContent, + SubagentSessionID: call.SubagentSessionID, + } + if len(call.ResultEvents) > 0 { + msg.ToolCalls[i].ResultEvents = make([]db.ToolResultEvent, len(call.ResultEvents)) + for j, ev := range call.ResultEvents { + msg.ToolCalls[i].ResultEvents[j] = db.ToolResultEvent{ + ToolUseID: ev.ToolUseID, + AgentID: ev.AgentID, + SubagentSessionID: ev.SubagentSessionID, + Source: ev.Source, + Status: ev.Status, + Content: ev.Content, + ContentLength: ev.ContentLength, + Timestamp: ev.Timestamp, + EventIndex: ev.EventIndex, + } + } + } + } + } + return msg +} + +func canonicalJSON(v any) ([]byte, error) { + var buf bytes.Buffer + if err := writeCanonicalJSON(&buf, reflect.ValueOf(v)); err != nil { + return nil, fmt.Errorf("encoding canonical artifact JSON: %w", err) + } + buf.WriteByte('\n') + return buf.Bytes(), nil +} + +func writeCanonicalJSON(buf *bytes.Buffer, v reflect.Value) error { + if !v.IsValid() { + buf.WriteString("null") + return nil + } + if v.Kind() == reflect.Interface { + if v.IsNil() { + buf.WriteString("null") + return nil + } + return writeCanonicalJSON(buf, v.Elem()) + } + if v.Kind() == reflect.Pointer { + if v.IsNil() { + buf.WriteString("null") + return nil + } + return writeCanonicalJSON(buf, v.Elem()) + } + if v.Type() == reflect.TypeFor[json.RawMessage]() { + raw := v.Interface().(json.RawMessage) + if len(raw) == 0 { + buf.WriteString("null") + return nil + } + dec := json.NewDecoder(bytes.NewReader(raw)) + dec.UseNumber() + var decoded any + if err := dec.Decode(&decoded); err != nil { + return err + } + return writeCanonicalJSON(buf, reflect.ValueOf(decoded)) + } + if v.Type() == reflect.TypeFor[json.Number]() { + buf.WriteString(v.Interface().(json.Number).String()) + return nil + } + switch v.Kind() { + case reflect.Bool: + buf.WriteString(strconv.FormatBool(v.Bool())) + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + buf.WriteString(strconv.FormatInt(v.Int(), 10)) + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr: + buf.WriteString(strconv.FormatUint(v.Uint(), 10)) + case reflect.Float32, reflect.Float64: + data, err := json.Marshal(v.Interface()) + if err != nil { + return err + } + buf.Write(data) + case reflect.String: + data, err := json.Marshal(v.String()) + if err != nil { + return err + } + buf.Write(data) + case reflect.Slice, reflect.Array: + buf.WriteByte('[') + for i := 0; i < v.Len(); i++ { + if i > 0 { + buf.WriteByte(',') + } + if err := writeCanonicalJSON(buf, v.Index(i)); err != nil { + return err + } + } + buf.WriteByte(']') + case reflect.Map: + return writeCanonicalMap(buf, v) + case reflect.Struct: + return writeCanonicalStruct(buf, v) + default: + return fmt.Errorf("unsupported canonical JSON kind %s", v.Kind()) + } + return nil +} + +func writeCanonicalMap(buf *bytes.Buffer, v reflect.Value) error { + if v.IsNil() { + buf.WriteString("null") + return nil + } + if v.Type().Key().Kind() != reflect.String { + return fmt.Errorf("unsupported canonical map key type %s", v.Type().Key()) + } + keys := make([]string, 0, v.Len()) + for _, key := range v.MapKeys() { + keys = append(keys, key.String()) + } + sort.Strings(keys) + buf.WriteByte('{') + for i, key := range keys { + if i > 0 { + buf.WriteByte(',') + } + keyData, err := json.Marshal(key) + if err != nil { + return err + } + buf.Write(keyData) + buf.WriteByte(':') + if err := writeCanonicalJSON(buf, v.MapIndex(reflect.ValueOf(key))); err != nil { + return err + } + } + buf.WriteByte('}') + return nil +} + +type canonicalField struct { + name string + value reflect.Value +} + +func writeCanonicalStruct(buf *bytes.Buffer, v reflect.Value) error { + fields := make([]canonicalField, 0, v.NumField()) + t := v.Type() + for i := 0; i < v.NumField(); i++ { + field := t.Field(i) + if field.PkgPath != "" { + continue + } + name, omitEmpty, skip := jsonField(field) + if skip { + continue + } + value := v.Field(i) + if omitEmpty && isCanonicalEmpty(value) { + continue + } + fields = append(fields, canonicalField{name: name, value: value}) + } + sort.Slice(fields, func(i, j int) bool { + return fields[i].name < fields[j].name + }) + + buf.WriteByte('{') + for i, field := range fields { + if i > 0 { + buf.WriteByte(',') + } + name, err := json.Marshal(field.name) + if err != nil { + return err + } + buf.Write(name) + buf.WriteByte(':') + if err := writeCanonicalJSON(buf, field.value); err != nil { + return err + } + } + buf.WriteByte('}') + return nil +} + +func jsonField(field reflect.StructField) (name string, omitEmpty bool, skip bool) { + name = field.Name + tag := field.Tag.Get("json") + if tag == "-" { + return "", false, true + } + if tag == "" { + return name, false, false + } + parts := strings.Split(tag, ",") + if parts[0] != "" { + name = parts[0] + } + for _, opt := range parts[1:] { + if opt == "omitempty" { + omitEmpty = true + } + } + return name, omitEmpty, false +} + +func isCanonicalEmpty(v reflect.Value) bool { + if !v.IsValid() { + return true + } + switch v.Kind() { + case reflect.Array: + return v.Len() == 0 + case reflect.Map, reflect.Slice, reflect.String: + return v.Len() == 0 + case reflect.Bool: + return !v.Bool() + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + return v.Int() == 0 + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr: + return v.Uint() == 0 + case reflect.Float32, reflect.Float64: + return v.Float() == 0 + case reflect.Interface, reflect.Pointer: + return v.IsNil() + } + return false +} + +func hashHex(data []byte) string { + sum := sha256.Sum256(data) + return hex.EncodeToString(sum[:]) +} + +// reconcileArtifactConflict handles a same-name, different-content pair. For a +// recognized artifact whose source validates and destination does not, the +// destination is repaired in place; a corrupt source is skipped instead of +// mirrored. Everything else keeps the write-once conflict error. +func isTempArtifactEntry(name string) bool { + return strings.HasPrefix(name, tempFilePrefix) +} + +// IsFolderTarget reports whether target is a local filesystem target rather +// than a future HTTP or object-store target. +func IsFolderTarget(target string) bool { + if target == "" || strings.Contains(target, "://") { + return false + } + if isWindowsDrivePath(target) { + return true + } + _, _, err := net.SplitHostPort(target) + return err != nil +} + +func isWindowsDrivePath(target string) bool { + if len(target) < 3 || target[1] != ':' { + return false + } + c := target[0] + if (c < 'A' || c > 'Z') && (c < 'a' || c > 'z') { + return false + } + return target[2] == '\\' || target[2] == '/' +} diff --git a/internal/artifact/sync_test.go b/internal/artifact/sync_test.go new file mode 100644 index 000000000..418c952f1 --- /dev/null +++ b/internal/artifact/sync_test.go @@ -0,0 +1,3918 @@ +package artifact + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/docbank" +) + +func TestEnsureOriginPersists(t *testing.T) { + database := testDB(t) + + first, err := EnsureOrigin(database) + require.NoError(t, err) + require.NotEmpty(t, first) + require.NotEqual(t, "local", first) + + second, err := EnsureOrigin(database) + require.NoError(t, err) + assert.Equal(t, first, second) +} + +func TestCheckpointFloorBootstrapsFromLiveAndQuarantinedNodes(t *testing.T) { + _, store := newTestDocbankStore(t, docbank.Config{}) + database := testDB(t) + origin := contractOrigin + for sequence := 1; sequence <= checkpointFloorPageSize+2; sequence++ { + body := fmt.Appendf(nil, `{"v":1,"origin":%q,"sequence":%d,"sessions":{}}`, origin, sequence) + createCheckpointBody(t, store, sequence, body) + } + quarantinedName := fmt.Sprintf("cp-%010d.json", checkpointFloorPageSize+2) + require.NoError(t, store.Quarantine(t.Context(), + requireContractRef(t, origin, KindCheckpoints, quarantinedName), + "test quarantine")) + + sequence, err := reserveCheckpointSequenceFromStore( + t.Context(), database, store, origin, + ) + require.NoError(t, err) + assert.Equal(t, 131, sequence) + + // A fresh vault may report no sequence after reset or quarantine expiry, but + // the SQLite floor remains authoritative and may never be lowered. + _, emptyStore := newTestDocbankStore(t, docbank.Config{}) + sequence, err = reserveCheckpointSequenceFromStore( + t.Context(), database, emptyStore, origin, + ) + require.NoError(t, err) + assert.Equal(t, 132, sequence) + + // If both SQLite and the vault are lost simultaneously, local prevention is + // impossible; peer common-checkpoint conflict handling is the final backstop. +} + +func TestCheckpointFloorTraversesStoreOnlyBeforeBootstrap(t *testing.T) { + database := testDB(t) + store := &countingCheckpointFloorStore{floor: 40} + + sequence, err := reserveCheckpointSequenceFromStore( + t.Context(), database, store, contractOrigin, + ) + require.NoError(t, err) + assert.Equal(t, 41, sequence) + sequence, err = reserveCheckpointSequenceFromStore( + t.Context(), database, store, contractOrigin, + ) + require.NoError(t, err) + assert.Equal(t, 42, sequence) + assert.Equal(t, 1, store.calls, "durable floor avoids repeated vault traversal") +} + +type countingCheckpointFloorStore struct { + ArtifactStore + floor int + calls int +} + +func (s *countingCheckpointFloorStore) checkpointFloor(context.Context, string) (int, error) { + s.calls++ + return s.floor, nil +} + +type recordingSyncStateValueReader struct { + states map[string]string + keys []string + calls int +} + +type countingQueuedExportStore struct { + database *db.DB + queueQueries int + sessionLoads int + messageLoads int + usageLoads int +} + +type countingCanonicalExportDB struct { + *db.DB + sessionLoads int + messageLoads int + usageLoads int +} + +func (s *countingCanonicalExportDB) GetSessionFull( + ctx context.Context, id string, +) (*db.Session, error) { + s.sessionLoads++ + return s.DB.GetSessionFull(ctx, id) +} + +func (s *countingCanonicalExportDB) GetAllMessages( + ctx context.Context, id string, +) ([]db.Message, error) { + s.messageLoads++ + return s.DB.GetAllMessages(ctx, id) +} + +func (s *countingCanonicalExportDB) GetUsageEvents( + ctx context.Context, id string, +) ([]db.UsageEvent, error) { + s.usageLoads++ + return s.DB.GetUsageEvents(ctx, id) +} + +type recordedArtifactCreate struct { + Ref Ref + Identity Identity + Created bool +} + +type recordingArtifactStore struct { + ArtifactStore + creates []recordedArtifactCreate +} + +type checkpointReadCountingStore struct { + ArtifactStore + lists int + opens int +} + +type dependencyOpenCountingStore struct { + ArtifactStore + opens map[Kind]int +} + +func (s *dependencyOpenCountingStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + s.opens[ref.Kind]++ + return s.ArtifactStore.Open(ctx, ref) +} + +type catalogCheckpointStore struct { + ArtifactStore + entry Entry + statCalls int + openCalls int +} + +type checkpointStatOverrideStore struct { + ArtifactStore + identity *Identity + err error +} + +type originTraversalGateStore struct { + ArtifactStore + processedFirstPage bool +} + +func (s *originTraversalGateStore) Origins(context.Context) (OriginIterator, error) { + page := 0 + return &testOriginIterator{next: func(context.Context, int) ([]string, error) { + page++ + if page == 1 { + origins := make([]string, artifactImportPageSize) + for i := range origins { + origins[i] = fmt.Sprintf("peer-%06x", i+1) + } + return origins, nil + } + if !s.processedFirstPage { + return nil, errors.New("second origin page requested before first page was processed") + } + return []string{"peer-ffffff"}, io.EOF + }}, nil +} + +func (s *originTraversalGateStore) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + s.processedFirstPage = true + return s.ArtifactStore.Entries(ctx, origin, kind) +} + +type checkpointHeaderFirstStore struct { + ArtifactStore + listedAll bool +} + +func (s *checkpointHeaderFirstStore) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + iterator, err := s.ArtifactStore.Entries(ctx, origin, kind) + if err != nil || kind != KindCheckpoints { + return iterator, err + } + return &testEntryIterator{ + next: func(ctx context.Context, limit int) ([]Entry, error) { + entries, err := iterator.Next(ctx, limit) + if errors.Is(err, io.EOF) { + s.listedAll = true + } + return entries, err + }, + close: iterator.Close, + }, nil +} + +func (s *checkpointHeaderFirstStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + if ref.Kind == KindCheckpoints && !s.listedAll { + return Entry{}, nil, errors.New("checkpoint body opened before all headers were enumerated") + } + return s.ArtifactStore.Open(ctx, ref) +} + +type cancelAfterMetadataOpenStore struct { + ArtifactStore + cancel context.CancelFunc + opens int +} + +type exactImportCountingStore struct { + ArtifactStore + origins atomic.Int32 + entries atomic.Int32 + opens atomic.Int32 +} + +type exactImportGateStore struct { + ArtifactStore + entered chan struct{} + release chan struct{} + once sync.Once + active atomic.Int32 + maximum atomic.Int32 + opens atomic.Int32 +} + +func (s *exactImportGateStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + active := s.active.Add(1) + defer s.active.Add(-1) + for { + maximum := s.maximum.Load() + if active <= maximum || s.maximum.CompareAndSwap(maximum, active) { + break + } + } + s.opens.Add(1) + blocked := false + s.once.Do(func() { + blocked = true + close(s.entered) + }) + if blocked { + select { + case <-ctx.Done(): + return Entry{}, nil, ctx.Err() + case <-s.release: + } + } + return s.ArtifactStore.Open(ctx, ref) +} + +func (s *exactImportCountingStore) Origins(ctx context.Context) (OriginIterator, error) { + s.origins.Add(1) + return s.ArtifactStore.Origins(ctx) +} + +func (s *exactImportCountingStore) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + s.entries.Add(1) + return s.ArtifactStore.Entries(ctx, origin, kind) +} + +func (s *exactImportCountingStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + s.opens.Add(1) + return s.ArtifactStore.Open(ctx, ref) +} + +type trackingRepairStore struct { + ArtifactStore + repairs int +} + +type failingImportScheduler struct { + err error + calls int +} + +func (s *failingImportScheduler) RecordChanged(context.Context, Entry) error { + s.calls++ + return s.err +} + +type transientImportStore struct { + ArtifactStore + calls atomic.Int32 +} + +func (s *transientImportStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + if s.calls.Add(1) == 1 { + return Entry{}, nil, syscall.EAGAIN + } + return s.ArtifactStore.Open(ctx, ref) +} + +func (s *trackingRepairStore) RepairContent( + context.Context, Identity, io.Reader, +) error { + s.repairs++ + return nil +} + +func (s *cancelAfterMetadataOpenStore) Open( + _ context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + entry, reader, err := s.ArtifactStore.Open(context.Background(), ref) + if err == nil && ref.Kind == KindMeta { + s.opens++ + if s.opens == 1 { + s.cancel() + } + } + return entry, reader, err +} + +func (s *checkpointStatOverrideStore) Stat(ctx context.Context, ref Ref) (Entry, error) { + if ref.Kind == Kind(KindCheckpoints) { + if s.err != nil { + return Entry{}, s.err + } + if s.identity != nil { + return Entry{Ref: ref, Identity: *s.identity}, nil + } + } + return s.ArtifactStore.Stat(ctx, ref) +} + +func (s *catalogCheckpointStore) Stat(context.Context, Ref) (Entry, error) { + s.statCalls++ + return s.entry, nil +} + +func (s *catalogCheckpointStore) Open( + context.Context, Ref, +) (Entry, VerifiedReader, error) { + s.openCalls++ + return Entry{}, nil, errors.New("checkpoint body must not be opened") +} + +func (s *checkpointReadCountingStore) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + if kind == Kind(KindCheckpoints) { + s.lists++ + } + return s.ArtifactStore.Entries(ctx, origin, kind) +} + +func (s *checkpointReadCountingStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + if ref.Kind == Kind(KindCheckpoints) { + s.opens++ + } + return s.ArtifactStore.Open(ctx, ref) +} + +func (s *recordingArtifactStore) Create( + ctx context.Context, + ref Ref, + identity Identity, + mediaType string, + body io.Reader, +) (CreateResult, error) { + result, err := s.ArtifactStore.Create(ctx, ref, identity, mediaType, body) + if err == nil { + s.creates = append(s.creates, recordedArtifactCreate{ + Ref: ref, Identity: identity, Created: result.Created, + }) + } + return result, err +} + +func TestExportToStorePublishesDependenciesBeforeCheckpointAndSkipsUnchanged(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + store := &recordingArtifactStore{ArtifactStore: filesystem} + + result, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.Equal(t, 1, result.ExportedSessions) + assert.True(t, result.CheckpointCreated) + require.Len(t, store.creates, 3) + assert.Equal(t, Kind(KindSegments), store.creates[0].Ref.Kind) + assert.Equal(t, Kind(KindManifests), store.creates[1].Ref.Kind) + assert.Equal(t, Kind(KindCheckpoints), store.creates[2].Ref.Kind) + assert.True(t, store.creates[0].Created) + assert.True(t, store.creates[1].Created) + assert.True(t, store.creates[2].Created) + checkpointBytes := readContractArtifact(t, store, store.creates[2].Ref) + var published checkpoint + require.NoError(t, json.Unmarshal(checkpointBytes, &published)) + canonicalCheckpoint, err := canonicalJSON(published) + require.NoError(t, err) + assert.Equal(t, canonicalCheckpoint, checkpointBytes, + "streaming checkpoint encoding must remain byte-compatible") + + store.creates = nil + result, err = ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + Full: true, + }) + require.NoError(t, err) + assert.Zero(t, result.ExportedSessions) + assert.False(t, result.CheckpointCreated) + require.Len(t, store.creates, 2, + "full export verifies immutable dependencies without minting a version") + assert.Equal(t, Kind(KindSegments), store.creates[0].Ref.Kind) + assert.Equal(t, Kind(KindManifests), store.creates[1].Ref.Kind) + assert.False(t, store.creates[0].Created) + assert.False(t, store.creates[1].Created) +} + +func TestExportToStoreFullRepairsMissingDependencyWithoutNewCheckpoint(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + first, err := ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + require.True(t, first.CheckpointCreated) + segments, err := firstStoreEntryPage(t.Context(), filesystem, contractOrigin, KindSegments, 10) + require.NoError(t, err) + require.Len(t, segments.Items, 1) + require.NoError(t, filesystem.Trash(t.Context(), segments.Items[0].Ref)) + _, err = filesystem.Stat(t.Context(), segments.Items[0].Ref) + require.ErrorIs(t, err, ErrArtifactNotFound) + headBefore, ok, err := database.GetArtifactCheckpointHead(t.Context(), contractOrigin) + require.NoError(t, err) + require.True(t, ok) + + result, err := ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, Full: true, + }) + require.NoError(t, err) + assert.False(t, result.CheckpointCreated) + _, err = filesystem.Stat(t.Context(), segments.Items[0].Ref) + require.NoError(t, err, "full export recreates a missing dependency even when its manifest survives") + headAfter, ok, err := database.GetArtifactCheckpointHead(t.Context(), contractOrigin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, headBefore, headAfter) +} + +func TestExportToStoreRecreatesMissingRecordedCheckpointAfterVaultReset(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + first, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + _, err = ExportToStore(t.Context(), database, first, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + require.NoError(t, first.Close()) + + replacement, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, replacement.Close()) }) + result, err := ExportToStore(t.Context(), database, replacement, ExportOptions{ + Origin: contractOrigin, + Full: true, + }) + require.NoError(t, err) + assert.True(t, result.CheckpointCreated) + assert.Equal(t, 1, result.CheckpointSequence, + "vault reset recreates the recorded immutable checkpoint without consuming a new sequence") + checkpointRef := requireContractRef(t, contractOrigin, KindCheckpoints, + "cp-0000000001.json") + checkpointBytes := readContractArtifact(t, replacement, checkpointRef) + var published checkpoint + require.NoError(t, json.Unmarshal(checkpointBytes, &published)) + assert.Equal(t, map[string]string{ + contractOrigin + "~sess-1": published.Sessions[contractOrigin+"~sess-1"], + }, published.Sessions) +} + +func TestExportToStoreChangedBatchDoesNotScanCheckpointHistory(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + store := &checkpointReadCountingStore{ArtifactStore: filesystem} + _, err = ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + store.lists = 0 + store.opens = 0 + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{{ + SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "changed", + }})) + + result, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.True(t, result.CheckpointCreated) + assert.Equal(t, 2, result.CheckpointSequence) + assert.Zero(t, store.lists, + "a recorded pre-apply head avoids checkpoint-history traversal") + assert.Zero(t, store.opens, + "normal changed export does not read an old checkpoint body") +} + +func TestExportToStoreUnchangedCheckpointUsesCatalogIdentityOnly(t *testing.T) { + for _, size := range []int64{128, 64 << 20} { + t.Run(fmt.Sprintf("size-%d", size), func(t *testing.T) { + database := testDB(t) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json") + identity := Identity{SHA256: strings64("a"), Size: size} + require.NoError(t, database.RecordArtifactCheckpointHead(t.Context(), db.ArtifactCheckpointHead{ + Origin: contractOrigin, Sequence: 1, + SessionMapSHA256: strings64("b"), CheckpointSHA256: identity.SHA256, + CheckpointSize: identity.Size, + }, nil)) + store := &catalogCheckpointStore{entry: Entry{Ref: ref, Identity: identity}} + + result, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.False(t, result.CheckpointCreated) + assert.Equal(t, 1, store.statCalls) + assert.Zero(t, store.openCalls, + "unchanged periodic export cannot drain the checkpoint body") + }) + } +} + +func TestExportToStoreRecordedCheckpointStatRecovery(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + _, err = ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + + operationalFailure := errors.New("injected checkpoint stat failure") + _, err = ExportToStore(t.Context(), database, &checkpointStatOverrideStore{ + ArtifactStore: filesystem, err: operationalFailure, + }, ExportOptions{Origin: contractOrigin}) + require.ErrorIs(t, err, operationalFailure) + + mismatch := Identity{SHA256: strings64("c"), Size: 1} + result, err := ExportToStore(t.Context(), database, &checkpointStatOverrideStore{ + ArtifactStore: filesystem, identity: &mismatch, + }, ExportOptions{Origin: contractOrigin}) + require.NoError(t, err) + assert.True(t, result.CheckpointCreated) + assert.Equal(t, 1, result.CheckpointSequence, + "catalog identity mismatch quarantines and reconstructs the recorded checkpoint") +} + +type maxReadReader struct { + reader io.Reader + max int +} + +func (r *maxReadReader) Read(p []byte) (int, error) { + if len(p) > r.max { + r.max = len(p) + } + return r.reader.Read(p) +} + +func TestExportCheckpointBootstrapStreamsLargeSessionMap(t *testing.T) { + sessions := make(map[string]string, 2000) + for i := range 2000 { + sessions[fmt.Sprintf("%s~session-%04d", contractOrigin, i)] = strings64("a") + } + body, err := canonicalJSON(checkpoint{ + Version: formatVersion, Origin: contractOrigin, Sequence: 42, Sessions: sessions, + }) + require.NoError(t, err) + reader := &maxReadReader{reader: strings.NewReader(string(body))} + head, err := decodeCanonicalCheckpointHead(reader, contractOrigin, + "cp-0000000042.json", identityForBytes(t, body)) + require.NoError(t, err) + mapBytes, err := canonicalJSON(sessions) + require.NoError(t, err) + assert.Equal(t, hashHex(mapBytes), head.SessionMapSHA256) + assert.Less(t, reader.max, len(body)/4, + "bootstrap must tokenize the checkpoint instead of reading its full body") +} + +type hiddenCheckpointHeadDB struct { + *db.DB +} + +func (s *hiddenCheckpointHeadDB) GetArtifactCheckpointHead( + context.Context, string, +) (db.ArtifactCheckpointHead, bool, error) { + return db.ArtifactCheckpointHead{}, false, nil +} + +func TestExportToStoreBootstrapsMissingDatabaseHeadFromLatestCheckpoint(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + _, err = ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + store := &checkpointReadCountingStore{ArtifactStore: filesystem} + + result, err := ExportToStore(t.Context(), &hiddenCheckpointHeadDB{DB: database}, + store, ExportOptions{Origin: contractOrigin}) + require.NoError(t, err) + assert.False(t, result.CheckpointCreated) + assert.Equal(t, 1, result.CheckpointSequence) + assert.Positive(t, store.lists) + assert.Equal(t, 1, store.opens) + page, err := firstStoreEntryPage(t.Context(), filesystem, contractOrigin, KindCheckpoints, 10) + require.NoError(t, err) + assert.Len(t, page.Items, 1) +} + +type closeErrorVerifiedReader struct { + VerifiedReader + err error +} + +func (r *closeErrorVerifiedReader) Close() error { + return errors.Join(r.VerifiedReader.Close(), r.err) +} + +type checkpointCloseErrorStore struct { + ArtifactStore + err error +} + +type checkpointOpenFailureStore struct { + ArtifactStore + err error +} + +func (s *checkpointOpenFailureStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + if ref.Kind == Kind(KindCheckpoints) { + return Entry{}, nil, s.err + } + return s.ArtifactStore.Open(ctx, ref) +} + +func TestExportCheckpointBootstrapPropagatesOperationalOpenErrors(t *testing.T) { + for _, test := range []struct { + name string + err error + }{ + {name: "canceled", err: context.Canceled}, + {name: "unavailable", err: errors.New("checkpoint store unavailable")}, + } { + t.Run(test.name, func(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + _, err = ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + + _, err = ExportToStore(t.Context(), &hiddenCheckpointHeadDB{DB: database}, + &checkpointOpenFailureStore{ArtifactStore: filesystem, err: test.err}, + ExportOptions{Origin: contractOrigin}) + require.ErrorIs(t, err, test.err) + }) + } +} + +func TestLatestValidCheckpointFallsBackPastSemanticCandidateAndDefersFuture(t *testing.T) { + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + for _, candidate := range []struct { + name string + body string + }{ + {name: "cp-0000000001.json", body: `{"origin":"contract-a1b2c3","seq":1,"sessions":{},"v":1}` + "\n"}, + {name: "cp-0000000002.json", body: `{"origin":"contract-a1b2c3","seq":99,"sessions":{},"v":1}` + "\n"}, + } { + body := []byte(candidate.body) + ref := requireContractRef(t, contractOrigin, KindCheckpoints, candidate.name) + _, err := filesystem.Create(t.Context(), ref, identityForBytes(t, body), + canonicalArtifactMediaType(KindCheckpoints), strings.NewReader(candidate.body)) + require.NoError(t, err) + } + + head, ok, err := latestValidCheckpointHead(t.Context(), filesystem, contractOrigin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 1, head.Sequence) + + future := []byte(`{"origin":"contract-a1b2c3","seq":3,"sessions":{},"v":2}` + "\n") + ref := requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000003.json") + _, err = filesystem.Create(t.Context(), ref, identityForBytes(t, future), + canonicalArtifactMediaType(KindCheckpoints), strings.NewReader(string(future))) + require.NoError(t, err) + _, _, err = latestValidCheckpointHead(t.Context(), filesystem, contractOrigin) + require.ErrorIs(t, err, errFutureArtifactVersion) +} + +func TestExportSpoolConstructionJoinsCleanupFailures(t *testing.T) { + chmodFailure := errors.New("injected spool setup failure") + cleanupFailure := errors.New("injected spool cleanup failure") + previousChmod := exportSpoolChmod + previousCleanup := exportSpoolCleanup + t.Cleanup(func() { + exportSpoolChmod = previousChmod + exportSpoolCleanup = previousCleanup + }) + exportSpoolChmod = func(*os.File) error { return chmodFailure } + exportSpoolCleanup = func(file *os.File) error { + return errors.Join(closeAndRemoveExportSpool(file), cleanupFailure) + } + + _, _, _, err := spoolArtifactPublicationMap( + t.Context(), testDB(t), contractOrigin, + ) + require.ErrorIs(t, err, chmodFailure) + require.ErrorIs(t, err, cleanupFailure) + + mapSpool, err := os.CreateTemp("", "agentsview-artifact-test-map-*") + require.NoError(t, err) + t.Cleanup(func() { _ = closeAndRemoveExportSpool(mapSpool) }) + _, err = io.WriteString(mapSpool, "{}\n") + require.NoError(t, err) + _, _, err = spoolArtifactCheckpoint(t.Context(), mapSpool, contractOrigin, 1) + require.ErrorIs(t, err, chmodFailure) + require.ErrorIs(t, err, cleanupFailure) +} + +func (s *checkpointCloseErrorStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + entry, reader, err := s.ArtifactStore.Open(ctx, ref) + if err != nil || ref.Kind != Kind(KindCheckpoints) { + return entry, reader, err + } + return entry, &closeErrorVerifiedReader{VerifiedReader: reader, err: s.err}, nil +} + +func TestExportToStorePropagatesCheckpointCloseErrorWithoutAcknowledging(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + _, err = ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + require.NoError(t, database.ReplaceSessionMessages("sess-1", []db.Message{{ + SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "changed", + }})) + closeFailure := errors.New("injected checkpoint close failure") + + _, err = ExportToStore(t.Context(), &hiddenCheckpointHeadDB{DB: database}, &checkpointCloseErrorStore{ + ArtifactStore: filesystem, err: closeFailure, + }, ExportOptions{Origin: contractOrigin}) + require.ErrorIs(t, err, closeFailure) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + assert.Len(t, pending, 1) +} + +func TestExportToStoreCancellationLeavesQueuePending(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err = ExportToStore(ctx, database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.ErrorIs(t, err, context.Canceled) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + assert.Len(t, pending, 1) +} + +type failingArtifactStore struct { + ArtifactStore + failKind Kind + failErr error + failed bool + calls []Kind +} + +func (s *failingArtifactStore) Create( + ctx context.Context, + ref Ref, + identity Identity, + mediaType string, + body io.Reader, +) (CreateResult, error) { + s.calls = append(s.calls, ref.Kind) + if !s.failed && ref.Kind == s.failKind { + s.failed = true + return CreateResult{}, s.failErr + } + return s.ArtifactStore.Create(ctx, ref, identity, mediaType, body) +} + +func TestExportToStoreFailureKeepsClaimAndCheckpointLast(t *testing.T) { + failure := errors.New("injected artifact create failure") + tests := []struct { + name string + failKind Kind + wantCalls []Kind + wantRetrySeq int + }{ + {name: "dependency", failKind: Kind(KindSegments), + wantCalls: []Kind{Kind(KindSegments)}, wantRetrySeq: 1}, + {name: "manifest", failKind: Kind(KindManifests), + wantCalls: []Kind{Kind(KindSegments), Kind(KindManifests)}, wantRetrySeq: 1}, + {name: "checkpoint", failKind: Kind(KindCheckpoints), + wantCalls: []Kind{Kind(KindSegments), Kind(KindManifests), Kind(KindCheckpoints)}, wantRetrySeq: 2}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + store := &failingArtifactStore{ + ArtifactStore: filesystem, failKind: tt.failKind, failErr: failure, + } + + _, err = ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.ErrorIs(t, err, failure) + assert.Equal(t, tt.wantCalls, store.calls) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1, "failed export must retain its exact claim") + page, err := firstStoreEntryPage(t.Context(), store, contractOrigin, KindCheckpoints, 10) + require.NoError(t, err) + assert.Empty(t, page.Items, "checkpoint cannot precede a failed dependency") + + store.calls = nil + result, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.True(t, result.CheckpointCreated) + assert.Equal(t, tt.wantRetrySeq, result.CheckpointSequence) + pending, err = database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + assert.Empty(t, pending) + }) + } +} + +type staleClaimExportDB struct { + *db.DB + once sync.Once +} + +func (s *staleClaimExportDB) ApplyArtifactPublicationChanges( + ctx context.Context, + origin string, + changes []db.ArtifactPublicationChange, +) (int64, bool, error) { + s.once.Do(func() { + _ = s.ReplaceSessionMessages("sess-1", []db.Message{{ + SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "newer", + }}) + }) + return s.DB.ApplyArtifactPublicationChanges(ctx, origin, changes) +} + +func TestExportToStoreStaleGenerationDoesNotPublishOrAcknowledge(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + stale := &staleClaimExportDB{DB: database} + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + + _, err = ExportToStore(t.Context(), stale, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.ErrorIs(t, err, db.ErrArtifactExportClaimStale) + _, ok, err := database.GetArtifactCheckpointHead(t.Context(), contractOrigin) + require.NoError(t, err) + assert.False(t, ok) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Greater(t, pending[0].Generation, int64(1)) + page, err := firstStoreEntryPage(t.Context(), filesystem, contractOrigin, KindCheckpoints, 10) + require.NoError(t, err) + assert.Empty(t, page.Items) +} + +type mutateAfterCheckpointStore struct { + ArtifactStore + database *db.DB + once sync.Once +} + +func (s *mutateAfterCheckpointStore) Create( + ctx context.Context, + ref Ref, + identity Identity, + mediaType string, + body io.Reader, +) (CreateResult, error) { + result, err := s.ArtifactStore.Create(ctx, ref, identity, mediaType, body) + if err == nil && ref.Kind == Kind(KindCheckpoints) { + s.once.Do(func() { + _ = s.database.ReplaceSessionMessages("sess-1", []db.Message{{ + SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "newest", + }}) + }) + } + return result, err +} + +func TestExportToStoreStaleGenerationAfterCheckpointDoesNotAdvanceHead(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + store := &mutateAfterCheckpointStore{ArtifactStore: filesystem, database: database} + + _, err = ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.ErrorIs(t, err, db.ErrArtifactExportClaimStale) + _, ok, err := database.GetArtifactCheckpointHead(t.Context(), contractOrigin) + require.NoError(t, err) + assert.False(t, ok) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + firstPage, err := firstStoreEntryPage(t.Context(), filesystem, contractOrigin, KindCheckpoints, 10) + require.NoError(t, err) + require.Len(t, firstPage.Items, 1, + "the immutable stale checkpoint is harmless while no head references it") + + result, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.Equal(t, 2, result.CheckpointSequence) + head, ok, err := database.GetArtifactCheckpointHead(t.Context(), contractOrigin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 2, head.Sequence) + pending, err = database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + assert.Empty(t, pending) +} + +func TestExportToStorePublicationRevisionRejectsPhysicallyCreatedStaleCheckpoint(t *testing.T) { + path := filepath.Join(t.TempDir(), "archive.db") + first, err := db.Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, first.Close()) }) + second, err := db.Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, second.Close()) }) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + ctx := t.Context() + seedSession(t, first, "session-a", "project") + claimA, err := first.ArtifactExportClaims(ctx, []string{"session-a"}) + require.NoError(t, err) + require.Len(t, claimA, 1) + sessionA, err := first.GetSessionFull(ctx, "session-a") + require.NoError(t, err) + messagesA, err := first.GetAllMessages(ctx, "session-a") + require.NoError(t, err) + usageA, err := first.GetUsageEvents(ctx, "session-a") + require.NoError(t, err) + manifestA, _, err := exportLoadedSessionToStore( + ctx, filesystem, contractOrigin, sessionA, messagesA, usageA, productionArtifactLimits(), + ) + require.NoError(t, err) + revisionA, changed, err := first.ApplyArtifactPublicationChanges(ctx, contractOrigin, + []db.ArtifactPublicationChange{{ + SessionID: "session-a", Generation: claimA[0].Generation, + ManifestHash: manifestA, SourceFingerprint: manifestA, + }}) + require.NoError(t, err) + require.True(t, changed) + mapA, digestA, snapshotA, err := spoolArtifactPublicationMap(ctx, first, contractOrigin) + require.NoError(t, err) + require.Equal(t, revisionA, snapshotA) + t.Cleanup(func() { _ = closeAndRemoveExportSpool(mapA) }) + + seedSession(t, second, "session-b", "project") + claimB, err := second.ArtifactExportClaims(ctx, []string{"session-b"}) + require.NoError(t, err) + require.Len(t, claimB, 1) + sessionB, err := second.GetSessionFull(ctx, "session-b") + require.NoError(t, err) + messagesB, err := second.GetAllMessages(ctx, "session-b") + require.NoError(t, err) + usageB, err := second.GetUsageEvents(ctx, "session-b") + require.NoError(t, err) + manifestB, _, err := exportLoadedSessionToStore( + ctx, filesystem, contractOrigin, sessionB, messagesB, usageB, productionArtifactLimits(), + ) + require.NoError(t, err) + revisionB, changed, err := second.ApplyArtifactPublicationChanges(ctx, contractOrigin, + []db.ArtifactPublicationChange{{ + SessionID: "session-b", Generation: claimB[0].Generation, + ManifestHash: manifestB, SourceFingerprint: manifestB, + }}) + require.NoError(t, err) + require.True(t, changed) + mapB, digestB, snapshotB, err := spoolArtifactPublicationMap(ctx, second, contractOrigin) + require.NoError(t, err) + require.Equal(t, revisionB, snapshotB) + sequenceB, err := reserveCheckpointSequenceFromStore(ctx, second, filesystem, contractOrigin) + require.NoError(t, err) + checkpointB, identityB, err := spoolArtifactCheckpoint(ctx, mapB, contractOrigin, sequenceB) + require.NoError(t, err) + refB := requireContractRef(t, contractOrigin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequenceB)) + _, err = filesystem.Create(ctx, refB, identityB, + canonicalArtifactMediaType(KindCheckpoints), checkpointB) + require.NoError(t, err) + require.NoError(t, closeAndRemoveExportSpool(mapB)) + require.NoError(t, closeAndRemoveExportSpool(checkpointB)) + require.NoError(t, second.RecordArtifactCheckpointHead(ctx, db.ArtifactCheckpointHead{ + Origin: contractOrigin, Sequence: sequenceB, PublicationRevision: snapshotB, + SessionMapSHA256: digestB, CheckpointSHA256: identityB.SHA256, + CheckpointSize: identityB.Size, + }, claimB)) + + sequenceA, err := reserveCheckpointSequenceFromStore(ctx, first, filesystem, contractOrigin) + require.NoError(t, err) + require.Greater(t, sequenceA, sequenceB) + checkpointA, identityA, err := spoolArtifactCheckpoint(ctx, mapA, contractOrigin, sequenceA) + require.NoError(t, err) + refA := requireContractRef(t, contractOrigin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequenceA)) + _, err = filesystem.Create(ctx, refA, identityA, + canonicalArtifactMediaType(KindCheckpoints), checkpointA) + require.NoError(t, err) + require.NoError(t, closeAndRemoveExportSpool(checkpointA)) + err = first.RecordArtifactCheckpointHead(ctx, db.ArtifactCheckpointHead{ + Origin: contractOrigin, Sequence: sequenceA, PublicationRevision: snapshotA, + SessionMapSHA256: digestA, CheckpointSHA256: identityA.SHA256, + CheckpointSize: identityA.Size, + }, claimA) + require.ErrorIs(t, err, db.ErrArtifactExportClaimStale) + _, err = filesystem.Stat(ctx, refA) + require.NoError(t, err, "the stale immutable checkpoint may exist without becoming authoritative") + head, ok, err := first.GetArtifactCheckpointHead(ctx, contractOrigin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, sequenceB, head.Sequence) + pending, err := first.ArtifactExportClaims(ctx, []string{"session-a"}) + require.NoError(t, err) + require.Equal(t, claimA, pending) + + result, err := ExportToStore(ctx, first, filesystem, ExportOptions{Origin: contractOrigin}) + require.NoError(t, err) + assert.False(t, result.CheckpointCreated) + pending, err = first.ArtifactExportClaims(ctx, []string{"session-a"}) + require.NoError(t, err) + assert.Empty(t, pending) +} + +func (s *countingQueuedExportStore) PendingArtifactExports( + ctx context.Context, limit int, +) ([]db.ArtifactExportQueueItem, error) { + s.queueQueries++ + return s.database.PendingArtifactExports(ctx, limit) +} + +func (s *countingQueuedExportStore) GetSessionFull( + ctx context.Context, id string, +) (*db.Session, error) { + s.sessionLoads++ + return s.database.GetSessionFull(ctx, id) +} + +func (s *countingQueuedExportStore) GetAllMessages( + ctx context.Context, id string, +) ([]db.Message, error) { + s.messageLoads++ + return s.database.GetAllMessages(ctx, id) +} + +func (s *countingQueuedExportStore) GetUsageEvents( + ctx context.Context, id string, +) ([]db.UsageEvent, error) { + s.usageLoads++ + return s.database.GetUsageEvents(ctx, id) +} + +func TestArtifactExportCardinalityLoadsOnlyDirtyBatch(t *testing.T) { + for _, archiveSize := range []int{20, 2000} { + t.Run(fmt.Sprintf("archive-%d", archiveSize), func(t *testing.T) { + database := testDB(t) + for i := range archiveSize { + require.NoError(t, database.UpsertSession(db.Session{ + ID: fmt.Sprintf("peer-%04d", i), Project: "project", + Machine: "peer-a1b2c3", Agent: "claude", + })) + } + require.NoError(t, database.UpsertSession(db.Session{ + ID: "dirty", Project: "project", Machine: "local", Agent: "claude", + })) + require.NoError(t, database.ReplaceSessionMessages("dirty", []db.Message{{ + SessionID: "dirty", Ordinal: 0, Role: "user", Content: "changed", + }})) + require.NoError(t, database.ReplaceSessionUsageEvents("dirty", []db.UsageEvent{{ + SessionID: "dirty", Source: "event", Model: "model", DedupKey: "one", + }})) + + store := &countingQueuedExportStore{database: database} + visited := 0 + err := forEachQueuedArtifactExport(t.Context(), store, 1, func(work queuedArtifactExport) error { + visited++ + assert.Equal(t, "dirty", work.Item.SessionID) + require.NotNil(t, work.Session) + assert.Len(t, work.Messages, 1) + assert.Len(t, work.UsageEvents, 1) + return nil + }) + require.NoError(t, err) + assert.Equal(t, 1, visited) + assert.Equal(t, 1, store.queueQueries) + assert.Equal(t, 1, store.sessionLoads) + assert.Equal(t, 1, store.messageLoads) + assert.Equal(t, 1, store.usageLoads) + }) + } +} + +func TestExportToStoreCardinalityIgnoresUnrelatedArchiveBodies(t *testing.T) { + for _, archiveSize := range []int{20, 2000} { + t.Run(fmt.Sprintf("archive-%d", archiveSize), func(t *testing.T) { + database := testDB(t) + for i := range archiveSize { + require.NoError(t, database.UpsertSession(db.Session{ + ID: fmt.Sprintf("peer-%04d", i), Project: "project", + Machine: "peer-a1b2c3", Agent: "claude", + })) + } + seedSession(t, database, "dirty", "project") + counted := &countingCanonicalExportDB{DB: database} + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + + result, err := ExportToStore(t.Context(), counted, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.Equal(t, 1, result.ExportedSessions) + assert.Equal(t, 1, counted.sessionLoads) + assert.Equal(t, 1, counted.messageLoads) + assert.Equal(t, 1, counted.usageLoads) + }) + } +} + +func TestExportToStoreIncrementalBatchIsBoundedAndFullStreamsAllBodies(t *testing.T) { + database := testDB(t) + const total = artifactExportBatchSize + 5 + for i := range total { + seedSession(t, database, fmt.Sprintf("session-%03d", i), "project") + } + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + + result, err := ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.Equal(t, artifactExportBatchSize, result.ExportedSessions) + pending, err := database.PendingArtifactExports(t.Context(), 1024) + require.NoError(t, err) + assert.Len(t, pending, 5) + firstBytes := readContractArtifact(t, filesystem, + requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json")) + var first checkpoint + require.NoError(t, json.Unmarshal(firstBytes, &first)) + assert.Len(t, first.Sessions, artifactExportBatchSize) + + result, err = ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + Full: true, + }) + require.NoError(t, err) + assert.Equal(t, 5, result.ExportedSessions, + "full pass idempotently recreates clean dependencies and creates only missing manifests") + pending, err = database.PendingArtifactExports(t.Context(), 1024) + require.NoError(t, err) + assert.Empty(t, pending) + secondBytes := readContractArtifact(t, filesystem, + requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000002.json")) + var second checkpoint + require.NoError(t, json.Unmarshal(secondBytes, &second)) + assert.Len(t, second.Sessions, total) +} + +func TestExportToStoreFullDrainsMoreThanOneClaimPage(t *testing.T) { + database := testDB(t) + const total = 1025 + for i := range total { + require.NoError(t, database.UpsertSession(db.Session{ + ID: fmt.Sprintf("session-%04d", i), Project: "project", + Machine: "local", Agent: "claude", CreatedAt: "2026-06-14T01:02:03Z", + })) + } + counted := &countingCanonicalExportDB{DB: database} + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + + result, err := ExportToStore(t.Context(), counted, filesystem, ExportOptions{ + Origin: contractOrigin, Full: true, + }) + require.NoError(t, err) + assert.Equal(t, total, result.ExportedSessions) + assert.Equal(t, total, counted.sessionLoads, + "full export loads each body once instead of reloading the archive per page") + pending, err := database.PendingArtifactExports(t.Context(), 1024) + require.NoError(t, err) + assert.Empty(t, pending) + head, ok, err := database.GetArtifactCheckpointHead(t.Context(), contractOrigin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 2, head.Sequence) + body := readContractArtifact(t, filesystem, requireContractRef( + t, contractOrigin, KindCheckpoints, "cp-0000000002.json", + )) + var published checkpoint + require.NoError(t, json.Unmarshal(body, &published)) + assert.Len(t, published.Sessions, total) +} + +func TestExportToStoreExplicitSessionIDsClaimBeyondOldestQueuePage(t *testing.T) { + database := testDB(t) + const total = 1025 + for i := range total { + require.NoError(t, database.UpsertSession(db.Session{ + ID: fmt.Sprintf("session-%04d", i), Project: "project", + Machine: "local", Agent: "claude", + })) + } + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + + result, err := ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, SessionIDs: []string{"session-1024"}, + }) + require.NoError(t, err) + assert.Equal(t, 1, result.ExportedSessions) + pending, err := database.PendingArtifactExports(t.Context(), 1024) + require.NoError(t, err) + assert.Len(t, pending, 1024) + body := readContractArtifact(t, filesystem, + requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json")) + var published checkpoint + require.NoError(t, json.Unmarshal(body, &published)) + assert.Equal(t, map[string]string{ + contractOrigin + "~session-1024": published.Sessions[contractOrigin+"~session-1024"], + }, published.Sessions) +} + +func TestExportToStorePublishesEmptyAndDeletionSets(t *testing.T) { + t.Run("empty", func(t *testing.T) { + database := testDB(t) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + + result, err := ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, Full: true, + }) + require.NoError(t, err) + assert.True(t, result.CheckpointCreated) + body := readContractArtifact(t, filesystem, + requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000001.json")) + assert.Equal(t, + `{"origin":"contract-a1b2c3","seq":1,"sessions":{},"v":1}`+"\n", + string(body)) + }) + + t.Run("deletion", func(t *testing.T) { + database := testDB(t) + seedSession(t, database, "sess-1", "project") + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + _, err = ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + require.NoError(t, database.SoftDeleteSession("sess-1")) + + result, err := ExportToStore(t.Context(), database, filesystem, ExportOptions{ + Origin: contractOrigin, + }) + require.NoError(t, err) + assert.True(t, result.CheckpointCreated) + body := readContractArtifact(t, filesystem, + requireContractRef(t, contractOrigin, KindCheckpoints, "cp-0000000002.json")) + assert.Contains(t, string(body), `"sessions":{}`) + }) +} + +func TestExactImportDefersCheckpointWithAbsentDependency(t *testing.T) { + origin := "peer-a1b2c3" + source := testDB(t) + seedSession(t, source, "session-1", "alpha") + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + _, err = ExportToStore(t.Context(), source, store, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + + segments, err := firstStoreEntryPage(t.Context(), store, origin, KindSegments, 10) + require.NoError(t, err) + require.Len(t, segments.Items, 1) + require.NoError(t, store.Trash(t.Context(), segments.Items[0].Ref)) + + target := testDB(t) + result, err := importResultFromTestStore( + t.Context(), target, store, "local-d4e5f6", + ) + require.NoError(t, err) + assert.Equal(t, 1, result.Deferred) + assert.False(t, result.Changed()) + landed, landedMap, ok, err := target.GetArtifactCheckpointLanding(t.Context(), origin) + require.NoError(t, err) + assert.False(t, ok) + assert.Empty(t, landed) + assert.Empty(t, landedMap) +} + +func TestImportCoordinatorRetriesExactCheckpointAfterRestartAndDependencyArrival(t *testing.T) { + origin := "peer-a1b2c3" + source := testDB(t) + seedSession(t, source, "session-1", "alpha") + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + _, err = ExportToStore(t.Context(), source, store, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + + segments, err := firstStoreEntryPage(t.Context(), store, origin, KindSegments, 10) + require.NoError(t, err) + require.Len(t, segments.Items, 1) + segment := segments.Items[0] + segmentBody := readContractArtifact(t, store, segment.Ref) + require.NoError(t, store.Trash(t.Context(), segment.Ref)) + checkpoints, err := firstStoreEntryPage(t.Context(), store, origin, KindCheckpoints, 10) + require.NoError(t, err) + require.Len(t, checkpoints.Items, 1) + + databasePath := filepath.Join(t.TempDir(), "target.db") + target, err := db.Open(databasePath) + require.NoError(t, err) + t.Cleanup(func() { + if target != nil { + require.NoError(t, target.Close()) + } + }) + coordinator := NewStoreImportCoordinator(target, store, "local-d4e5f6") + require.NoError(t, coordinator.RecordChanged(t.Context(), checkpoints.Items[0])) + result, err := coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.Equal(t, 1, result.Deferred) + assert.False(t, result.Changed()) + + require.NoError(t, target.Close()) + target = nil + target, err = db.Open(databasePath) + require.NoError(t, err) + pending, err := target.PendingArtifactImports(t.Context(), formatVersion, 10) + require.NoError(t, err) + require.Len(t, pending, 1, "unfinished checkpoint work must survive restart") + + created, err := store.Create( + t.Context(), segment.Ref, segment.Identity, "application/x-ndjson", + bytes.NewReader(segmentBody), + ) + require.NoError(t, err) + coordinator = NewStoreImportCoordinator(target, store, "local-d4e5f6") + require.NoError(t, coordinator.RecordChanged(t.Context(), created.Entry)) + result, err = coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.Equal(t, 1, result.Sessions) + assert.Equal(t, 2, result.Messages) + assert.Zero(t, result.Deferred) + pending, err = target.PendingArtifactImports(t.Context(), formatVersion, 10) + require.NoError(t, err) + assert.Empty(t, pending) + landed, err := target.GetSession(t.Context(), origin+"~session-1") + require.NoError(t, err) + require.NotNil(t, landed) +} + +type corruptUntilRepairedStore struct { + ArtifactStore + corruptRef Ref + repaired bool + originLists int +} + +func (s *corruptUntilRepairedStore) Origins(ctx context.Context) (OriginIterator, error) { + s.originLists++ + return s.ArtifactStore.Origins(ctx) +} + +func (s *corruptUntilRepairedStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + if ref == s.corruptRef && !s.repaired { + entry, err := s.Stat(ctx, ref) + if err != nil { + return Entry{}, nil, err + } + return entry, &alwaysCorruptVerifiedReader{}, nil + } + return s.ArtifactStore.Open(ctx, ref) +} + +func (s *corruptUntilRepairedStore) RepairContent( + ctx context.Context, identity Identity, trusted io.Reader, +) error { + repairer, ok := s.ArtifactStore.(interface { + RepairContent(context.Context, Identity, io.Reader) error + }) + if !ok { + return errors.New("wrapped store cannot repair content") + } + if err := repairer.RepairContent(ctx, identity, trusted); err != nil { + return err + } + s.repaired = true + return nil +} + +type alwaysCorruptVerifiedReader struct{} + +func (*alwaysCorruptVerifiedReader) Read([]byte) (int, error) { + return 0, ErrArtifactCorrupt +} + +func (*alwaysCorruptVerifiedReader) Verify() error { return ErrArtifactCorrupt } + +func (*alwaysCorruptVerifiedReader) Close() error { return nil } + +func TestExactImportRepairsPhysicalCorruptionFromTrustedPeer(t *testing.T) { + origin := "peer-a1b2c3" + source := testDB(t) + seedSession(t, source, "session-1", "alpha") + _, canonical := newTestDocbankStore(t, docbank.Config{}) + _, err := ExportToStore(t.Context(), source, canonical, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + segments, err := firstStoreEntryPage(t.Context(), canonical, origin, KindSegments, 10) + require.NoError(t, err) + require.Len(t, segments.Items, 1) + segment := segments.Items[0] + _, trustedReader, err := canonical.Open(t.Context(), segment.Ref) + require.NoError(t, err) + trusted, err := io.ReadAll(trustedReader) + require.NoError(t, err) + require.NoError(t, trustedReader.Verify()) + require.NoError(t, trustedReader.Close()) + + store := &corruptUntilRepairedStore{ + ArtifactStore: canonical, + corruptRef: segment.Ref, + } + target := testDB(t) + _, err = importResultFromTestStore(t.Context(), target, store, "local-d4e5f6") + require.ErrorIs(t, err, ErrArtifactCorrupt) + pending, err := target.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, db.ArtifactRepair{ + Origin: origin, + Kind: string(KindSegments), + Name: segment.Ref.Name, + SHA256: segment.Identity.SHA256, + Size: segment.Identity.Size, + }, db.ArtifactRepair{ + Origin: pending[0].Origin, + Kind: pending[0].Kind, + Name: pending[0].Name, + SHA256: pending[0].SHA256, + Size: pending[0].Size, + }) + + coordinator := NewStoreImportCoordinator(target, store, "local-d4e5f6") + require.NoError(t, RepairArtifactFromTrustedPeer( + t.Context(), target, store, pending[0], bytes.NewReader(trusted), coordinator, + )) + pending, err = target.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, err) + assert.Empty(t, pending) + coordinator.requestDrain() + + result, err := coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.Equal(t, 1, result.Sessions) + assert.Equal(t, 2, result.Messages) + assert.Equal(t, 1, store.originLists, + "repair retry must open the queued checkpoint without another origin scan") + + result, err = coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.False(t, result.Changed()) + assert.Equal(t, 1, store.originLists, "a drained coordinator must not run another import") +} + +func TestImportCoordinatorIgnoresLandedHistory(t *testing.T) { + for _, historySize := range []int{10, 10_000} { + t.Run(fmt.Sprintf("history-%d", historySize), func(t *testing.T) { + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + store := &exactImportCountingStore{ArtifactStore: base} + coordinator := NewStoreImportCoordinator( + testDB(t), store, "local-d4e5f6", + ) + require.NoError(t, coordinator.requestDrain()) + + result, err := coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.False(t, result.Changed()) + assert.Zero(t, store.origins.Load()) + assert.Zero(t, store.entries.Load()) + assert.Zero(t, store.opens.Load()) + }) + } +} + +func TestImportCoordinatorProcessesChangedMetadata(t *testing.T) { + for _, historySize := range []int{10, 10_000} { + t.Run(fmt.Sprintf("history-%d", historySize), func(t *testing.T) { + origin := "peer-a1b2c3" + localOrigin := "local-d4e5f6" + database := testDB(t) + for i := range historySize { + require.NoError(t, database.MarkMetadataEventApplied( + t.Context(), origin, fmt.Sprintf("history-%010d", i), strings.Repeat("a", 64), + )) + } + gid := origin + "~session-1" + require.NoError(t, database.UpsertSession(db.Session{ + ID: gid, Project: "project", Machine: origin, Agent: "claude", + })) + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + ref := createStoreMetadataEvent(t, base, origin, replayRenameEvent( + t, origin, gid, replayTestHLC(0, 0), "renamed", + )) + entry, err := base.Stat(t.Context(), ref) + require.NoError(t, err) + store := &exactImportCountingStore{ArtifactStore: base} + coordinator := NewStoreImportCoordinator(database, store, localOrigin) + require.NoError(t, coordinator.RecordChanged(t.Context(), entry)) + + result, err := coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.Equal(t, 1, result.Metadata) + assert.Zero(t, store.origins.Load()) + assert.Zero(t, store.entries.Load()) + assert.Equal(t, int32(1), store.opens.Load()) + updated, err := database.GetSession(t.Context(), gid) + require.NoError(t, err) + require.NotNil(t, updated) + require.NotNil(t, updated.DisplayName) + assert.Equal(t, "renamed", *updated.DisplayName) + }) + } +} + +func TestStoreImportCoordinatorSerializesFinalizersAndRetainsMidRunWork(t *testing.T) { + origin := "peer-a1b2c3" + localOrigin := "local-d4e5f6" + database := testDB(t) + gid := origin + "~session-1" + require.NoError(t, database.UpsertSession(db.Session{ + ID: gid, Project: "project", Machine: origin, Agent: "claude", + })) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + ref := createStoreMetadataEvent(t, filesystem, origin, replayRenameEvent( + t, origin, gid, replayTestHLC(0, 0), "serialized", + )) + entry, err := filesystem.Stat(t.Context(), ref) + require.NoError(t, err) + store := &exactImportGateStore{ + ArtifactStore: filesystem, + entered: make(chan struct{}), + release: make(chan struct{}), + } + coordinator := NewStoreImportCoordinator(database, store, localOrigin) + require.NoError(t, coordinator.RecordChanged(t.Context(), entry)) + + firstDone := make(chan error, 1) + go func() { + _, err := coordinator.Finalize(t.Context()) + firstDone <- err + }() + <-store.entered + coordinator.requestDrain() + secondStarted := make(chan struct{}) + secondDone := make(chan error, 1) + go func() { + close(secondStarted) + _, err := coordinator.Finalize(t.Context()) + secondDone <- err + }() + <-secondStarted + assert.Never(t, func() bool { + return store.maximum.Load() > 1 + }, 100*time.Millisecond, time.Millisecond, + "finalizers must not overlap exact artifact reads") + close(store.release) + require.NoError(t, <-firstDone) + require.NoError(t, <-secondDone) + assert.Equal(t, int32(1), store.maximum.Load()) + assert.Equal(t, int32(1), store.opens.Load(), + "the mid-run generation drains after the first exact claim is acknowledged") +} + +func TestStoreImportCoordinatorRetriesTransientFinalizeOnceAndDrains(t *testing.T) { + origin := "peer-a1b2c3" + localOrigin := "local-d4e5f6" + database := testDB(t) + gid := origin + "~session-1" + require.NoError(t, database.UpsertSession(db.Session{ + ID: gid, Project: "project", Machine: origin, Agent: "claude", + })) + filesystem, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, filesystem.Close()) }) + ref := createStoreMetadataEvent(t, filesystem, origin, replayRenameEvent( + t, origin, gid, replayTestHLC(0, 0), "retry", + )) + entry, err := filesystem.Stat(t.Context(), ref) + require.NoError(t, err) + store := &transientImportStore{ArtifactStore: filesystem} + coordinator := NewStoreImportCoordinator(database, store, localOrigin) + require.NoError(t, coordinator.RecordChanged(t.Context(), entry)) + + _, err = coordinator.Finalize(t.Context()) + require.ErrorIs(t, err, syscall.EAGAIN) + result, err := coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.Equal(t, 1, result.Metadata) + assert.Equal(t, int32(2), store.calls.Load(), + "one failed exact open and one successful retry are expected") + + result, err = coordinator.Finalize(t.Context()) + require.NoError(t, err) + assert.False(t, result.Changed()) + assert.Equal(t, int32(2), store.calls.Load(), "successful retry must drain pending work") +} + +func TestRepairArtifactFromTrustedPeerRejectsTypedNilRetryCoordinator(t *testing.T) { + database := testDB(t) + repair := db.ArtifactRepair{ + Origin: "peer-a1b2c3", + Kind: string(KindRaw), + Name: strings.Repeat("a", 64), + SHA256: strings.Repeat("a", 64), + Size: 1, + } + require.NoError(t, database.EnqueueArtifactRepair(t.Context(), repair)) + store := &trackingRepairStore{} + var retry *StoreImportCoordinator + + err := RepairArtifactFromTrustedPeer( + t.Context(), database, store, repair, strings.NewReader("x"), retry, + ) + require.Error(t, err) + assert.Zero(t, store.repairs, "invalid retry coordination must fail before repair side effects") + pending, pendingErr := database.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, pendingErr) + assert.Len(t, pending, 1, "an unobservable repair must not acknowledge its durable claim") +} + +func TestRepairArtifactFromTrustedPeerLeavesClaimPendingWhenSchedulingFails(t *testing.T) { + database := testDB(t) + repair := db.ArtifactRepair{ + Origin: "peer-a1b2c3", + Kind: string(KindRaw), + Name: strings.Repeat("a", 64), + SHA256: strings.Repeat("a", 64), + Size: 1, + } + require.NoError(t, database.EnqueueArtifactRepair(t.Context(), repair)) + store := &trackingRepairStore{} + scheduleErr := errors.New("retry scheduler unavailable") + scheduler := &failingImportScheduler{err: scheduleErr} + + err := RepairArtifactFromTrustedPeer( + t.Context(), database, store, repair, strings.NewReader("x"), scheduler, + ) + assert.ErrorIs(t, err, scheduleErr) + assert.Equal(t, 1, store.repairs, "physical repair completes before retry scheduling") + assert.Equal(t, 1, scheduler.calls) + pending, pendingErr := database.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, pendingErr) + assert.Len(t, pending, 1, "failed scheduling must prevent durable claim acknowledgement") +} + +func TestExactImportReplaysMetadataAfterSessionContent(t *testing.T) { + origin := "peer-a1b2c3" + source := testDB(t) + seedSession(t, source, "session-1", "alpha") + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + _, err = ExportToStore(t.Context(), source, store, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + recorder := NewMetadataRecorder(source, MetadataRecorderOptions{ + Store: store, + Origin: origin, + Now: fixedHLCTime, + }) + _, err = recorder.Append(t.Context(), MetadataEventInput{ + SessionID: "session-1", + Op: MetadataOpRename, + Value: json.RawMessage(`{"display_name":"Renamed by peer"}`), + }) + require.NoError(t, err) + + target := testDB(t) + result, err := importResultFromTestStore( + t.Context(), target, store, "local-d4e5f6", + ) + require.NoError(t, err) + assert.Equal(t, 1, result.Sessions) + assert.Equal(t, 1, result.Metadata) + imported, err := target.GetSessionFull(t.Context(), origin+"~session-1") + require.NoError(t, err) + require.NotNil(t, imported) + require.NotNil(t, imported.DisplayName) + assert.Equal(t, "Renamed by peer", *imported.DisplayName) +} + +func TestCheckpointLandingStatusUsesExactRecordedCheckpointMap(t *testing.T) { + origin := "peer-a1b2c3" + store := newTestArtifactStore(t) + source := testDB(t) + seedSession(t, source, "alpha", "project") + seedSession(t, source, "bravo", "project") + _, err := ExportToStore(t.Context(), source, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + _, cp, err := latestStoreCheckpointSummary(t.Context(), store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + + target := testDB(t) + for gid, manifestHash := range cp.Sessions { + require.NoError(t, target.SetSyncState(importStateKey(origin, gid), manifestHash)) + } + status, err := CheckpointLandingStatusFromStore(t.Context(), store, origin, target, false) + require.NoError(t, err) + require.True(t, status.Found) + assert.Zero(t, status.LandedSessionCount, "legacy sync state is not exact landing provenance") + + require.NoError(t, target.RecordArtifactCheckpointLanding( + t.Context(), + db.ArtifactCheckpointLanding{Origin: origin, Sequence: cp.Sequence}, + cp.Sessions, + )) + + status, err = CheckpointLandingStatusFromStore(t.Context(), store, origin, target, false) + require.NoError(t, err) + require.True(t, status.Found) + assert.Equal(t, cp.Sequence, status.Sequence) + assert.Equal(t, len(cp.Sessions), status.LandedSessionCount) +} + +func TestCheckpointLandingStatusUsesExactLocalPublicationMap(t *testing.T) { + origin := "local-a1b2c3" + store := newTestArtifactStore(t) + database := testDB(t) + seedSession(t, database, "alpha", "project") + seedSession(t, database, "bravo", "project") + _, err := ExportToStore(t.Context(), database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + + status, err := CheckpointLandingStatusFromStore(t.Context(), store, origin, database, true) + require.NoError(t, err) + require.True(t, status.Found) + assert.Equal(t, 2, status.LandedSessionCount) +} + +type legacyCheckpointStatusStore struct { + values map[string]string +} + +func (s *legacyCheckpointStatusStore) SyncStateValues(keys []string) (map[string]string, error) { + result := make(map[string]string, len(keys)) + for _, key := range keys { + result[key] = s.values[key] + } + return result, nil +} + +type publicationRevisionRaceStore struct { + *db.DB +} + +func (s *publicationRevisionRaceStore) StreamArtifactPublications( + ctx context.Context, origin string, visit func(db.ArtifactPublication) error, +) (int64, error) { + revision, err := s.DB.StreamArtifactPublications(ctx, origin, visit) + return revision + 1, err +} + +type streamingLandingStatusStore struct { + *db.DB + streamed int +} + +func (s *streamingLandingStatusStore) GetArtifactCheckpointLanding( + context.Context, string, +) (db.ArtifactCheckpointLanding, map[string]string, bool, error) { + return db.ArtifactCheckpointLanding{}, nil, false, + errors.New("checkpoint status must not materialize the landing map") +} + +func (s *streamingLandingStatusStore) StreamArtifactCheckpointLanding( + ctx context.Context, origin string, visit func(string, string) error, +) (db.ArtifactCheckpointLanding, bool, error) { + landing, manifests, found, err := s.DB.GetArtifactCheckpointLanding(ctx, origin) + if err != nil || !found { + return landing, found, err + } + keys := make([]string, 0, len(manifests)) + for gid := range manifests { + keys = append(keys, gid) + } + sort.Strings(keys) + for _, gid := range keys { + if err := visit(gid, manifests[gid]); err != nil { + return db.ArtifactCheckpointLanding{}, false, err + } + s.streamed++ + } + return landing, true, nil +} + +func TestCheckpointLandingStatusDoesNotFallbackToLegacySyncState(t *testing.T) { + origin := "peer-a1b2c3" + store := newTestArtifactStore(t) + source := testDB(t) + seedSession(t, source, "alpha", "project") + _, err := ExportToStore(t.Context(), source, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + _, cp, err := latestStoreCheckpointSummary(t.Context(), store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + legacy := &legacyCheckpointStatusStore{values: map[string]string{}} + for gid, manifestHash := range cp.Sessions { + legacy.values[importStateKey(origin, gid)] = manifestHash + } + + status, err := CheckpointLandingStatusFromStore(t.Context(), store, origin, legacy, false) + require.NoError(t, err) + assert.Zero(t, status.LandedSessionCount) +} + +func TestCheckpointLandingStatusRejectsMixedPublicationRevisions(t *testing.T) { + origin := "local-a1b2c3" + store := newTestArtifactStore(t) + database := testDB(t) + seedSession(t, database, "alpha", "project") + _, err := ExportToStore(t.Context(), database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + + status, err := CheckpointLandingStatusFromStore( + t.Context(), store, origin, &publicationRevisionRaceStore{DB: database}, true, + ) + require.NoError(t, err) + assert.Zero(t, status.LandedSessionCount, + "rows from a different publication revision are not coherent landing evidence") +} + +func TestCheckpointLandingStatusStreamsLargeForeignLanding(t *testing.T) { + origin := "peer-a1b2c3" + artifactStore := newTestArtifactStore(t) + manifests := make(map[string]string, artifactImportPageSize*3) + for i := range artifactImportPageSize * 3 { + manifests[fmt.Sprintf("%s~session-%04d", origin, i)] = fmt.Sprintf("%064x", i+1) + } + cp := checkpoint{ + Version: formatVersion, Origin: origin, Sequence: 1, Sessions: manifests, + } + data, err := canonicalJSON(cp) + require.NoError(t, err) + ref, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + createContractArtifact(t, artifactStore, ref, data) + database := testDB(t) + require.NoError(t, database.RecordArtifactCheckpointLanding( + t.Context(), db.ArtifactCheckpointLanding{Origin: origin, Sequence: 1}, manifests, + )) + store := &streamingLandingStatusStore{DB: database} + + status, err := CheckpointLandingStatusFromStore(t.Context(), artifactStore, origin, store, false) + require.NoError(t, err) + assert.Equal(t, len(manifests), status.LandedSessionCount) + assert.Equal(t, len(manifests), store.streamed) +} + +func TestCheckpointLandingStatusHonorsCanceledRequestContext(t *testing.T) { + origin := "peer-a1b2c3" + store := newTestArtifactStore(t) + source := testDB(t) + seedSession(t, source, "alpha", "project") + _, err := ExportToStore(t.Context(), source, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + target := testDB(t) + _, cp, err := latestStoreCheckpointSummary(t.Context(), store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + require.NoError(t, target.RecordArtifactCheckpointLanding( + t.Context(), db.ArtifactCheckpointLanding{Origin: origin, Sequence: cp.Sequence}, cp.Sessions, + )) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err = CheckpointLandingStatusFromStore(ctx, store, origin, target, false) + require.ErrorIs(t, err, context.Canceled) +} + +func TestExactImportQuarantinesInvalidManifestAndSegment(t *testing.T) { + tests := []struct { + name string + invalidRef func(*testing.T, ArtifactStore, string) (Ref, string) + }{ + { + name: "manifest", + invalidRef: func(t *testing.T, store ArtifactStore, origin string) (Ref, string) { + data := []byte("not-json\n") + hash := hashHex(data) + ref, err := NewRef(origin, KindManifests, hash+".json") + require.NoError(t, err) + identity, err := NewIdentity(hash, int64(len(data))) + require.NoError(t, err) + _, err = store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(KindManifests), bytes.NewReader(data)) + require.NoError(t, err) + return ref, hash + }, + }, + { + name: "segment", + invalidRef: func(t *testing.T, store ArtifactStore, origin string) (Ref, string) { + segmentData := []byte("not-ndjson\n") + segmentHash := hashHex(segmentData) + segmentRef, err := NewRef(origin, KindSegments, segmentHash+".ndjson") + require.NoError(t, err) + identity, err := NewIdentity(segmentHash, int64(len(segmentData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), segmentRef, identity, + canonicalArtifactMediaType(KindSegments), bytes.NewReader(segmentData)) + require.NoError(t, err) + m := manifest{ + Version: formatVersion, Origin: origin, NativeSessionID: "session-1", + Session: manifestSession{ID: "session-1", Machine: origin, Agent: "claude", Project: "alpha"}, + Segments: []string{segmentHash}, + } + manifestData, err := canonicalJSON(m) + require.NoError(t, err) + manifestHash := hashHex(manifestData) + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + manifestIdentity, err := NewIdentity(manifestHash, int64(len(manifestData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), manifestRef, manifestIdentity, + canonicalArtifactMediaType(KindManifests), bytes.NewReader(manifestData)) + require.NoError(t, err) + return segmentRef, manifestHash + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + origin := "peer-a1b2c3" + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + invalidRef, manifestHash := tt.invalidRef(t, store, origin) + cp := checkpoint{ + Version: formatVersion, Origin: origin, Sequence: 1, + Sessions: map[string]string{origin + "~session-1": manifestHash}, + } + checkpointData, err := canonicalJSON(cp) + require.NoError(t, err) + checkpointRef, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + checkpointIdentity, err := NewIdentity(hashHex(checkpointData), int64(len(checkpointData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), checkpointRef, checkpointIdentity, + canonicalArtifactMediaType(KindCheckpoints), bytes.NewReader(checkpointData)) + require.NoError(t, err) + + result, err := importResultFromTestStore( + t.Context(), testDB(t), store, "local-d4e5f6", + ) + require.NoError(t, err) + assert.Equal(t, 1, result.Deferred) + _, err = store.Stat(t.Context(), invalidRef) + assert.ErrorIs(t, err, ErrArtifactNotFound) + }) + } +} + +func TestExactImportProcessesOriginPagesIncrementally(t *testing.T) { + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + store := &originTraversalGateStore{ArtifactStore: base} + + _, err = importResultFromTestStore(t.Context(), testDB(t), store, "local-d4e5f6") + require.NoError(t, err) + assert.True(t, store.processedFirstPage) +} + +func TestExactImportEnumeratesCheckpointHeadersBeforeBodies(t *testing.T) { + origin := "peer-a1b2c3" + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + for sequence := 1; sequence <= artifactImportPageSize+1; sequence++ { + cp := checkpoint{ + Version: formatVersion, Origin: origin, Sequence: sequence, + Sessions: map[string]string{}, + } + data, marshalErr := canonicalJSON(cp) + require.NoError(t, marshalErr) + ref, refErr := NewRef(origin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequence)) + require.NoError(t, refErr) + identity, identityErr := NewIdentity(hashHex(data), int64(len(data))) + require.NoError(t, identityErr) + _, createErr := base.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(KindCheckpoints), bytes.NewReader(data)) + require.NoError(t, createErr) + } + store := &checkpointHeaderFirstStore{ArtifactStore: base} + target := testDB(t) + + _, err = importResultFromTestStore(t.Context(), target, store, "local-d4e5f6") + require.NoError(t, err) + assert.True(t, store.listedAll) + landing, _, found, err := target.GetArtifactCheckpointLanding(t.Context(), origin) + require.NoError(t, err) + require.True(t, found) + assert.Equal(t, artifactImportPageSize+1, landing.Sequence) +} + +func TestExactImportNewestCheckpointSkipsHistoricalClosures(t *testing.T) { + origin := "peer-a1b2c3" + source := testDB(t) + seedSession(t, source, "session-1", "alpha") + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + _, err = ExportToStore(t.Context(), source, base, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + firstPage, err := firstStoreEntryPage(t.Context(), base, origin, KindCheckpoints, 1) + require.NoError(t, err) + require.Len(t, firstPage.Items, 1) + _, reader, err := base.Open(t.Context(), firstPage.Items[0].Ref) + require.NoError(t, err) + data, err := io.ReadAll(reader) + require.NoError(t, err) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + var cp checkpoint + require.NoError(t, json.Unmarshal(data, &cp)) + for sequence := 2; sequence <= artifactImportPageSize*2+1; sequence++ { + cp.Sequence = sequence + checkpointData, marshalErr := canonicalJSON(cp) + require.NoError(t, marshalErr) + ref, refErr := NewRef(origin, KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequence)) + require.NoError(t, refErr) + identity, identityErr := NewIdentity(hashHex(checkpointData), int64(len(checkpointData))) + require.NoError(t, identityErr) + _, createErr := base.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(KindCheckpoints), bytes.NewReader(checkpointData)) + require.NoError(t, createErr) + } + store := &dependencyOpenCountingStore{ + ArtifactStore: base, + opens: make(map[Kind]int), + } + + result, err := importResultFromTestStore(t.Context(), testDB(t), store, "local-d4e5f6") + require.NoError(t, err) + assert.Equal(t, 1, result.Sessions) + assert.LessOrEqual(t, store.opens[KindManifests], 2, + "a complete newest checkpoint must not open historical manifests") + assert.LessOrEqual(t, store.opens[KindSegments], 2, + "a complete newest checkpoint must not open historical segments") + assert.LessOrEqual(t, store.opens[KindCheckpoints], 1, + "checkpoint bodies must be inspected newest-first") +} + +func createFutureMetadataEntries( + t *testing.T, store ArtifactStore, origin string, count int, +) { + t.Helper() + for i := range count { + stamp := HLCTimestamp{ + WallTime: fixedHLCTime().Add(time.Duration(i) * time.Nanosecond), + } + event := metadataEvent{ + Version: formatVersion + 1, + HLC: stamp.String(), Origin: origin, + SessionGID: origin + "~future", Op: "future-op", + } + data, err := canonicalJSON(event) + require.NoError(t, err) + hash := hashHex(data) + ref, err := NewRef(origin, KindMeta, + stamp.OrderingKey(hash)+metadataEventExtension) + require.NoError(t, err) + identity, err := NewIdentity(hash, int64(len(data))) + require.NoError(t, err) + _, err = store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(KindMeta), bytes.NewReader(data)) + require.NoError(t, err) + } +} + +func TestExactImportCancelsBetweenMetadataEntries(t *testing.T) { + origin := "peer-a1b2c3" + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + createFutureMetadataEntries(t, base, origin, 2) + ctx, cancel := context.WithCancel(t.Context()) + store := &cancelAfterMetadataOpenStore{ArtifactStore: base, cancel: cancel} + + _, err = importResultFromTestStore(ctx, testDB(t), store, "local-d4e5f6") + require.ErrorIs(t, err, context.Canceled) + assert.Equal(t, 1, store.opens, "cancellation must stop before opening another event") +} + +func TestExactImportDefersFutureManifestWithoutQuarantine(t *testing.T) { + origin := "peer-a1b2c3" + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + m := manifest{Version: formatVersion + 1, Origin: origin, NativeSessionID: "session-1"} + manifestData, err := canonicalJSON(m) + require.NoError(t, err) + manifestHash := hashHex(manifestData) + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + manifestIdentity, err := NewIdentity(manifestHash, int64(len(manifestData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), manifestRef, manifestIdentity, + canonicalArtifactMediaType(KindManifests), bytes.NewReader(manifestData)) + require.NoError(t, err) + cp := checkpoint{ + Version: formatVersion, Origin: origin, Sequence: 1, + Sessions: map[string]string{origin + "~session-1": manifestHash}, + } + checkpointData, err := canonicalJSON(cp) + require.NoError(t, err) + checkpointRef, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + checkpointIdentity, err := NewIdentity(hashHex(checkpointData), int64(len(checkpointData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), checkpointRef, checkpointIdentity, + canonicalArtifactMediaType(KindCheckpoints), bytes.NewReader(checkpointData)) + require.NoError(t, err) + + result, err := importResultFromTestStore( + t.Context(), testDB(t), store, "local-d4e5f6", + ) + require.NoError(t, err) + assert.Equal(t, 1, result.Deferred) + _, err = store.Stat(t.Context(), manifestRef) + assert.NoError(t, err, "future protocol artifacts must remain live") +} + +func TestExactImportDefersFutureSegmentWithoutQuarantine(t *testing.T) { + origin := "peer-a1b2c3" + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + segmentData, err := canonicalJSON(segmentMessage{ + Version: formatVersion + 1, Ordinal: 0, Role: "user", Content: "future", + }) + require.NoError(t, err) + segmentHash := hashHex(segmentData) + segmentRef, err := NewRef(origin, KindSegments, segmentHash+".ndjson") + require.NoError(t, err) + segmentIdentity, err := NewIdentity(segmentHash, int64(len(segmentData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), segmentRef, segmentIdentity, + canonicalArtifactMediaType(KindSegments), bytes.NewReader(segmentData)) + require.NoError(t, err) + m := manifest{ + Version: formatVersion, Origin: origin, NativeSessionID: "session-1", + Session: manifestSession{ID: "session-1", Machine: origin, Agent: "claude", Project: "alpha"}, + Segments: []string{segmentHash}, + } + manifestData, err := canonicalJSON(m) + require.NoError(t, err) + manifestHash := hashHex(manifestData) + manifestRef, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + manifestIdentity, err := NewIdentity(manifestHash, int64(len(manifestData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), manifestRef, manifestIdentity, + canonicalArtifactMediaType(KindManifests), bytes.NewReader(manifestData)) + require.NoError(t, err) + cp := checkpoint{ + Version: formatVersion, Origin: origin, Sequence: 1, + Sessions: map[string]string{origin + "~session-1": manifestHash}, + } + checkpointData, err := canonicalJSON(cp) + require.NoError(t, err) + checkpointRef, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + checkpointIdentity, err := NewIdentity(hashHex(checkpointData), int64(len(checkpointData))) + require.NoError(t, err) + _, err = store.Create(t.Context(), checkpointRef, checkpointIdentity, + canonicalArtifactMediaType(KindCheckpoints), bytes.NewReader(checkpointData)) + require.NoError(t, err) + + result, err := importResultFromTestStore( + t.Context(), testDB(t), store, "local-d4e5f6", + ) + require.NoError(t, err) + assert.Equal(t, 1, result.Deferred) + _, err = store.Stat(t.Context(), segmentRef) + assert.NoError(t, err, "future protocol segments must remain live") +} + +func TestExactImportCountsFutureCheckpointAsDeferred(t *testing.T) { + origin := "peer-a1b2c3" + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + future := checkpoint{ + Version: formatVersion + 1, Origin: origin, Sequence: 1, + Sessions: map[string]string{origin + "~future": strings.Repeat("f", 64)}, + } + data, err := canonicalJSON(future) + require.NoError(t, err) + ref, err := NewRef(origin, KindCheckpoints, "cp-0000000001.json") + require.NoError(t, err) + identity, err := NewIdentity(hashHex(data), int64(len(data))) + require.NoError(t, err) + _, err = store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(KindCheckpoints), bytes.NewReader(data)) + require.NoError(t, err) + + result, err := importResultFromTestStore( + t.Context(), testDB(t), store, "local-d4e5f6", + ) + require.NoError(t, err) + assert.Equal(t, 1, result.Deferred) + _, err = store.Stat(t.Context(), ref) + assert.NoError(t, err) +} + +func TestExactImportLandsUpdatedLocallyTrashedSession(t *testing.T) { + origin := "peer-a1b2c3" + source := testDB(t) + seedSession(t, source, "session-1", "alpha") + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + _, err = ExportToStore(t.Context(), source, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + target := testDB(t) + _, err = importResultFromTestStore(t.Context(), target, store, "local-d4e5f6") + require.NoError(t, err) + importedID := origin + "~session-1" + require.NoError(t, target.SoftDeleteSession(importedID)) + + seedSession(t, source, "session-1", "updated") + _, err = ExportToStore(t.Context(), source, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + result, err := importResultFromTestStore(t.Context(), target, store, "local-d4e5f6") + require.NoError(t, err) + assert.False(t, result.Changed(), "import must not resurrect a locally trashed replica") + trashed, err := target.GetSessionFull(t.Context(), importedID) + require.NoError(t, err) + require.NotNil(t, trashed) + assert.NotNil(t, trashed.DeletedAt) + assert.Equal(t, "alpha", trashed.Project, + "the newer peer manifest must not rewrite locally trashed content") + landing, manifests, ok, err := target.GetArtifactCheckpointLanding(t.Context(), origin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 2, landing.Sequence) + assert.NotEmpty(t, manifests[importedID]) +} + +func TestQueuedArtifactExportTreatsOwnershipLossAsDeletion(t *testing.T) { + database := testDB(t) + require.NoError(t, database.UpsertSession(db.Session{ + ID: "moved", Project: "project", Machine: "local", Agent: "claude", + })) + require.NoError(t, database.ReplaceSessionMessages("moved", []db.Message{{ + SessionID: "moved", Ordinal: 0, Role: "user", Content: "local content", + }})) + require.NoError(t, database.UpsertSession(db.Session{ + ID: "moved", Project: "project", Machine: "peer-a1b2c3", Agent: "claude", + })) + + store := &countingQueuedExportStore{database: database} + visited := 0 + err := forEachQueuedArtifactExport(t.Context(), store, 1, func(work queuedArtifactExport) error { + visited++ + assert.Nil(t, work.Session) + assert.Empty(t, work.Messages) + assert.Empty(t, work.UsageEvents) + return nil + }) + require.NoError(t, err) + assert.Equal(t, 1, visited) + assert.Equal(t, 1, store.sessionLoads) + assert.Zero(t, store.messageLoads) + assert.Zero(t, store.usageLoads) +} + +func (r *recordingSyncStateValueReader) SyncStateValues( + keys []string, +) (map[string]string, error) { + r.calls++ + r.keys = append([]string(nil), keys...) + result := make(map[string]string) + for _, key := range keys { + if value := r.states[key]; value != "" { + result[key] = value + } + } + return result, nil +} + +func TestImportedSessionIDsReadsOnlyCandidateProvenance(t *testing.T) { + reader := &recordingSyncStateValueReader{states: map[string]string{ + "artifact_import:desk-a1b2c3:desk-a1b2c3~one": "manifest-one", + "artifact_import:laptop-d4e5f6:laptop-d4e5f6~two": "manifest-two", + }} + + got, err := ImportedSessionIDs(reader, []string{ + "desk-a1b2c3~one", + "local-session", + "phone-112233~missing", + }) + require.NoError(t, err) + assert.Equal(t, map[string]struct{}{"desk-a1b2c3~one": {}}, got) + assert.Equal(t, []string{ + "artifact_import:desk-a1b2c3:desk-a1b2c3~one", + "artifact_import:phone-112233:phone-112233~missing", + }, reader.keys) + assert.Equal(t, 1, reader.calls) + + empty, err := ImportedSessionIDs(reader, nil) + require.NoError(t, err) + assert.Empty(t, empty) + assert.Equal(t, 1, reader.calls, + "a no-candidate push must not query artifact provenance") +} + +func TestLoadImportStatesReadsCheckpointInOneBatch(t *testing.T) { + reader := &recordingSyncStateValueReader{states: map[string]string{ + "artifact_import:desk-a1b2c3:desk-a1b2c3~one": "manifest-one", + "artifact_import:desk-a1b2c3:desk-a1b2c3~two": "manifest-two", + }} + + states, err := loadImportStates(reader, "desk-a1b2c3", []string{ + "desk-a1b2c3~one", "desk-a1b2c3~two", "desk-a1b2c3~three", + }) + require.NoError(t, err) + assert.Equal(t, map[string]string{ + "artifact_import:desk-a1b2c3:desk-a1b2c3~one": "manifest-one", + "artifact_import:desk-a1b2c3:desk-a1b2c3~two": "manifest-two", + }, states) + assert.Equal(t, []string{ + "artifact_import:desk-a1b2c3:desk-a1b2c3~one", + "artifact_import:desk-a1b2c3:desk-a1b2c3~two", + "artifact_import:desk-a1b2c3:desk-a1b2c3~three", + }, reader.keys) + assert.Equal(t, 1, reader.calls) +} + +func TestIsFolderTargetAcceptsWindowsDrivePaths(t *testing.T) { + tests := []struct { + name string + target string + want bool + }{ + {"windows backslash path", `C:\Users\runner\artifacts`, true}, + {"windows slash path", `C:/Users/runner/artifacts`, true}, + {"posix path", "/tmp/agentsview-artifacts", true}, + {"relative path", "artifacts", true}, + {"http peer", "https://peer.example.test/artifacts", false}, + {"s3 target", "s3://bucket/artifacts", false}, + {"host port", "localhost:8080", false}, + {"empty", "", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, IsFolderTarget(tt.target)) + }) + } +} + +func TestAdoptOriginPersistsConfigOrigin(t *testing.T) { + database := testDB(t) + + require.NoError(t, AdoptOrigin(database, "desk-a1b2c3")) + + stored, err := StoredOrigin(database) + require.NoError(t, err) + assert.Equal(t, "desk-a1b2c3", stored) + + // EnsureOrigin and its callers now agree with the adopted origin instead + // of generating a divergent DB-only value. + ensured, err := EnsureOrigin(database) + require.NoError(t, err) + assert.Equal(t, "desk-a1b2c3", ensured) +} + +func TestAdoptOriginIsIdempotent(t *testing.T) { + database := testDB(t) + + require.NoError(t, AdoptOrigin(database, "desk-a1b2c3")) + require.NoError(t, AdoptOrigin(database, "desk-a1b2c3")) + + stored, err := StoredOrigin(database) + require.NoError(t, err) + assert.Equal(t, "desk-a1b2c3", stored) +} + +func TestAdoptOriginOverwritesDivergentDBOrigin(t *testing.T) { + database := testDB(t) + + // Simulate the pre-fix state: the recorder generated a DB-only origin + // before the authoritative config origin existed. + stale, err := EnsureOrigin(database) + require.NoError(t, err) + require.NotEqual(t, "desk-a1b2c3", stale) + + require.NoError(t, AdoptOrigin(database, "desk-a1b2c3")) + + stored, err := StoredOrigin(database) + require.NoError(t, err) + assert.Equal(t, "desk-a1b2c3", stored) +} + +func TestAdoptOriginRejectsInvalidOrigin(t *testing.T) { + database := testDB(t) + + err := AdoptOrigin(database, "../outside") + require.Error(t, err) + assert.Contains(t, err.Error(), "adopting artifact origin") + + stored, err := StoredOrigin(database) + require.NoError(t, err) + assert.Empty(t, stored) +} + +func TestEnsureOriginRejectsInvalidPersistedOrigin(t *testing.T) { + database := testDB(t) + require.NoError(t, database.SetSyncState(originStateKey, "../outside")) + + origin, err := EnsureOrigin(database) + require.Error(t, err) + assert.Empty(t, origin) + assert.Contains(t, err.Error(), "stored artifact origin") + assert.Contains(t, err.Error(), "invalid artifact origin") +} + +func TestSyncFolderRoundTripImportsForeignSession(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + aRes, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + assert.Equal(t, "laptop-a1b2c3", aRes.Origin) + assert.Equal(t, 1, aRes.ExportedSessions) + + bRes, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Equal(t, 1, bRes.ImportedSessions) + assert.Equal(t, 2, bRes.ImportedMessages) + + bRes, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Zero(t, bRes.ImportedSessions) + assert.Zero(t, bRes.ImportedMessages) + + got, err := bDB.GetSessionFull(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + require.NotNil(t, got) + assert.Equal(t, "laptop-a1b2c3", got.Machine) + assert.Equal(t, "alpha", got.Project) + + msgs, err := bDB.GetAllMessages(ctx, got.ID) + require.NoError(t, err) + require.Len(t, msgs, 2) + assert.Equal(t, "user", msgs[0].Role) + assert.Equal(t, "hello", msgs[0].Content) + assert.Equal(t, "assistant", msgs[1].Role) + assert.Equal(t, "world", msgs[1].Content) +} + +func TestSyncFolderRoundTripPreservesSessionName(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + sessionName := "Parser Provided Title" + seedSession(t, aDB, "sess-1", "alpha", func(s *db.Session) { + s.SessionName = &sessionName + }) + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + + bRes, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + require.Equal(t, 1, bRes.ImportedSessions) + + got, err := bDB.GetSessionFull(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + require.NotNil(t, got) + require.NotNil(t, got.SessionName) + assert.Equal(t, sessionName, *got.SessionName) +} + +func TestSyncFolderInitBaselineMetadataConvergesCuration(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, AdoptOrigin(aDB, "laptop-a1b2c3")) + require.NoError(t, AdoptOrigin(bDB, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + require.NoError(t, aDB.ReplaceSessionMessages("sess-1", []db.Message{ + { + SessionID: "sess-1", + Ordinal: 0, + Role: "user", + Content: "hello", + ContentLength: 5, + SourceUUID: "uuid-question", + }, + { + SessionID: "sess-1", + Ordinal: 1, + Role: "assistant", + Content: "world", + ContentLength: 5, + SourceUUID: "uuid-answer", + }, + })) + displayName := "Already renamed" + require.NoError(t, aDB.RenameSession("sess-1", &displayName)) + starred, err := aDB.StarSession("sess-1") + require.NoError(t, err) + require.True(t, starred) + msgs, err := aDB.GetAllMessages(ctx, "sess-1") + require.NoError(t, err) + require.Len(t, msgs, 2) + note := "already pinned" + _, err = aDB.PinMessage("sess-1", msgs[1].ID, ¬e) + require.NoError(t, err) + + _, err = SyncFolder(ctx, aDB, SyncOptions{ + DataDir: aData, + Target: share, + BaselineMetadata: true, + }) + require.NoError(t, err) + res, err := SyncFolder(ctx, bDB, SyncOptions{ + DataDir: bData, + Target: share, + }) + require.NoError(t, err) + assert.Equal(t, 1, res.ImportedSessions) + assert.Equal(t, 2, res.ImportedMessages) + assert.Equal(t, 3, res.ImportedMetadata) + + gid := "laptop-a1b2c3~sess-1" + got, err := bDB.GetSessionFull(ctx, gid) + require.NoError(t, err) + require.NotNil(t, got) + require.NotNil(t, got.DisplayName) + assert.Equal(t, displayName, *got.DisplayName) + stars, err := bDB.ListStarredSessionIDs(ctx) + require.NoError(t, err) + assert.Equal(t, []string{gid}, stars) + pins, err := bDB.ListPinnedMessages(ctx, gid, "") + require.NoError(t, err) + require.Len(t, pins, 1) + assert.Equal(t, 1, pins[0].Ordinal) + require.NotNil(t, pins[0].Note) + assert.Equal(t, note, *pins[0].Note) + + op, ok, err := bDB.MetadataReplayStateOp(ctx, gid, "display_name") + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, MetadataOpRename, op) + op, ok, err = bDB.MetadataReplayStateOp(ctx, gid, "starred") + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, MetadataOpStar, op) + op, ok, err = bDB.MetadataReplayStateOp(ctx, gid, "pin:source_uuid:uuid-answer") + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, MetadataOpPin, op) +} + +func TestSyncFolderInitBaselineMetadataConvergesTrash(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, AdoptOrigin(aDB, "laptop-a1b2c3")) + require.NoError(t, AdoptOrigin(bDB, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + _, err := SyncFolder(ctx, aDB, SyncOptions{ + DataDir: aData, + Target: share, + }) + require.NoError(t, err) + _, err = SyncFolder(ctx, bDB, SyncOptions{ + DataDir: bData, + Target: share, + }) + require.NoError(t, err) + + require.NoError(t, aDB.SoftDeleteSession("sess-1")) + + _, err = SyncFolder(ctx, aDB, SyncOptions{ + DataDir: aData, + Target: share, + BaselineMetadata: true, + }) + require.NoError(t, err) + res, err := SyncFolder(ctx, bDB, SyncOptions{ + DataDir: bData, + Target: share, + }) + require.NoError(t, err) + assert.Equal(t, 1, res.ImportedMetadata) + + gid := "laptop-a1b2c3~sess-1" + got, err := bDB.GetSessionFull(ctx, gid) + require.NoError(t, err) + require.NotNil(t, got) + require.NotNil(t, got.DeletedAt) + op, ok, err := bDB.MetadataReplayStateOp(ctx, gid, "deleted_at") + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, MetadataOpSoftDelete, op) +} + +func TestSyncFolderInitBaselinesCurationOfTrashedSession(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, AdoptOrigin(aDB, "laptop-a1b2c3")) + require.NoError(t, AdoptOrigin(bDB, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + require.NoError(t, aDB.ReplaceSessionMessages("sess-1", []db.Message{{ + SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", + ContentLength: 5, SourceUUID: "uuid-question", + }})) + displayName := "Renamed before trash" + require.NoError(t, aDB.RenameSession("sess-1", &displayName)) + starred, err := aDB.StarSession("sess-1") + require.NoError(t, err) + require.True(t, starred) + msgs, err := aDB.GetAllMessages(ctx, "sess-1") + require.NoError(t, err) + require.Len(t, msgs, 1) + note := "pinned before trash" + _, err = aDB.PinMessage("sess-1", msgs[0].ID, ¬e) + require.NoError(t, err) + require.NoError(t, aDB.SoftDeleteSession("sess-1")) + + // The session sits in trash when the machine first opts in: its curation + // must still baseline, or a later restore reaches peers without it. + _, err = SyncFolder(ctx, aDB, SyncOptions{ + DataDir: aData, + Target: share, + BaselineMetadata: true, + }) + require.NoError(t, err) + _, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + + // Restoring on A publishes the session content; B then converges on a + // visible session that kept its pre-init name, star, and pin. + _, err = aDB.RestoreSession("sess-1") + require.NoError(t, err) + repository, err := OpenRepository(ctx, aData) + require.NoError(t, err) + recorder := NewMetadataRecorder(aDB, MetadataRecorderOptions{ + Store: repository.Content(), + Origin: "laptop-a1b2c3", + }) + _, err = recorder.Append(ctx, MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpRestore, + }) + require.NoError(t, err) + require.NoError(t, repository.Close()) + _, err = SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + _, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + + gid := "laptop-a1b2c3~sess-1" + got, err := bDB.GetSessionFull(ctx, gid) + require.NoError(t, err) + require.NotNil(t, got) + assert.Nil(t, got.DeletedAt) + require.NotNil(t, got.DisplayName) + assert.Equal(t, displayName, *got.DisplayName) + stars, err := bDB.ListStarredSessionIDs(ctx) + require.NoError(t, err) + assert.Equal(t, []string{gid}, stars) + pins, err := bDB.ListPinnedMessages(ctx, gid, "") + require.NoError(t, err) + require.Len(t, pins, 1) + require.NotNil(t, pins[0].Note) + assert.Equal(t, note, *pins[0].Note) +} + +func TestSyncFolderImportClearsSourceFileState(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + sourcePath := filepath.Join(t.TempDir(), "shared-session.jsonl") + peerHash := strings64("a") + localHash := strings64("b") + lastEntryUUID := "entry-99" + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha", func(s *db.Session) { + s.FilePath = &sourcePath + s.FileSize = new(int64) + *s.FileSize = 4096 + s.FileMtime = new(int64) + *s.FileMtime = 200 + s.NextOrdinal = 99 + s.LastEntryUUID = &lastEntryUUID + s.FileInode = new(int64) + *s.FileInode = 12345 + s.FileDevice = new(int64) + *s.FileDevice = 67890 + s.FileHash = &peerHash + }) + seedSession(t, bDB, "local-sess", "alpha", func(s *db.Session) { + s.FilePath = &sourcePath + s.FileMtime = new(int64) + *s.FileMtime = 100 + s.FileHash = &localHash + }) + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + res, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + require.Equal(t, 1, res.ImportedSessions) + + imported, err := bDB.GetSessionFull(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + require.NotNil(t, imported) + assert.Nil(t, imported.FilePath) + assert.Nil(t, imported.FileSize) + assert.Nil(t, imported.FileMtime) + assert.Zero(t, imported.NextOrdinal) + assert.Nil(t, imported.LastEntryUUID) + assert.Nil(t, imported.FileInode) + assert.Nil(t, imported.FileDevice) + assert.Nil(t, imported.FileHash) + + ids, err := bDB.ListSessionIDsByFilePath(sourcePath, "claude") + require.NoError(t, err) + assert.Equal(t, []string{"local-sess"}, ids) + gotHash, ok := bDB.GetFileHashByPath(sourcePath) + require.True(t, ok) + assert.Equal(t, localHash, gotHash) +} + +func TestSyncFolderRoundTripPreservesSessionSignals(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + // Signal columns are written outside UpsertSession, so seed them through + // the same writer paths the live app uses. These include fields the Session + // JSON drops (json:"-"): has_tool_calls, has_context_data, and the quality + // scalars. Secret-scan state is seeded too, but unlike the others it is + // deliberately not carried across import (asserted below). + require.NoError(t, aDB.UpdateSessionSignals("sess-1", db.SessionSignalUpdate{ + HasToolCalls: true, + HasContextData: true, + Outcome: "success", + QualitySignals: db.QualitySignals{ + Version: 3, + ShortPromptCount: 2, + UnstructuredStart: true, + RunawayToolLoopCount: 1, + }, + })) + require.NoError(t, aDB.ReplaceSessionSecretFindings("sess-1", nil, 0, "rules-v7")) + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + + bRes, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + require.Equal(t, 1, bRes.ImportedSessions) + + got, err := bDB.GetSessionFull(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + require.NotNil(t, got) + assert.True(t, got.HasToolCalls, "has_tool_calls should survive the round trip") + assert.True(t, got.HasContextData, "has_context_data should survive the round trip") + assert.Equal(t, "success", got.Outcome) + // Secret findings are not carried in the manifest, so the imported session + // is treated as unscanned: the source rules version is dropped so + // `secrets scan --backfill` rescans it with local rules. + assert.Empty(t, got.SecretsRulesVersion, + "secret-scan state must not be restored without findings") + + qs := got.StoredQualitySignals() + require.NotNil(t, qs, "quality signals should survive the round trip") + assert.Equal(t, 3, qs.Version) + assert.Equal(t, 2, qs.ShortPromptCount) + assert.True(t, qs.UnstructuredStart) + assert.Equal(t, 1, qs.RunawayToolLoopCount) +} + +func TestSyncFolderRoundTripRewritesForeignRelationshipIDs(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "source-1", "alpha") + seedSession(t, aDB, "parent-1", "alpha") + seedSession(t, aDB, "child-1", "alpha") + parentID := "parent-1" + seedSession(t, aDB, "sess-1", "alpha", func(s *db.Session) { + s.SourceSessionID = "source-1" + s.ParentSessionID = &parentID + }) + require.NoError(t, aDB.ReplaceSessionMessages("sess-1", []db.Message{ + { + SessionID: "sess-1", + Ordinal: 0, + Role: "assistant", + Content: "delegating", + ContentLength: 10, + ToolCalls: []db.ToolCall{{ + ToolName: "Task", + Category: "Task", + ToolUseID: "toolu_1", + SubagentSessionID: "child-1", + ResultEvents: []db.ToolResultEvent{{ + ToolUseID: "toolu_1", + AgentID: "agent-1", + SubagentSessionID: "child-1", + Source: "tool_result", + Status: "success", + Content: "done", + ContentLength: 4, + EventIndex: 0, + }}, + }}, + }, + })) + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + _, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + + got, err := bDB.GetSessionFull(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + require.NotNil(t, got) + assert.Equal(t, "laptop-a1b2c3~source-1", got.SourceSessionID) + require.NotNil(t, got.ParentSessionID) + assert.Equal(t, "laptop-a1b2c3~parent-1", *got.ParentSessionID) + + msgs, err := bDB.GetAllMessages(ctx, got.ID) + require.NoError(t, err) + require.Len(t, msgs, 1) + require.Len(t, msgs[0].ToolCalls, 1) + assert.Equal(t, "laptop-a1b2c3~child-1", msgs[0].ToolCalls[0].SubagentSessionID) + require.Len(t, msgs[0].ToolCalls[0].ResultEvents, 1) + assert.Equal(t, "laptop-a1b2c3~child-1", + msgs[0].ToolCalls[0].ResultEvents[0].SubagentSessionID) +} + +// Recent Edits and edit file grouping read the persisted +// tool_calls.file_path column, so the segment format must carry it: +// imported foreign sessions get no parse-time re-derivation and the +// one-time file_path backfill has already run on existing databases. +func TestSyncFolderRoundTripPreservesToolCallFilePath(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + require.NoError(t, aDB.ReplaceSessionMessages("sess-1", []db.Message{ + { + SessionID: "sess-1", + Ordinal: 0, + Role: "assistant", + Content: "editing", + ContentLength: 7, + ToolCalls: []db.ToolCall{{ + ToolName: "Edit", + Category: "Edit", + ToolUseID: "toolu_1", + InputJSON: `{"file_path":"src/app.go"}`, + FilePath: "src/app.go", + }}, + }, + })) + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + _, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + + msgs, err := bDB.GetAllMessages(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + require.Len(t, msgs, 1) + require.Len(t, msgs[0].ToolCalls, 1) + assert.Equal(t, "src/app.go", msgs[0].ToolCalls[0].FilePath) +} + +// TestSyncFolderImportLeavesScannedSessionBackfillable verifies that a session +// scanned for secrets at the source (rules version, a finding row, a nonzero +// leak count) is imported as unscanned, because the manifest carries no finding +// rows. The imported session must have no leak count and no rules version, so it +// stays a `secrets scan --backfill` candidate even when the source rules version +// is current on the importing machine. Stamping it scanned-at-source-version +// would skip a secret-bearing session, leaving a leak count with no revealable +// findings. +func TestSyncFolderImportLeavesScannedSessionBackfillable(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + // The source session is fully scanned: a finding row, a nonzero leak count, + // and a rules version that is also current on the importing machine below. + const rulesVersion = "rules-current" + findings := []db.SecretFinding{{ + SessionID: "sess-1", + RuleName: "aws-access-key", + Confidence: "definite", + LocationKind: "message", + MessageOrdinal: 1, + MatchStart: 4, + MatchEnd: 24, + RedactedMatch: "AKIA…MPLE", + RulesVersion: rulesVersion, + }} + require.NoError(t, aDB.ReplaceSessionSecretFindings("sess-1", findings, 1, rulesVersion)) + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + + bRes, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + require.Equal(t, 1, bRes.ImportedSessions) + + const importedID = "laptop-a1b2c3~sess-1" + got, err := bDB.GetSessionFull(ctx, importedID) + require.NoError(t, err) + require.NotNil(t, got) + // Imported session must be unscanned so its state is consistent with the zero + // findings carried in the manifest. + assert.Empty(t, got.SecretsRulesVersion, "imported session must not be stamped scanned") + assert.Zero(t, got.SecretLeakCount, "imported session must not claim leaks without findings") + + // With the source rules version current on the importing machine, backfill + // must still treat the imported session as a candidate (secrets_rules_version + // "" != current) instead of skipping it. + cands, err := bDB.SecretScanCandidates(ctx, db.SecretScanCandidateFilter{ + CurrentVersion: rulesVersion, + OnlyStale: true, + }) + require.NoError(t, err) + assert.Contains(t, cands, importedID, + "secret-bearing imported session must be a backfill candidate, not skipped") +} + +// TestSyncFolderSourceLeakCountChangeKeepsLocalFindings verifies that a +// source-side secret rescan that changes only secret_leak_count (not message +// content) does not alter the artifact manifest hash, so the importer neither +// re-imports the session nor clears the findings it scanned locally. +// secret_leak_count is the only secret field carried in the Session JSON, and +// import discards secret-scan state, so it must not influence the +// content-addressed manifest. +func TestSyncFolderSourceLeakCountChangeKeepsLocalFindings(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + // First round trip: A exports sess-1 (no secrets yet), B imports it. + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + bRes, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + require.Equal(t, 1, bRes.ImportedSessions) + + const importedID = "laptop-a1b2c3~sess-1" + + // B scans the imported session locally and records a finding. + bFinding := []db.SecretFinding{{ + SessionID: importedID, RuleName: "aws-access-key", Confidence: "definite", + LocationKind: "message", MessageOrdinal: 0, MatchStart: 4, MatchEnd: 24, + RedactedMatch: "AKIA…MPLE", RulesVersion: "rules-b", + }} + require.NoError(t, bDB.ReplaceSessionSecretFindings(importedID, bFinding, 1, "rules-b")) + + // A rescans sess-1: only secret_leak_count changes (0 -> 1); the message + // content A exports is untouched. + aFinding := []db.SecretFinding{{ + SessionID: "sess-1", RuleName: "aws-access-key", Confidence: "definite", + LocationKind: "message", MessageOrdinal: 0, MatchStart: 4, MatchEnd: 24, + RedactedMatch: "AKIA…MPLE", RulesVersion: "rules-a", + }} + require.NoError(t, aDB.ReplaceSessionSecretFindings("sess-1", aFinding, 1, "rules-a")) + + // Second round trip after the source-only rescan. + _, err = SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + bRes, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Zero(t, bRes.ImportedSessions, + "a source-only leak-count change must not re-import the session") + + // B's locally scanned findings and scan state survive. + got, err := bDB.SessionSecretFindings(ctx, importedID) + require.NoError(t, err) + assert.Len(t, got, 1, "importer's local secret findings must not be cleared") + + sess, err := bDB.GetSessionFull(ctx, importedID) + require.NoError(t, err) + require.NotNil(t, sess) + assert.Equal(t, 1, sess.SecretLeakCount, "importer's leak count preserved") + assert.Equal(t, "rules-b", sess.SecretsRulesVersion, "importer's scan version preserved") +} + +func TestSyncFolderNotifiesWhenImportWritesData(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + + changes := 0 + _, err = SyncFolder(ctx, bDB, SyncOptions{ + DataDir: bData, + Target: share, + OnDataChanged: func() { changes++ }, + }) + require.NoError(t, err) + assert.Equal(t, 1, changes) + + _, err = SyncFolder(ctx, bDB, SyncOptions{ + DataDir: bData, + Target: share, + OnDataChanged: func() { changes++ }, + }) + require.NoError(t, err) + assert.Equal(t, 1, changes) +} + +func TestSyncFolderUsesProvidedOrigin(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + dataDir := t.TempDir() + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + + res, err := SyncFolder(ctx, database, SyncOptions{ + DataDir: dataDir, + Target: share, + Origin: "configured-a1b2c3", + }) + require.NoError(t, err) + + assert.Equal(t, "configured-a1b2c3", res.Origin) + persisted, err := database.GetSyncState(originStateKey) + require.NoError(t, err) + assert.Empty(t, persisted) + manifests := globArtifacts(t, share, "configured-a1b2c3", "manifests", "*"+manifestExtension) + assert.Len(t, manifests, 1) +} + +func TestImportMaintainsFTS(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + exportDB := testDB(t) + importDB := testDB(t) + if !importDB.HasFTS() { + t.Skip("FTS unavailable") + } + seedSession(t, exportDB, "sess-1", "alpha") + + _, err := ExportToStore(ctx, exportDB, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + imported, messages, err := importFromTestStore(ctx, importDB, store, "desktop-d4e5f6") + require.NoError(t, err) + require.Equal(t, 1, imported) + require.Equal(t, 2, messages) + + page, err := importDB.Search(ctx, db.SearchFilter{Query: "world", Limit: 10}) + require.NoError(t, err) + require.Len(t, page.Results, 1) + assert.Equal(t, origin+"~sess-1", page.Results[0].SessionID) +} + +func TestImportPreservesPinsAndStatsOnRewrite(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + exportDB := testDB(t) + importDB := testDB(t) + seedSession(t, exportDB, "sess-1", "alpha") + + _, err := ExportToStore(ctx, exportDB, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + imported, messages, err := importFromTestStore(ctx, importDB, store, "desktop-d4e5f6") + require.NoError(t, err) + require.Equal(t, 1, imported) + require.Equal(t, 2, messages) + + gid := origin + "~sess-1" + importedMsgs, err := importDB.GetAllMessages(ctx, gid) + require.NoError(t, err) + require.Len(t, importedMsgs, 2) + note := "keep this pin" + _, err = importDB.PinMessage(gid, importedMsgs[1].ID, ¬e) + require.NoError(t, err) + + require.NoError(t, exportDB.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + {SessionID: "sess-1", Ordinal: 1, Role: "assistant", Content: "planet", ContentLength: 6}, + })) + _, err = ExportToStore(ctx, exportDB, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + imported, messages, err = importFromTestStore(ctx, importDB, store, "desktop-d4e5f6") + require.NoError(t, err) + require.Equal(t, 1, imported) + require.Equal(t, 2, messages) + + pins, err := importDB.ListPinnedMessages(ctx, gid, "") + require.NoError(t, err) + require.Len(t, pins, 1) + assert.Equal(t, 1, pins[0].Ordinal) + require.NotNil(t, pins[0].Note) + assert.Equal(t, note, *pins[0].Note) + + allPins, err := importDB.ListPinnedMessages(ctx, "", "") + require.NoError(t, err) + require.Len(t, allPins, 1) + assert.Equal(t, gid, allPins[0].SessionID) + require.NotNil(t, allPins[0].Content) + assert.Equal(t, "planet", *allPins[0].Content) + + stats, err := importDB.GetStats(ctx, false, false) + require.NoError(t, err) + assert.Equal(t, 1, stats.SessionCount) + assert.Equal(t, 2, stats.MessageCount) + assert.Equal(t, 1, stats.ProjectCount) + assert.Equal(t, 1, stats.MachineCount) +} + +func TestImportDoesNotAdvanceStateForExcludedOrTrashedSessions(t *testing.T) { + ctx := context.Background() + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + exportDB := testDB(t) + seedSession(t, exportDB, "sess-1", "alpha") + _, err := ExportToStore(ctx, exportDB, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + + gid := origin + "~sess-1" + tests := []struct { + name string + seed func(*testing.T, *db.DB) + }{ + { + name: "excluded", + seed: func(t *testing.T, database *db.DB) { + t.Helper() + seedSession(t, database, gid, "alpha", func(s *db.Session) { + s.Machine = origin + }) + require.NoError(t, database.DeleteSession(gid)) + }, + }, + { + name: "trashed", + seed: func(t *testing.T, database *db.DB) { + t.Helper() + seedSession(t, database, gid, "alpha", func(s *db.Session) { + s.Machine = origin + }) + require.NoError(t, database.SoftDeleteSession(gid)) + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + importDB := testDB(t) + tt.seed(t, importDB) + + imported, messages, err := importFromTestStore(ctx, importDB, store, "desktop-d4e5f6") + require.NoError(t, err) + assert.Zero(t, imported) + assert.Zero(t, messages) + + state, err := importDB.GetSyncState(importStateKey(origin, gid)) + require.NoError(t, err) + assert.Empty(t, state) + }) + } +} + +func TestSyncFolderRetriesIncompleteForeignArtifacts(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + segments := globArtifacts(t, share, "laptop-a1b2c3", "segments", "*"+segmentExtension) + require.Len(t, segments, 1) + require.NoError(t, os.Remove(segments[0])) + + res, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Zero(t, res.ImportedSessions) + assert.Zero(t, res.ImportedMessages) + + got, err := bDB.GetSessionFull(ctx, "laptop-a1b2c3~sess-1") + require.NoError(t, err) + assert.Nil(t, got) + + _, err = SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + res, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Equal(t, 1, res.ImportedSessions) + assert.Equal(t, 2, res.ImportedMessages) +} + +func TestSyncFolderSkipsCheckpointWithMissingManifest(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + manifests := globArtifacts(t, share, "laptop-a1b2c3", "manifests", "*"+manifestExtension) + require.Len(t, manifests, 1) + require.NoError(t, os.Remove(manifests[0])) + + res, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Zero(t, res.ImportedSessions) + assert.Zero(t, res.ImportedMessages) +} + +func TestSyncFolderRejectsOverlappingRoots(t *testing.T) { + ctx := context.Background() + + tests := []struct { + name string + target func(string) string + }{ + { + name: "target is data dir", + target: func(dataDir string) string { + return dataDir + }, + }, + { + name: "target is artifact store", + target: func(dataDir string) string { + return filepath.Join(dataDir, "artifacts") + }, + }, + { + name: "target inside artifact store", + target: func(dataDir string) string { + return filepath.Join(dataDir, "artifacts", "share") + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + database := testDB(t) + dataDir := t.TempDir() + + _, err := SyncFolder(ctx, database, SyncOptions{ + DataDir: dataDir, + Target: tt.target(dataDir), + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "must not overlap") + }) + } +} + +func TestSyncFolderRejectsSymlinkedAncestorResolvingInsideArtifactStore(t *testing.T) { + dataDir := t.TempDir() + localRoot := filepath.Join(dataDir, "artifacts") + targetDir := filepath.Join(localRoot, "shared") + require.NoError(t, os.MkdirAll(targetDir, 0o755)) + alias := filepath.Join(t.TempDir(), "alias") + require.NoError(t, os.Symlink(localRoot, alias)) + + err := validateDisjointRoots(localRoot, filepath.Join(alias, "shared")) + require.Error(t, err) + assert.Contains(t, err.Error(), "must not overlap") +} + +type batchNotificationTransport struct{ exchangeErr error } + +func (t batchNotificationTransport) Prepare(context.Context, ArtifactStore) error { return nil } + +func (t batchNotificationTransport) Exchange(context.Context, ArtifactStore) error { + return t.exchangeErr +} + +func TestSyncBatchNotificationRunsOnlyAfterSuccess(t *testing.T) { + for _, suppressed := range []bool{false, true} { + mode := "success" + want := int32(1) + if suppressed { + mode = "shutdown" + want = 0 + } + t.Run(mode, func(t *testing.T) { + repository, err := OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + var notifications atomic.Int32 + notify := func(ctx context.Context) { + if !artifactMaintenanceSuppressed(ctx) { + notifications.Add(1) + } + } + database := testDB(t) + seedSession(t, database, "notify-session", "project") + ctx := t.Context() + if suppressed { + ctx = SuppressArtifactMaintenance(ctx) + } + + _, err = syncContentWithTransport( + ctx, database, nil, repository.Content(), notify, + SyncOptions{Origin: "local-a1b2c3"}, batchNotificationTransport{}, + ) + require.NoError(t, err) + assert.Equal(t, want, notifications.Load()) + }) + } +} + +func TestSyncBatchNotificationSkipsFailedExchange(t *testing.T) { + repository, err := OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + database := testDB(t) + seedSession(t, database, "notify-session", "project") + var notifications atomic.Int32 + + _, err = syncContentWithTransport( + t.Context(), database, nil, repository.Content(), + func(context.Context) { notifications.Add(1) }, + SyncOptions{Origin: "local-a1b2c3"}, + batchNotificationTransport{exchangeErr: errors.New("exchange failed")}, + ) + require.Error(t, err) + assert.Zero(t, notifications.Load()) +} + +func TestSyncUnchangedArchiveDoesNotRepublish(t *testing.T) { + for _, archiveSize := range []int{0, 500} { + t.Run(fmt.Sprintf("archive-%d", archiveSize), func(t *testing.T) { + database := testDB(t) + for i := range archiveSize { + require.NoError(t, database.UpsertSession(db.Session{ + ID: fmt.Sprintf("peer-%04d", i), Project: "project", + Machine: "peer-a1b2c3", Agent: "claude", + })) + } + seedSession(t, database, "session-0000", "project") + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + store := &recordingArtifactStore{ArtifactStore: base} + opts := SyncOptions{Origin: "local-a1b2c3"} + + initial, err := syncContentWithTransport( + t.Context(), database, nil, store, nil, opts, batchNotificationTransport{}, + ) + require.NoError(t, err) + assert.Equal(t, 1, initial.ExportedSessions) + + store.creates = nil + unchanged, err := syncContentWithTransport( + t.Context(), database, nil, store, nil, opts, batchNotificationTransport{}, + ) + require.NoError(t, err) + assert.Zero(t, unchanged.ExportedSessions) + assert.Empty(t, store.creates, + "unchanged sync work must not grow with the total archive") + + require.NoError(t, database.ReplaceSessionMessages("session-0000", []db.Message{{ + SessionID: "session-0000", Ordinal: 0, Role: "user", Content: "changed", + }})) + store.creates = nil + changed, err := syncContentWithTransport( + t.Context(), database, nil, store, nil, opts, batchNotificationTransport{}, + ) + require.NoError(t, err) + assert.Equal(t, 1, changed.ExportedSessions) + createdKinds := make(map[Kind]int) + for _, create := range store.creates { + createdKinds[create.Ref.Kind]++ + } + assert.Equal(t, map[Kind]int{ + KindSegments: 1, KindManifests: 1, KindCheckpoints: 1, + }, createdKinds) + }) + } +} + +func TestSyncWithStoreRejectsFolderWithoutRepositoryOwner(t *testing.T) { + dataDir := t.TempDir() + repository, err := OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + _, err = SyncWithStore(t.Context(), testDB(t), repository.Content(), SyncOptions{ + DataDir: t.TempDir(), Target: t.TempDir(), Origin: "local-a1b2c3", + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactUnsupported) +} + +func TestValidateSyncTargetRejectsSecretURLComponentsWithoutDisclosure(t *testing.T) { + const secret = "target-secret" + for _, target := range []string{ + "https://user:" + secret + "@example.invalid/archive", + "https://example.invalid/archive?token=" + secret, + "https://example.invalid/archive#" + secret, + "s3://user:" + secret + "@bucket/archive", + "s3://bucket/archive?token=" + secret, + "s3://bucket/archive#" + secret, + } { + err := ValidateSyncTarget(target) + require.ErrorIs(t, err, ErrArtifactInvalid) + assert.NotContains(t, err.Error(), secret) + } + require.NoError(t, ValidateSyncTarget(filepath.Join(t.TempDir(), "folder?with#marks")), + "non-URL folder targets retain platform-native path semantics") +} + +func TestSyncWithRepositoryFolderUsesRetainedRepositoryIdentity(t *testing.T) { + vaultDataDir := t.TempDir() + repository, err := OpenRepository(t.Context(), vaultDataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + legacyDataDir := t.TempDir() + + t.Run("actual vault and hierarchy are rejected", func(t *testing.T) { + targets := map[string]string{ + "vault": filepath.Join(vaultDataDir, repositoryDirectory), + "ancestor": vaultDataDir, + "descendant": filepath.Join(vaultDataDir, repositoryDirectory, "share"), + } + alias := filepath.Join(t.TempDir(), "vault-alias") + if err := os.Symlink(filepath.Join(vaultDataDir, repositoryDirectory), alias); err == nil { + targets["symlink alias"] = alias + } + for name, target := range targets { + t.Run(name, func(t *testing.T) { + _, err := SyncWithRepository(t.Context(), testDB(t), repository, SyncOptions{ + DataDir: legacyDataDir, Target: target, Origin: "local-a1b2c3", + }) + require.Error(t, err) + assert.ErrorContains(t, err, "must not overlap") + }) + } + }) + + t.Run("unrelated legacy path is allowed", func(t *testing.T) { + target := filepath.Join(legacyDataDir, repositoryDirectory) + _, err := SyncWithRepository(t.Context(), testDB(t), repository, SyncOptions{ + DataDir: legacyDataDir, Target: target, Origin: "local-a1b2c3", + }) + require.NoError(t, err) + }) +} + +func TestExportEmitsNewManifestAfterDataVersionChange(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + + result, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + require.Equal(t, 1, result.ExportedSessions) + _, cp, err := latestStoreCheckpointSummary(ctx, store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + firstHash := cp.Sessions[origin+"~sess-1"] + require.NotEmpty(t, firstHash) + + require.NoError(t, database.SetSessionDataVersion("sess-1", 42)) + result, err = ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + assert.Equal(t, 1, result.ExportedSessions) + _, cp, err = latestStoreCheckpointSummary(ctx, store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + nextHash := cp.Sessions[origin+"~sess-1"] + require.NotEmpty(t, nextHash) + assert.NotEqual(t, firstHash, nextHash) + + ref, err := NewRef(origin, KindManifests, nextHash+".json") + require.NoError(t, err) + m, err := decodeManifestWithLimits(readContractArtifact(t, store, ref), productionArtifactLimits()) + require.NoError(t, err) + assert.Equal(t, 42, m.DataVersion) +} + +func TestExportIncludesLocalOwnedSessionClasses(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "file-sess", "alpha") + seedSession(t, database, "claude-ai-sess", "bravo", func(s *db.Session) { + s.Agent = "claude-ai" + }) + seedSession(t, database, "upload-sess", "charlie", func(s *db.Session) { + s.Agent = "upload" + }) + seedSession(t, database, "orphan-sess", "delta", func(s *db.Session) { + s.SourceSessionID = "missing-source" + }) + + result, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + assert.Equal(t, 4, result.ExportedSessions) + + _, cp, err := latestStoreCheckpointSummary(ctx, store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + assert.Contains(t, cp.Sessions, origin+"~file-sess") + assert.Contains(t, cp.Sessions, origin+"~claude-ai-sess") + assert.Contains(t, cp.Sessions, origin+"~upload-sess") + assert.Contains(t, cp.Sessions, origin+"~orphan-sess") +} + +func TestExportScrubsUnstableArtifactIDs(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + origin := "laptop-a1b2c3" + seedSession(t, database, "sess-1", "alpha") + require.NoError(t, database.ReplaceSessionUsageEvents("sess-1", []db.UsageEvent{ + { + SessionID: "sess-1", + Source: "fixture", + Model: "claude-test", + InputTokens: 1, + OccurredAt: "2026-06-14T01:02:04Z", + DedupKey: "usage-1", + }, + })) + + _, err := ExportToStore(ctx, database, store, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + _, cp, err := latestStoreCheckpointSummary(ctx, store, origin) + require.NoError(t, err) + require.NotNil(t, cp) + manifestHash := cp.Sessions[origin+"~sess-1"] + require.NotEmpty(t, manifestHash) + + ref, err := NewRef(origin, KindManifests, manifestHash+".json") + require.NoError(t, err) + m, err := decodeManifestWithLimits(readContractArtifact(t, store, ref), productionArtifactLimits()) + require.NoError(t, err) + require.Len(t, m.UsageEvents, 1) + manifestData, err := canonicalJSON(m) + require.NoError(t, err) + assert.NotContains(t, string(manifestData), `"ID"`) + assert.NotContains(t, string(manifestData), `"SessionID"`) + + msgs := testStoreManifestMessages(t, store, origin, m) + require.Len(t, msgs, 2) + for _, msg := range msgs { + assert.Zero(t, msg.ID) + assert.Empty(t, msg.SessionID) + } +} + +func TestExportRejectsInvalidOriginBeforeCreatingPaths(t *testing.T) { + ctx := context.Background() + database := testDB(t) + store := newTestArtifactStore(t) + seedSession(t, database, "sess-1", "alpha") + + result, err := ExportToStore(ctx, database, store, ExportOptions{Origin: "../outside", Full: true}) + require.Error(t, err) + assert.Zero(t, result.ExportedSessions) + assert.Contains(t, err.Error(), "invalid artifact origin") + assertNoPublishedArtifacts(t, store, "outside-a1b2c3") +} + +func TestSyncFolderHealsCorruptSegmentInShare(t *testing.T) { + ctx := context.Background() + share := t.TempDir() + aData := t.TempDir() + bData := t.TempDir() + aDB := testDB(t) + bDB := testDB(t) + + require.NoError(t, aDB.SetSyncState(originStateKey, "laptop-a1b2c3")) + require.NoError(t, bDB.SetSyncState(originStateKey, "desktop-d4e5f6")) + seedSession(t, aDB, "sess-1", "alpha") + + _, err := SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + + // Corrupt A's segment in the share, as an interrupted file-sync write would. + segments := globArtifacts(t, share, "laptop-a1b2c3", "segments", "*"+segmentExtension) + require.Len(t, segments, 1) + require.NoError(t, os.Remove(segments[0])) + require.NoError(t, os.WriteFile(segments[0], compressPeerTestData(t, []byte("tampered")), 0o644)) + + // B tolerates the corrupt share: sync succeeds, the poison is not + // mirrored into B's local store, and repeat runs stay healthy. + res, err := SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Zero(t, res.ImportedSessions) + segmentRef, err := FromWireRef("laptop-a1b2c3", KindSegments, filepath.Base(segments[0])) + require.NoError(t, err) + bRepository, err := OpenRepository(ctx, bData) + require.NoError(t, err) + _, err = bRepository.Content().Stat(ctx, segmentRef) + assert.ErrorIs(t, err, ErrArtifactNotFound) + require.NoError(t, bRepository.Close()) + _, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + + // A still holds the valid copy and repairs the share on its next sync. + aRepository, err := OpenRepository(ctx, aData) + require.NoError(t, err) + _, reader, err := aRepository.Content().Open(ctx, segmentRef) + require.NoError(t, err) + var want bytes.Buffer + require.NoError(t, EncodeWire(ctx, segmentRef, reader, &want)) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + require.NoError(t, aRepository.Close()) + _, err = SyncFolder(ctx, aDB, SyncOptions{DataDir: aData, Target: share}) + require.NoError(t, err) + got, err := os.ReadFile(segments[0]) + require.NoError(t, err) + require.Equal(t, want.Bytes(), got, "publisher should repair the corrupt share copy") + + // With the share healed, B converges. + res, err = SyncFolder(ctx, bDB, SyncOptions{DataDir: bData, Target: share}) + require.NoError(t, err) + assert.Equal(t, 1, res.ImportedSessions) + assert.Equal(t, 2, res.ImportedMessages) +} + +func globArtifacts(t *testing.T, root, origin, kind, pattern string) []string { + t.Helper() + paths, err := filepath.Glob(filepath.Join(root, origin, kind, pattern)) + require.NoError(t, err) + return paths +} + +func testDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(filepath.Join(t.TempDir(), "test.db")) + require.NoError(t, err) + t.Cleanup(func() { database.Close() }) + return database +} + +func seedSession(t *testing.T, database *db.DB, id, project string, opts ...func(*db.Session)) { + t.Helper() + sess := db.Session{ + ID: id, + Project: project, + Machine: "local", + Agent: "claude", + MessageCount: 2, + UserMessageCount: 1, + FirstMessage: new("hello"), + StartedAt: new("2026-06-14T01:02:03Z"), + EndedAt: new("2026-06-14T01:03:03Z"), + SessionName: new("Test Session"), + CreatedAt: "2026-06-14T01:02:03Z", + } + for _, opt := range opts { + opt(&sess) + } + require.NoError(t, database.UpsertSession(sess)) + require.NoError(t, database.ReplaceSessionMessages(id, []db.Message{ + {SessionID: id, Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + {SessionID: id, Ordinal: 1, Role: "assistant", Content: "world", ContentLength: 5}, + })) +} diff --git a/internal/artifact/transport.go b/internal/artifact/transport.go new file mode 100644 index 000000000..5d917a745 --- /dev/null +++ b/internal/artifact/transport.go @@ -0,0 +1,1040 @@ +package artifact + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "io/fs" + "log" + "os" + "path/filepath" + "strings" + "sync" +) + +func openArtifactRoot(path, role string) (*os.Root, error) { + info, err := os.Lstat(path) + if err != nil { + return nil, err + } + if !info.IsDir() { + return nil, fmt.Errorf("artifact %s root %s is not a directory", role, path) + } + root, err := os.OpenRoot(path) + if err != nil { + return nil, fmt.Errorf("opening artifact %s root: %w", role, err) + } + openedInfo, err := root.Stat(".") + if err != nil { + _ = root.Close() + return nil, fmt.Errorf("stating artifact %s root: %w", role, err) + } + currentInfo, err := os.Lstat(path) + if err != nil { + _ = root.Close() + return nil, err + } + if !currentInfo.IsDir() || !os.SameFile(openedInfo, currentInfo) { + _ = root.Close() + return nil, fmt.Errorf("artifact %s root %s changed while opening", role, path) + } + return root, nil +} + +func quarantineArtifactRoot(root *os.Root, rel string) { + dst := rel + quarantineSuffix + _ = root.Remove(dst) + err := root.Rename(rel, dst) + path := filepath.Join(root.Name(), rel) + switch { + case err == nil: + log.Printf("artifact: quarantined corrupt artifact %s", path) + case !errors.Is(err, fs.ErrNotExist): + log.Printf("artifact: quarantining %s: %v", path, err) + } +} + +func openArtifactSubroot(parent *os.Root, rel, role string) (*os.Root, error) { + info, err := parent.Lstat(rel) + if err != nil { + return nil, err + } + if !info.IsDir() { + return nil, fmt.Errorf("artifact %s %s is not a directory", role, rel) + } + root, err := parent.OpenRoot(rel) + if err != nil { + return nil, fmt.Errorf("opening artifact %s: %w", role, err) + } + openedInfo, err := root.Stat(".") + if err != nil { + _ = root.Close() + return nil, fmt.Errorf("stating artifact %s: %w", role, err) + } + currentInfo, err := parent.Lstat(rel) + if err != nil { + _ = root.Close() + return nil, err + } + if !currentInfo.IsDir() || !os.SameFile(openedInfo, currentInfo) { + _ = root.Close() + return nil, fmt.Errorf("artifact %s %s changed while opening", role, rel) + } + return root, nil +} + +func openRootRegularFile(root *os.Root, name string) (*os.File, fs.FileInfo, error) { + before, err := root.Lstat(name) + if err != nil { + return nil, nil, err + } + if !before.Mode().IsRegular() { + return nil, nil, fmt.Errorf("artifact source %s is not a regular file", + filepath.Join(root.Name(), name)) + } + file, err := root.Open(name) + if err != nil { + return nil, nil, err + } + opened, err := file.Stat() + if err != nil { + _ = file.Close() + return nil, nil, err + } + after, err := root.Lstat(name) + if err != nil { + _ = file.Close() + return nil, nil, err + } + if !after.Mode().IsRegular() || !os.SameFile(before, opened) || !os.SameFile(opened, after) { + _ = file.Close() + return nil, nil, fmt.Errorf("artifact source %s changed while opening", + filepath.Join(root.Name(), name)) + } + return file, opened, nil +} + +func createRootTemp(root *os.Root, dir string) (*os.File, string, error) { + for range 100 { + var suffix [8]byte + if _, err := rand.Read(suffix[:]); err != nil { + return nil, "", err + } + rel := filepath.Join(dir, tempFilePrefix+hex.EncodeToString(suffix[:])) + file, err := root.OpenFile(rel, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0o600) + if err == nil { + return file, rel, nil + } + if !errors.Is(err, fs.ErrExist) { + return nil, "", err + } + } + return nil, "", errors.New("creating temporary artifact file: too many collisions") +} + +const ( + transportPageSize = 512 + transportPageMaxBytes = 1 << 20 + quarantineSuffix = ".corrupt" +) + +var transportKinds = [...]Kind{ + KindSegments, + KindRaw, + KindManifests, + KindMeta, + KindCheckpoints, +} + +var errArtifactPathConflict = errors.New("artifact path conflict") + +func readTransportPage(body io.Reader) ([]byte, error) { + page, err := io.ReadAll(io.LimitReader(body, transportPageMaxBytes+1)) + if err != nil { + return nil, err + } + if len(page) > transportPageMaxBytes { + return nil, fmt.Errorf("%w: transport page response exceeds %d bytes", + ErrArtifactInvalid, transportPageMaxBytes) + } + return page, nil +} + +// Transport exchanges immutable artifacts through the canonical logical store +// boundary. A transport never receives or derives the store's physical root. +type Transport interface { + Prepare(context.Context, ArtifactStore) error + Exchange(context.Context, ArtifactStore) error +} + +// transportRepairStore is implemented by a store decorator that binds the +// durable SQLite repair queue and StoreImportCoordinator to a canonical store. +// The point lookup keeps equal-name repair checks bounded; completion owns the +// exact-identity repair, acknowledgement, and coalesced import scheduling. +type transportRepairStore interface { + PendingTransportRepair(context.Context, Ref) (Entry, bool, error) + RepairTransportArtifact(context.Context, Entry, io.Reader) error + AcknowledgeTransportRepair(context.Context, Entry) error +} + +type transportChangeStore interface { + RecordTransportChanged(context.Context, Entry) error +} + +// folderTransport owns only its external wire directory. It deliberately does +// not wrap that directory in filesystemStore: the share has its own layout, +// encoding, confinement, and no-follow rules. +type folderTransport struct { + mu sync.Mutex + target string + root *os.Root + closed bool +} + +func (t *folderTransport) Prepare(ctx context.Context, _ ArtifactStore) error { + if err := requireTransportContext(ctx); err != nil { + return err + } + t.mu.Lock() + defer t.mu.Unlock() + if t.closed { + return fs.ErrClosed + } + if strings.TrimSpace(t.target) == "" || t.root == nil { + return fmt.Errorf("%w: artifact sync target is required", ErrArtifactInvalid) + } + _, err := t.root.Stat(".") + return err +} + +func (t *folderTransport) Exchange( + ctx context.Context, local ArtifactStore, +) (retErr error) { + if err := validateTransportStore(ctx, local); err != nil { + return err + } + t.mu.Lock() + defer t.mu.Unlock() + if t.closed || t.root == nil { + return fs.ErrClosed + } + target := t.root + + // Pull the external share first. Each directory page is processed before + // the next one is read, and local membership is a point lookup. + if err := visitFolderOrigins(ctx, target, func(origin string) error { + for _, kind := range transportKinds { + if err := visitFolderKind(ctx, target, origin, kind, func(wire WireRef) error { + ref, err := FromWireRef(wire.Origin, wire.Kind, wire.Name) + if err != nil { + return err + } + entry, found, err := findStoreEntry(ctx, local, ref) + if err != nil { + return err + } + if found { + repaired, err := repairQueuedTransportArtifact(ctx, local, wire, + func(consume func(io.Reader) error) error { + file, _, openErr := openRootRegularFile(target, folderWirePath(wire)) + if openErr != nil { + return openErr + } + defer file.Close() + return consume(file) + }) + if err != nil { + return err + } + if repaired { + return nil + } + if kind == KindCheckpoints { + return compareFolderCheckpoint(ctx, local, entry, target, wire) + } + return nil + } + if err := receiveFolderWire(ctx, local, target, wire); err != nil { + if errors.Is(err, ErrArtifactInvalid) || errors.Is(err, ErrArtifactCorrupt) { + return nil + } + return err + } + return nil + }); err != nil { + return fmt.Errorf("fetching %s artifacts for %s: %w", kind, origin, err) + } + } + return nil + }); err != nil { + return err + } + + // Publish local pages in dependency-before-checkpoint order. Remote + // membership remains a confined point lookup, so no full share index is + // materialized. + return visitTransportStoreOrigins(ctx, local, func(origin string) error { + for _, kind := range transportKinds { + if err := visitStoreKind(ctx, local, origin, kind, func(entry Entry) error { + wire, err := ToWireRef(entry.Ref) + if err != nil { + return err + } + has, err := folderHasWire(target, wire) + if err != nil { + return err + } + if has { + repaired, err := repairQueuedTransportArtifact(ctx, local, wire, + func(consume func(io.Reader) error) error { + file, _, openErr := openRootRegularFile(target, folderWirePath(wire)) + if openErr != nil { + return openErr + } + defer file.Close() + return consume(file) + }) + if err != nil { + return err + } + if repaired { + return nil + } + if kind == KindCheckpoints { + return compareFolderCheckpoint(ctx, local, entry, target, wire) + } + matches, err := folderWireMatchesEntry(ctx, target, wire, entry) + if errors.Is(err, fs.ErrPermission) { + return nil + } + if err == nil && matches { + return nil + } + if err != nil && !errors.Is(err, ErrArtifactInvalid) && !errors.Is(err, ErrArtifactCorrupt) { + return err + } + quarantineArtifactRoot(target, folderWirePath(wire)) + return publishFolderWire(ctx, local, target, entry) + } + return publishFolderWire(ctx, local, target, entry) + }); err != nil { + return fmt.Errorf("publishing %s artifacts for %s: %w", kind, origin, err) + } + } + return nil + }) +} + +func (t *folderTransport) Close() error { + if t == nil { + return nil + } + t.mu.Lock() + defer t.mu.Unlock() + if t.closed { + return nil + } + t.closed = true + if t.root == nil { + return nil + } + err := t.root.Close() + t.root = nil + return err +} + +func openFolderTransport(target string) (*folderTransport, error) { + if strings.TrimSpace(target) == "" { + return nil, fmt.Errorf("%w: artifact sync target is required", ErrArtifactInvalid) + } + if err := os.MkdirAll(target, 0o755); err != nil { + return nil, fmt.Errorf("creating artifact sync target: %w", err) + } + canonical, err := canonicalArtifactPath(target) + if err != nil { + return nil, fmt.Errorf("resolving artifact sync target: %w", err) + } + root, err := openArtifactRoot(canonical, "target exchange") + if err != nil { + return nil, err + } + return &folderTransport{target: canonical, root: root}, nil +} + +func requireTransportContext(ctx context.Context) error { + if ctx == nil { + return fmt.Errorf("%w: context is required", ErrArtifactInvalid) + } + return ctx.Err() +} + +func validateTransportStore(ctx context.Context, store ArtifactStore) error { + if err := requireTransportContext(ctx); err != nil { + return err + } + if store == nil || isTypedNil(store) { + return fmt.Errorf("%w: artifact store is required", ErrArtifactInvalid) + } + return nil +} + +func openStoreOriginIterator(ctx context.Context, store ArtifactStore) (OriginIterator, error) { + return store.Origins(ctx) +} + +func openStoreEntryIterator( + ctx context.Context, store ArtifactStore, origin string, kind Kind, +) (EntryIterator, error) { + return store.Entries(ctx, origin, kind) +} + +func visitTransportStoreOrigins( + ctx context.Context, store ArtifactStore, visit func(string) error, +) (retErr error) { + iterator, err := openStoreOriginIterator(ctx, store) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + for { + origins, nextErr := iterator.Next(ctx, transportPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return nextErr + } + for _, origin := range origins { + if err := ctx.Err(); err != nil { + return err + } + if err := visit(origin); err != nil { + return err + } + } + if errors.Is(nextErr, io.EOF) { + return nil + } + } +} + +type artifactItem struct { + kind string + name string +} + +func indexItems(index OriginArtifactIndex) []artifactItem { + items := make([]artifactItem, 0, + len(index.Segments)+len(index.Raw)+len(index.Manifests)+len(index.Meta)+len(index.Checkpoints)) + for _, group := range []struct { + kind string + names []string + }{ + {KindSegments, index.Segments}, + {KindRaw, index.Raw}, + {KindManifests, index.Manifests}, + {KindMeta, index.Meta}, + {KindCheckpoints, index.Checkpoints}, + } { + for _, name := range group.names { + items = append(items, artifactItem{kind: group.kind, name: name}) + } + } + return items +} + +func visitStoreKind( + ctx context.Context, + store ArtifactStore, + origin string, + kind Kind, + visit func(Entry) error, +) (retErr error) { + iterator, err := openStoreEntryIterator(ctx, store, origin, kind) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + for { + entries, nextErr := iterator.Next(ctx, transportPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return nextErr + } + for _, entry := range entries { + if err := ctx.Err(); err != nil { + return err + } + if err := visit(entry); err != nil { + return err + } + } + if errors.Is(nextErr, io.EOF) { + return nil + } + } +} + +type storeWireIterator struct { + store ArtifactStore + origin string + kindIndex int + iterator EntryIterator + items []Entry + itemIndex int +} + +func newStoreWireIterator(store ArtifactStore, origin string) *storeWireIterator { + return &storeWireIterator{store: store, origin: origin} +} + +func (i *storeWireIterator) Next(ctx context.Context) (WireRef, Entry, bool, error) { + for { + if i.itemIndex < len(i.items) { + entry := i.items[i.itemIndex] + i.itemIndex++ + wire, err := ToWireRef(entry.Ref) + return wire, entry, err == nil, err + } + if i.kindIndex >= len(transportKinds) { + return WireRef{}, Entry{}, false, nil + } + if i.iterator == nil { + iterator, err := openStoreEntryIterator( + ctx, i.store, i.origin, transportKinds[i.kindIndex], + ) + if err != nil { + return WireRef{}, Entry{}, false, err + } + i.iterator = iterator + } + entries, nextErr := i.iterator.Next(ctx, transportPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return WireRef{}, Entry{}, false, nextErr + } + if len(entries) == 0 && !errors.Is(nextErr, io.EOF) { + return WireRef{}, Entry{}, false, + fmt.Errorf("%w: local entry iterator made no progress", ErrArtifactInvalid) + } + i.items = entries + i.itemIndex = 0 + if errors.Is(nextErr, io.EOF) { + if err := i.iterator.Close(); err != nil { + return WireRef{}, Entry{}, false, err + } + i.iterator = nil + i.kindIndex++ + } + } +} + +func (i *storeWireIterator) Close() error { + if i.iterator == nil { + return nil + } + return i.iterator.Close() +} + +func compareWireRefs(left, right WireRef) int { + if left.Origin != right.Origin { + return strings.Compare(left.Origin, right.Origin) + } + leftKind, rightKind := transportKindRank(left.Kind), transportKindRank(right.Kind) + if leftKind != rightKind { + if leftKind < rightKind { + return -1 + } + return 1 + } + return strings.Compare(left.Name, right.Name) +} + +func transportKindRank(kind Kind) int { + for index, candidate := range transportKinds { + if candidate == kind { + return index + } + } + return len(transportKinds) +} + +func findStoreEntry( + ctx context.Context, store ArtifactStore, ref Ref, +) (Entry, bool, error) { + entry, err := store.Stat(ctx, ref) + if errors.Is(err, ErrArtifactNotFound) { + return Entry{}, false, nil + } + if err != nil { + return Entry{}, false, err + } + return entry, true, nil +} + +func createTransportArtifactFromWire( + ctx context.Context, store ArtifactStore, wire WireRef, body io.Reader, +) (CreateResult, error) { + result, err := CreateFromWire(ctx, store, wire, body, transportWireLimits(wire.Kind)) + if err != nil || !result.Created { + return result, err + } + if err := recordTransportChanged(ctx, store, result.Entry); err != nil { + return result, err + } + return result, nil +} + +func recordTransportChanged(ctx context.Context, store ArtifactStore, entry Entry) error { + recorder, ok := store.(transportChangeStore) + if !ok { + return nil + } + return recorder.RecordTransportChanged(ctx, entry) +} + +func repairQueuedTransportArtifact( + ctx context.Context, + store ArtifactStore, + wire WireRef, + open func(func(io.Reader) error) error, +) (repaired bool, retErr error) { + queue, ok := store.(transportRepairStore) + if !ok { + return false, nil + } + ref, err := FromWireRef(wire.Origin, wire.Kind, wire.Name) + if err != nil { + return false, err + } + request, pending, err := queue.PendingTransportRepair(ctx, ref) + if err != nil || !pending { + return false, err + } + err = open(func(body io.Reader) (callbackErr error) { + spool, identity, err := spoolCanonicalWire( + ctx, wire, body, transportWireLimits(wire.Kind), + ) + if err != nil { + return err + } + defer func() { + callbackErr = errors.Join(callbackErr, closeAndRemoveTransportSpool(spool)) + }() + if identity != request.Identity { + return fmt.Errorf( + "%w: trusted repair identity differs for %s/%s/%s", + ErrArtifactConflict, ref.Origin, ref.Kind, ref.Name, + ) + } + if err := queue.RepairTransportArtifact(ctx, request, spool); err != nil { + return err + } + if err := recordTransportChanged(ctx, store, request); err != nil { + return err + } + return queue.AcknowledgeTransportRepair(ctx, request) + }) + return err == nil, err +} + +func visitFolderOrigins( + ctx context.Context, root *os.Root, visit func(string) error, +) error { + dir, err := root.Open(".") + if err != nil { + return err + } + defer dir.Close() + for { + if err := ctx.Err(); err != nil { + return err + } + entries, err := dir.ReadDir(transportPageSize) + if err != nil && !errors.Is(err, io.EOF) { + return err + } + for _, entry := range entries { + origin := entry.Name() + if validateOriginID(origin) != nil { + continue + } + info, statErr := root.Lstat(origin) + if statErr != nil { + return statErr + } + if !info.IsDir() { + return fmt.Errorf("artifact origin %s is not a directory", origin) + } + if err := visit(origin); err != nil { + return err + } + } + if errors.Is(err, io.EOF) { + return nil + } + } +} + +func visitFolderKind( + ctx context.Context, + root *os.Root, + origin string, + kind Kind, + visit func(WireRef) error, +) error { + rel := filepath.Join(origin, string(kind)) + info, err := root.Lstat(rel) + if errors.Is(err, fs.ErrNotExist) { + return nil + } + if err != nil { + return err + } + if !info.IsDir() { + return fmt.Errorf("artifact kind %s is not a directory", kind) + } + kindRoot, err := openArtifactSubroot(root, rel, "folder kind") + if err != nil { + return err + } + defer kindRoot.Close() + dir, err := kindRoot.Open(".") + if err != nil { + return err + } + defer dir.Close() + for { + if err := ctx.Err(); err != nil { + return err + } + entries, err := dir.ReadDir(transportPageSize) + if err != nil && !errors.Is(err, io.EOF) { + return err + } + for _, entry := range entries { + name := entry.Name() + if isTempArtifactEntry(name) { + continue + } + ref, refErr := FromWireRef(origin, kind, name) + if refErr != nil { + continue + } + fileInfo, statErr := kindRoot.Lstat(name) + if statErr != nil { + return statErr + } + if !fileInfo.Mode().IsRegular() { + return fmt.Errorf("artifact %s/%s is not a regular file", kind, name) + } + wire, refErr := ToWireRef(ref) + if refErr != nil { + return refErr + } + if err := visit(wire); err != nil { + return err + } + } + if errors.Is(err, io.EOF) { + return nil + } + } +} + +func folderWirePath(wire WireRef) string { + return filepath.Join(wire.Origin, string(wire.Kind), wire.Name) +} + +func folderHasWire(root *os.Root, wire WireRef) (bool, error) { + info, err := root.Lstat(folderWirePath(wire)) + if errors.Is(err, fs.ErrNotExist) { + return false, nil + } + if err != nil { + return false, err + } + if !info.Mode().IsRegular() { + return false, fmt.Errorf("artifact destination %s is not a regular file", folderWirePath(wire)) + } + return true, nil +} + +func folderWireMatchesEntry( + ctx context.Context, root *os.Root, wire WireRef, entry Entry, +) (bool, error) { + file, _, err := openRootRegularFile(root, folderWirePath(wire)) + if err != nil { + return false, err + } + defer file.Close() + identity, err := wireIdentity(ctx, wire, file, transportWireLimits(wire.Kind)) + if err != nil { + return false, err + } + return identity == entry.Identity, nil +} + +func receiveFolderWire( + ctx context.Context, store ArtifactStore, root *os.Root, wire WireRef, +) error { + file, _, err := openRootRegularFile(root, folderWirePath(wire)) + if err != nil { + return err + } + defer file.Close() + _, err = createTransportArtifactFromWire(ctx, store, wire, file) + return err +} + +func publishFolderWire( + ctx context.Context, store ArtifactStore, root *os.Root, entry Entry, +) (retErr error) { + wire, err := ToWireRef(entry.Ref) + if err != nil { + return err + } + spool, size, _, err := spoolWireArtifact(ctx, store, entry) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) }() + rel := folderWirePath(wire) + if err := root.MkdirAll(filepath.Dir(rel), 0o755); err != nil { + return err + } + tmp, tmpRel, err := createRootTemp(root, filepath.Dir(rel)) + if err != nil { + return err + } + defer func() { _ = root.Remove(tmpRel) }() + if _, err := io.CopyN(tmp, spool, size); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Chmod(0o644); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + if err := root.Link(tmpRel, rel); err != nil { + if errors.Is(err, fs.ErrExist) { + has, statErr := folderHasWire(root, wire) + if statErr != nil { + return statErr + } + if has { + matches, matchErr := folderWireMatchesEntry(ctx, root, wire, entry) + if matchErr == nil && matches { + return nil + } + return fmt.Errorf("%w: immutable artifact collision at %s: %v", + ErrArtifactConflict, rel, matchErr) + } + } + return fmt.Errorf("publishing immutable artifact %s: %w", rel, err) + } + return nil +} + +func compareFolderCheckpoint( + ctx context.Context, store ArtifactStore, local Entry, root *os.Root, wire WireRef, +) error { + ref, err := FromWireRef(wire.Origin, wire.Kind, wire.Name) + if err != nil { + return err + } + file, _, err := openRootRegularFile(root, folderWirePath(wire)) + if err != nil { + return err + } + defer file.Close() + return compareOrRepairCheckpoint(ctx, store, local, wire, file, + "target", ref.Origin, ref.Name) +} + +func spoolWireArtifact( + ctx context.Context, store ArtifactStore, expected Entry, +) (_ *os.File, _ int64, _ string, retErr error) { + entry, reader, err := store.Open(ctx, expected.Ref) + if err != nil { + return nil, 0, "", err + } + defer func() { retErr = errors.Join(retErr, reader.Close()) }() + if entry.Ref != expected.Ref || entry.Identity != expected.Identity { + return nil, 0, "", fmt.Errorf("%w: artifact changed while opening", ErrArtifactConflict) + } + spool, err := os.CreateTemp("", "agentsview-artifact-wire-*") + if err != nil { + return nil, 0, "", err + } + cleanup := true + defer func() { + if cleanup { + retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) + } + }() + if err := spool.Chmod(0o600); err != nil { + return nil, 0, "", err + } + hasher := sha256.New() + if err := EncodeWire(ctx, expected.Ref, reader, io.MultiWriter(spool, hasher)); err != nil { + return nil, 0, "", err + } + if err := reader.Verify(); err != nil { + return nil, 0, "", err + } + if err := ctx.Err(); err != nil { + return nil, 0, "", err + } + info, err := spool.Stat() + if err != nil { + return nil, 0, "", err + } + if err := spool.Sync(); err != nil { + return nil, 0, "", err + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return nil, 0, "", err + } + cleanup = false + return spool, info.Size(), hex.EncodeToString(hasher.Sum(nil)), nil +} + +func wireIdentity( + ctx context.Context, wire WireRef, src io.Reader, limits WireLimits, +) (Identity, error) { + hasher := sha256.New() + var size int64 + writer := writerFunc(func(p []byte) (int, error) { + n, err := hasher.Write(p) + size += int64(n) + return n, err + }) + if err := DecodeWire(ctx, wire, src, writer, limits); err != nil { + return Identity{}, err + } + return NewIdentity(hex.EncodeToString(hasher.Sum(nil)), size) +} + +func compareOrRepairCheckpoint( + ctx context.Context, + store ArtifactStore, + local Entry, + wire WireRef, + remote io.Reader, + remoteLabel, origin, name string, +) (retErr error) { + spool, identity, err := spoolCanonicalWire(ctx, wire, remote, transportWireLimits(wire.Kind)) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) }() + if identity != local.Identity { + return fmt.Errorf( + "%w: checkpoint %s/%s differs between the local store and the %s; was this origin's artifact store rebuilt or its origin id reused?", + errArtifactPathConflict, origin, name, remoteLabel, + ) + } + if err := verifyStoreEntry(ctx, store, local); err == nil { + return nil + } else if !errors.Is(err, ErrArtifactCorrupt) && !errors.Is(err, ErrArtifactInvalid) { + return err + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return err + } + return store.RepairContent(ctx, local.Identity, spool) +} + +func verifyStoreEntry(ctx context.Context, store ArtifactStore, expected Entry) (retErr error) { + entry, reader, err := store.Open(ctx, expected.Ref) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, reader.Close()) }() + if entry.Identity != expected.Identity { + return fmt.Errorf("%w: artifact identity changed while opening", ErrArtifactCorrupt) + } + if _, err := io.Copy(io.Discard, &wireContextReader{ctx: ctx, reader: reader}); err != nil { + return err + } + return reader.Verify() +} + +func spoolCanonicalWire( + ctx context.Context, wire WireRef, src io.Reader, limits WireLimits, +) (_ *os.File, _ Identity, retErr error) { + spool, err := os.CreateTemp("", "agentsview-artifact-canonical-*") + if err != nil { + return nil, Identity{}, err + } + cleanup := true + defer func() { + if cleanup { + retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) + } + }() + if err := spool.Chmod(0o600); err != nil { + return nil, Identity{}, err + } + hasher := sha256.New() + var size int64 + writer := writerFunc(func(p []byte) (int, error) { + n, err := io.MultiWriter(spool, hasher).Write(p) + size += int64(n) + return n, err + }) + if err := DecodeWire(ctx, wire, src, writer, limits); err != nil { + return nil, Identity{}, err + } + identity, err := NewIdentity(hex.EncodeToString(hasher.Sum(nil)), size) + if err != nil { + return nil, Identity{}, err + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return nil, Identity{}, err + } + cleanup = false + return spool, identity, nil +} + +type writerFunc func([]byte) (int, error) + +func (f writerFunc) Write(p []byte) (int, error) { return f(p) } + +func transportWireLimits(kind Kind) WireLimits { + var decoded int64 + switch kind { + case KindManifests, KindMeta: + decoded = manifestDecodedLimit + case KindSegments, KindCheckpoints: + decoded = segmentDecodedLimit + case KindRaw: + decoded = 1 << 40 + default: + decoded = 1 + } + return WireLimits{MaxEncodedBytes: decoded + (1 << 20), MaxDecodedBytes: decoded} +} + +func closeAndRemoveTransportSpool(file *os.File) error { + if file == nil { + return nil + } + name := file.Name() + return errors.Join(file.Close(), removeTransportSpool(name)) +} + +func removeTransportSpool(name string) error { + err := os.Remove(name) + if errors.Is(err, fs.ErrNotExist) { + return nil + } + return err +} diff --git a/internal/artifact/transport_folder_test.go b/internal/artifact/transport_folder_test.go new file mode 100644 index 000000000..eda06d9f6 --- /dev/null +++ b/internal/artifact/transport_folder_test.go @@ -0,0 +1,144 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "io" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +func openFolderTransportForTest(t *testing.T, target string) *folderTransport { + t.Helper() + transport, err := openFolderTransport(target) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, transport.Close()) }) + return transport +} + +type forbidArtifactOpenStore struct{ ArtifactStore } + +func (s forbidArtifactOpenStore) Open(ctx context.Context, ref Ref) (Entry, VerifiedReader, error) { + if ref.Kind == KindManifests || ref.Kind == KindSegments { + return Entry{}, nil, errors.New("unchanged artifact content must not be reopened") + } + return s.ArtifactStore.Open(ctx, ref) +} + +func TestFolderTransportUnchangedExchangeDoesNotReopenCommonContent(t *testing.T) { + ctx := context.Background() + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + target := t.TempDir() + transport := openFolderTransportForTest(t, target) + require.NoError(t, transport.Exchange(ctx, localStore)) + + require.NoError(t, transport.Exchange(ctx, forbidArtifactOpenStore{localStore}), + "name-set convergence must not reread payloads already held by both stores") +} + +func TestFolderTransportRejectsSymlinkedArtifactKind(t *testing.T) { + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + target := t.TempDir() + segments := filepath.Join(target, origin, KindSegments) + external := filepath.Join(t.TempDir(), KindSegments) + require.NoError(t, os.MkdirAll(filepath.Dir(segments), 0o755)) + require.NoError(t, os.MkdirAll(external, 0o755)) + require.NoError(t, os.Symlink(external, segments)) + + transport := openFolderTransportForTest(t, target) + err := transport.Exchange(context.Background(), localStore) + + require.Error(t, err) + targetSegments, globErr := filepath.Glob( + filepath.Join(target, origin, KindSegments, "*"), + ) + require.NoError(t, globErr) + assert.Empty(t, targetSegments, + "folder exchange must not follow an artifact-kind symlink") +} + +func TestPublishFolderWireRejectsDivergentImmutableCollision(t *testing.T) { + body := []byte("canonical body") + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + store := openTransportStore(t, t.TempDir()) + created, err := store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + wire, err := ToWireRef(ref) + require.NoError(t, err) + targetPath := t.TempDir() + require.NoError(t, os.MkdirAll( + filepath.Join(targetPath, wire.Origin, string(wire.Kind)), 0o755, + )) + require.NoError(t, os.WriteFile( + filepath.Join(targetPath, folderWirePath(wire)), []byte("collision"), 0o644, + )) + target, err := openArtifactRoot(targetPath, "test target") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, target.Close()) }) + + err = publishFolderWire(t.Context(), store, target, created.Entry) + require.ErrorIs(t, err, ErrArtifactConflict) + got, readErr := os.ReadFile(filepath.Join(targetPath, folderWirePath(wire))) + require.NoError(t, readErr) + assert.Equal(t, []byte("collision"), got) +} + +func TestTransportMemoryFolderExchangeRemainsBounded(t *testing.T) { + measure := func(size int64) (uint64, uint64) { + hasher := sha256.New() + _, err := io.CopyN(hasher, repeatedByteReader('x'), size) + require.NoError(t, err) + identity, err := NewIdentity(hex.EncodeToString(hasher.Sum(nil)), size) + require.NoError(t, err) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + source := openTransportStore(t, t.TempDir()) + _, err = source.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), + io.LimitReader(repeatedByteReader('x'), size), + ) + require.NoError(t, err) + target := t.TempDir() + transport := openFolderTransportForTest(t, target) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + require.NoError(t, transport.Exchange(t.Context(), source)) + runtime.ReadMemStats(&after) + outbound := after.TotalAlloc - before.TotalAlloc + + destination := openTransportStore(t, t.TempDir()) + runtime.GC() + runtime.ReadMemStats(&before) + require.NoError(t, transport.Exchange(t.Context(), destination)) + runtime.ReadMemStats(&after) + _, err = destination.Stat(t.Context(), ref) + require.NoError(t, err) + return outbound, after.TotalAlloc - before.TotalAlloc + } + + smallOut, smallIn := measure(1 << 20) + largeOut, largeIn := measure(24 << 20) + assert.Less(t, largeOut, smallOut+(4<<20), + "real folder Exchange upload allocation must remain bounded") + assert.Less(t, largeIn, smallIn+(4<<20), + "real folder Exchange download allocation must remain bounded") +} diff --git a/internal/artifact/transport_http.go b/internal/artifact/transport_http.go new file mode 100644 index 000000000..4c428b11a --- /dev/null +++ b/internal/artifact/transport_http.go @@ -0,0 +1,744 @@ +package artifact + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "log" + "net/http" + "net/url" + "os" + "strconv" + "strings" + "sync" + "time" +) + +const ( + artifactAPIPath = "/api/v1/artifacts" + httpTransportTimeout = 120 * time.Second + httpCursorCleanupLimit = 750 * time.Millisecond + httpTransportMaxErrLen = 512 + peerImportModeHeader = "X-Agentsview-Artifact-Import" + peerImportModeDeferred = "deferred" +) + +var errHTTPPeer = errors.New("artifact peer request failed") + +func IsHTTPTarget(target string) bool { + return strings.HasPrefix(target, "http://") || strings.HasPrefix(target, "https://") +} + +type httpTransport struct { + stateMu sync.Mutex + base string + origin string + token string + client *http.Client + preparedOrigins []string + preparedCursor string + hasPreparedOrigins bool + closed bool +} + +func newHTTPTransport(target, token string, allowInsecure bool) (*httpTransport, error) { + u, err := url.Parse(strings.TrimRight(target, "/")) + if err != nil { + return nil, fmt.Errorf("parsing artifact peer URL: %w", err) + } + if u.Scheme != "http" && u.Scheme != "https" { + return nil, errors.New("artifact peer target must use http:// or https://") + } + if u.Host == "" { + return nil, errors.New("artifact peer target is missing a host") + } + if u.User != nil || u.RawQuery != "" || u.Fragment != "" { + return nil, errors.New("artifact peer target must not contain credentials, query, or fragment") + } + if u.Scheme == "http" && !allowInsecure && !isLoopbackEndpointHost(u.Hostname()) { + return nil, errors.New("insecure artifact peer requires HTTPS") + } + if u.Scheme == "http" && allowInsecure && !isLoopbackEndpointHost(u.Hostname()) { + log.Print("warning: artifact sync uses plaintext HTTP; credentials and archive content are not encrypted in transit") + } + base := strings.TrimRight(u.String(), "/") + if !strings.HasSuffix(base, artifactAPIPath) { + base += artifactAPIPath + } + return &httpTransport{ + base: base, + origin: (&url.URL{Scheme: u.Scheme, Host: u.Host}).String(), + token: token, + client: &http.Client{ + Timeout: httpTransportTimeout, + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + }, nil +} + +func (t *httpTransport) Prepare(ctx context.Context, _ ArtifactStore) error { + if err := requireTransportContext(ctx); err != nil { + return err + } + t.stateMu.Lock() + if t.closed { + t.stateMu.Unlock() + return fs.ErrClosed + } + previousCursor := t.preparedCursor + t.preparedOrigins = nil + t.preparedCursor = "" + t.hasPreparedOrigins = false + t.stateMu.Unlock() + if previousCursor != "" { + if err := t.releaseCursor(previousCursor); err != nil { + return fmt.Errorf("releasing previous artifact peer cursor: %w", err) + } + } + page, err := t.getOriginsPage(ctx, "") + if err != nil { + return fmt.Errorf("connecting to artifact peer: %w", err) + } + t.stateMu.Lock() + if t.closed { + t.stateMu.Unlock() + return errors.Join(fs.ErrClosed, t.releaseCursor(page.NextCursor)) + } + t.preparedOrigins = append(t.preparedOrigins[:0], page.Origins...) + t.preparedCursor = page.NextCursor + t.hasPreparedOrigins = true + t.stateMu.Unlock() + return nil +} + +func (t *httpTransport) Exchange(ctx context.Context, local ArtifactStore) (retErr error) { + if err := validateTransportStore(ctx, local); err != nil { + return err + } + t.stateMu.Lock() + closed := t.closed + t.stateMu.Unlock() + if closed { + return fs.ErrClosed + } + remoteIt := t.exchangeOrigins() + localIt := &storeOriginIterator{store: local} + defer func() { + retErr = errors.Join(retErr, remoteIt.Close(ctx), localIt.Close()) + }() + remoteOrigin, remoteOK, err := remoteIt.Next(ctx) + if err != nil { + return fmt.Errorf("fetching artifact origins from peer: %w", err) + } + localOrigin, localOK, err := localIt.Next(ctx) + if err != nil { + return fmt.Errorf("listing local artifact origins: %w", err) + } + for remoteOK || localOK { + var origin string + switch { + case !localOK || (remoteOK && remoteOrigin < localOrigin): + origin = remoteOrigin + remoteOrigin, remoteOK, err = remoteIt.Next(ctx) + case !remoteOK || localOrigin < remoteOrigin: + origin = localOrigin + localOrigin, localOK, err = localIt.Next(ctx) + default: + origin = remoteOrigin + remoteOrigin, remoteOK, err = remoteIt.Next(ctx) + if err == nil { + localOrigin, localOK, err = localIt.Next(ctx) + } + } + if err != nil { + return fmt.Errorf("advancing artifact origins: %w", err) + } + if err := t.exchangeOrigin(ctx, local, origin); err != nil { + return fmt.Errorf("exchanging peer artifacts for %s: %w", origin, err) + } + } + if err := t.finalizePush(ctx); err != nil { + return fmt.Errorf("finalizing peer artifact batch: %w", err) + } + return nil +} + +func (t *httpTransport) exchangeOrigins() *httpOriginIterator { + t.stateMu.Lock() + defer t.stateMu.Unlock() + if t.hasPreparedOrigins { + iterator := &httpOriginIterator{ + transport: t, + items: append([]string(nil), t.preparedOrigins...), + cursor: t.preparedCursor, + } + t.preparedOrigins = nil + t.preparedCursor = "" + t.hasPreparedOrigins = false + if iterator.cursor == "" { + iterator.done = true + } + return iterator + } + return &httpOriginIterator{transport: t} +} + +func (t *httpTransport) Close() error { + if t == nil { + return nil + } + t.stateMu.Lock() + if t.closed { + t.stateMu.Unlock() + return nil + } + t.closed = true + cursor := t.preparedCursor + t.preparedOrigins = nil + t.preparedCursor = "" + t.hasPreparedOrigins = false + t.stateMu.Unlock() + return t.releaseCursor(cursor) +} + +func (t *httpTransport) exchangeOrigin( + ctx context.Context, local ArtifactStore, origin string, +) (retErr error) { + localIt := newStoreWireIterator(local, origin) + remoteIt := &httpWireIterator{transport: t, origin: origin} + defer func() { + retErr = errors.Join(retErr, localIt.Close(), remoteIt.Close()) + }() + localWire, localEntry, localOK, err := localIt.Next(ctx) + if err != nil { + return err + } + remoteWire, remoteOK, err := remoteIt.Next(ctx) + if err != nil { + return err + } + for localOK || remoteOK { + if err := ctx.Err(); err != nil { + return err + } + switch { + case !localOK || (remoteOK && compareWireRefs(remoteWire, localWire) < 0): + if err := t.receiveArtifact(ctx, local, remoteWire); err != nil { + if errors.Is(err, ErrArtifactNotFound) { + remoteWire, remoteOK, err = remoteIt.Next(ctx) + if err != nil { + return err + } + continue + } + return err + } + remoteWire, remoteOK, err = remoteIt.Next(ctx) + if err != nil { + return err + } + case !remoteOK || compareWireRefs(localWire, remoteWire) < 0: + if err := t.postEntry(ctx, local, localEntry); err != nil { + return err + } + localWire, localEntry, localOK, err = localIt.Next(ctx) + if err != nil { + return err + } + default: + repaired, err := repairQueuedTransportArtifact(ctx, local, remoteWire, + func(consume func(io.Reader) error) error { + return t.withArtifact(ctx, remoteWire, consume) + }) + if err != nil { + return err + } + if !repaired && localWire.Kind == KindCheckpoints { + if err := t.compareCheckpoint(ctx, local, localEntry, remoteWire); err != nil { + return err + } + } + localWire, localEntry, localOK, err = localIt.Next(ctx) + if err != nil { + return err + } + remoteWire, remoteOK, err = remoteIt.Next(ctx) + if err != nil { + return err + } + } + } + return nil +} + +type httpOriginsPage struct { + Origins []string `json:"origins"` + NextCursor string `json:"next_cursor,omitempty"` +} + +func (t *httpTransport) getOriginsPage(ctx context.Context, cursor string) (httpOriginsPage, error) { + u := t.base + "/origins" + query := url.Values{} + query.Set("limit", strconv.Itoa(transportPageSize)) + if cursor != "" { + query.Set("cursor", cursor) + } + u += "?" + query.Encode() + var page httpOriginsPage + if err := t.getJSON(ctx, u, &page); err != nil { + return httpOriginsPage{}, err + } + if len(page.Origins) > transportPageSize { + return httpOriginsPage{}, fmt.Errorf("%w: peer origin page exceeds %d entries", ErrArtifactInvalid, transportPageSize) + } + for index, origin := range page.Origins { + if err := validateOriginID(origin); err != nil { + return httpOriginsPage{}, fmt.Errorf("%w: invalid peer origin: %v", ErrArtifactInvalid, err) + } + if index > 0 && page.Origins[index-1] >= origin { + return httpOriginsPage{}, fmt.Errorf("%w: peer origins are not strictly increasing", ErrArtifactInvalid) + } + } + return page, nil +} + +type httpOriginIterator struct { + transport *httpTransport + cursor string + items []string + next int + done bool + guard boundedCursorCycleGuard + last string +} + +func (i *httpOriginIterator) Next(ctx context.Context) (string, bool, error) { + for { + if i.next < len(i.items) { + origin := i.items[i.next] + i.next++ + if i.last != "" && origin <= i.last { + return "", false, fmt.Errorf("%w: peer origins are not strictly increasing", ErrArtifactInvalid) + } + i.last = origin + return origin, true, nil + } + if i.done { + return "", false, nil + } + page, err := i.transport.getOriginsPage(ctx, i.cursor) + if err != nil { + return "", false, err + } + i.items = page.Origins + i.next = 0 + if page.NextCursor == "" { + i.cursor = "" + i.done = true + } else { + if i.guard.Observe(Cursor(page.NextCursor)) { + return "", false, errors.New("artifact peer origin cursor cycle") + } + i.cursor = page.NextCursor + } + } +} + +func (i *httpOriginIterator) Close(_ context.Context) error { + if i == nil || i.cursor == "" { + return nil + } + err := i.transport.releaseCursor(i.cursor) + i.cursor = "" + return err +} + +type storeOriginIterator struct { + store ArtifactStore + iterator OriginIterator + items []string + next int + done bool + last string +} + +func (i *storeOriginIterator) Next(ctx context.Context) (string, bool, error) { + for { + if i.next < len(i.items) { + origin := i.items[i.next] + i.next++ + if i.last != "" && origin <= i.last { + return "", false, fmt.Errorf("%w: local origins are not strictly increasing", ErrArtifactInvalid) + } + i.last = origin + return origin, true, nil + } + if i.done { + return "", false, nil + } + if i.iterator == nil { + iterator, err := openStoreOriginIterator(ctx, i.store) + if err != nil { + return "", false, err + } + i.iterator = iterator + } + origins, nextErr := i.iterator.Next(ctx, transportPageSize) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return "", false, nextErr + } + if len(origins) == 0 && !errors.Is(nextErr, io.EOF) { + return "", false, fmt.Errorf("%w: local origin iterator made no progress", ErrArtifactInvalid) + } + if len(origins) > transportPageSize { + return "", false, fmt.Errorf("%w: local origin page exceeds %d entries", ErrArtifactInvalid, transportPageSize) + } + for _, origin := range origins { + if err := validateOriginID(origin); err != nil { + return "", false, fmt.Errorf("%w: invalid local origin: %v", ErrArtifactInvalid, err) + } + } + i.items = origins + i.next = 0 + if errors.Is(nextErr, io.EOF) { + i.done = true + } + } +} + +func (i *storeOriginIterator) Close() error { + if i.iterator == nil { + return nil + } + return i.iterator.Close() +} + +type httpIndexPage struct { + OriginArtifactIndex + NextCursor string `json:"next_cursor,omitempty"` +} + +func (t *httpTransport) getIndexPage( + ctx context.Context, origin, cursor string, +) (httpIndexPage, error) { + u := t.base + "/" + url.PathEscape(origin) + "/index" + q := url.Values{} + q.Set("limit", strconv.Itoa(transportPageSize)) + if cursor != "" { + q.Set("cursor", cursor) + } + u += "?" + q.Encode() + var page httpIndexPage + if err := t.getJSON(ctx, u, &page); err != nil { + return httpIndexPage{}, err + } + if page.Origin != origin { + return httpIndexPage{}, fmt.Errorf("%w: peer index origin mismatch", ErrArtifactInvalid) + } + return page, nil +} + +func (t *httpTransport) receiveArtifact( + ctx context.Context, local ArtifactStore, wire WireRef, +) error { + return t.withArtifact(ctx, wire, func(body io.Reader) error { + _, err := createTransportArtifactFromWire(ctx, local, wire, body) + if errors.Is(err, ErrArtifactInvalid) || errors.Is(err, ErrArtifactCorrupt) { + log.Printf("artifact: skipping corrupt artifact %s/%s/%s from peer: %v", + wire.Origin, wire.Kind, wire.Name, err) + return nil + } + return err + }) +} + +func (t *httpTransport) compareCheckpoint( + ctx context.Context, store ArtifactStore, local Entry, wire WireRef, +) error { + return t.withArtifact(ctx, wire, func(body io.Reader) error { + return compareOrRepairCheckpoint(ctx, store, local, wire, body, + "remote", wire.Origin, local.Ref.Name) + }) +} + +func (t *httpTransport) withArtifact( + ctx context.Context, wire WireRef, consume func(io.Reader) error, +) error { + u := t.base + "/" + url.PathEscape(wire.Origin) + "/" + + url.PathEscape(string(wire.Kind)) + "/" + url.PathEscape(wire.Name) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil) + if err != nil { + return err + } + t.authorize(req) + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode == http.StatusNotFound { + return fmt.Errorf("%w: %s", ErrArtifactNotFound, resp.Status) + } + if resp.StatusCode != http.StatusOK { + return httpStatusError(resp) + } + return consume(resp.Body) +} + +func (t *httpTransport) postEntry( + ctx context.Context, local ArtifactStore, entry Entry, +) (retErr error) { + spool, size, _, err := spoolWireArtifact(ctx, local, entry) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) }() + wire, err := ToWireRef(entry.Ref) + if err != nil { + return err + } + return t.postArtifact(ctx, wire.Origin, string(wire.Kind), wire.Name, spool, size) +} + +func (t *httpTransport) postArtifact( + ctx context.Context, + origin, kind, name string, + body io.ReadSeeker, + size int64, +) error { + u := t.base + "/" + url.PathEscape(origin) + "/" + url.PathEscape(kind) + "/" + url.PathEscape(name) + readerAt := bodyReaderAt(body) + section := io.NewSectionReader(readerAt, 0, size) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, section) + if err != nil { + return err + } + req.ContentLength = size + req.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(io.NewSectionReader(readerAt, 0, size)), nil + } + req.Header.Set("Content-Type", "application/octet-stream") + req.Header.Set("Origin", t.origin) + req.Header.Set(peerImportModeHeader, peerImportModeDeferred) + t.authorize(req) + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated { + return httpStatusError(resp) + } + _, _ = io.Copy(io.Discard, resp.Body) + return nil +} + +type readSeekerAt interface { + io.ReadSeeker + io.ReaderAt +} + +func bodyReaderAt(body io.ReadSeeker) io.ReaderAt { + if at, ok := body.(io.ReaderAt); ok { + return at + } + return &lockedReadSeekerAt{body: body} +} + +type lockedReadSeekerAt struct { + mu sync.Mutex + body io.ReadSeeker +} + +func (r *lockedReadSeekerAt) ReadAt(p []byte, off int64) (int, error) { + r.mu.Lock() + defer r.mu.Unlock() + if _, err := r.body.Seek(off, io.SeekStart); err != nil { + return 0, err + } + return io.ReadFull(r.body, p) +} + +func (t *httpTransport) finalizePush(ctx context.Context) error { + req, err := http.NewRequestWithContext(ctx, http.MethodPost, t.base+"/finalize", nil) + if err != nil { + return err + } + req.Header.Set("Origin", t.origin) + t.authorize(req) + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode == http.StatusNotFound { + _, _ = io.Copy(io.Discard, resp.Body) + return nil + } + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return httpStatusError(resp) + } + _, _ = io.Copy(io.Discard, resp.Body) + return nil +} + +func (t *httpTransport) getJSON(ctx context.Context, u string, out any) error { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil) + if err != nil { + return err + } + t.authorize(req) + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return httpStatusError(resp) + } + body, err := readTransportPage(resp.Body) + if err != nil { + return err + } + return json.Unmarshal(body, out) +} + +func (t *httpTransport) authorize(req *http.Request) { + if t.token != "" { + req.Header.Set("Authorization", "Bearer "+t.token) + } +} + +func (t *httpTransport) releaseCursor(cursor string) error { + if cursor == "" { + return nil + } + ctx, cancel := context.WithTimeout(context.Background(), httpCursorCleanupLimit) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodDelete, + t.base+"/cursors/"+url.PathEscape(cursor), nil) + if err != nil { + return err + } + req.Header.Set("Origin", t.origin) + t.authorize(req) + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return httpStatusError(resp) + } + _, _ = io.Copy(io.Discard, resp.Body) + return nil +} + +func httpStatusError(resp *http.Response) error { + body, _ := io.ReadAll(io.LimitReader(resp.Body, httpTransportMaxErrLen)) + msg := strings.TrimSpace(string(body)) + if resp.StatusCode == http.StatusUnauthorized { + return fmt.Errorf("%w: peer rejected the bearer token (401)", errHTTPPeer) + } + if msg == "" { + return fmt.Errorf("%w: %s", errHTTPPeer, resp.Status) + } + return fmt.Errorf("%w: %s: %s", errHTTPPeer, resp.Status, msg) +} + +type httpWireIterator struct { + transport *httpTransport + origin string + cursor string + items []WireRef + next int + done bool + guard boundedCursorCycleGuard + last WireRef + hasLast bool +} + +func (i *httpWireIterator) Next(ctx context.Context) (WireRef, bool, error) { + for { + if i.next < len(i.items) { + wire := i.items[i.next] + i.next++ + return wire, true, nil + } + if i.done { + return WireRef{}, false, nil + } + page, err := i.transport.getIndexPage(ctx, i.origin, i.cursor) + if err != nil { + return WireRef{}, false, err + } + i.items, err = wireRefsFromIndex(page.OriginArtifactIndex) + if err != nil { + return WireRef{}, false, err + } + i.next = 0 + for _, item := range i.items { + if i.hasLast && compareWireRefs(i.last, item) >= 0 { + return WireRef{}, false, fmt.Errorf("%w: peer artifact index is not strictly increasing", ErrArtifactInvalid) + } + i.last = item + i.hasLast = true + } + if page.NextCursor == "" { + i.done = true + } else { + if i.guard.Observe(Cursor(page.NextCursor)) { + return WireRef{}, false, errors.New("artifact peer index cursor cycle") + } + i.cursor = page.NextCursor + } + } +} + +func (i *httpWireIterator) Close() error { + return i.transport.releaseCursor(i.cursor) +} + +func wireRefsFromIndex(index OriginArtifactIndex) ([]WireRef, error) { + total := len(index.Segments) + len(index.Raw) + len(index.Manifests) + len(index.Meta) + len(index.Checkpoints) + if total > transportPageSize { + return nil, fmt.Errorf("%w: peer artifact index page exceeds %d entries", ErrArtifactInvalid, transportPageSize) + } + groups := [...]struct { + kind Kind + names []string + }{ + {KindSegments, index.Segments}, + {KindRaw, index.Raw}, + {KindManifests, index.Manifests}, + {KindMeta, index.Meta}, + {KindCheckpoints, index.Checkpoints}, + } + items := make([]WireRef, 0, total) + for _, group := range groups { + for _, name := range group.names { + ref, err := FromWireRef(index.Origin, group.kind, name) + if err != nil { + return nil, err + } + wire, err := ToWireRef(ref) + if err != nil { + return nil, err + } + items = append(items, wire) + } + } + for index := 1; index < len(items); index++ { + if compareWireRefs(items[index-1], items[index]) >= 0 { + return nil, fmt.Errorf("%w: peer artifact index is not strictly increasing", ErrArtifactInvalid) + } + } + return items, nil +} + +var _ readSeekerAt = (*os.File)(nil) diff --git a/internal/artifact/transport_http_test.go b/internal/artifact/transport_http_test.go new file mode 100644 index 000000000..03a31cf27 --- /dev/null +++ b/internal/artifact/transport_http_test.go @@ -0,0 +1,1044 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "io" + "log" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "runtime" + "sort" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +// fakeArtifactPeer is an in-memory peer implementing the artifact API surface +// the HTTP transport exchanges against: origin listing, per-origin index, and +// artifact get/post. +type fakeArtifactPeer struct { + mu sync.Mutex + arts map[string][]byte // "origin/kind/name" -> bytes + posts []string + deferredPosts int + finalizeCalls int + supportsFinalize bool +} + +type pagedArtifactStoreObservations struct { + ArtifactStore + mu sync.Mutex + originPages int + listPages map[Kind]int + maxRequestedEntries int + maxReturnedEntries int +} + +func (s *pagedArtifactStoreObservations) Origins(ctx context.Context) (OriginIterator, error) { + iterator, err := s.ArtifactStore.Origins(ctx) + if err != nil { + return nil, err + } + return &testOriginIterator{ + next: func(ctx context.Context, limit int) ([]string, error) { + origins, err := iterator.Next(ctx, limit) + s.mu.Lock() + defer s.mu.Unlock() + s.originPages++ + s.maxRequestedEntries = max(s.maxRequestedEntries, limit) + s.maxReturnedEntries = max(s.maxReturnedEntries, len(origins)) + return origins, err + }, + close: iterator.Close, + }, nil +} + +func (s *pagedArtifactStoreObservations) Entries( + ctx context.Context, origin string, kind Kind, +) (EntryIterator, error) { + iterator, err := s.ArtifactStore.Entries(ctx, origin, kind) + if err != nil { + return nil, err + } + return &testEntryIterator{ + next: func(ctx context.Context, limit int) ([]Entry, error) { + entries, err := iterator.Next(ctx, limit) + s.mu.Lock() + defer s.mu.Unlock() + s.listPages[kind]++ + s.maxRequestedEntries = max(s.maxRequestedEntries, limit) + s.maxReturnedEntries = max(s.maxReturnedEntries, len(entries)) + return entries, err + }, + close: iterator.Close, + }, nil +} + +type queuedTransportRepairStore struct { + ArtifactStore + pending Entry + repairs int +} + +type cleanupFailingRepairStore struct { + ArtifactStore + pending Entry +} + +func (s *cleanupFailingRepairStore) PendingTransportRepair( + context.Context, Ref, +) (Entry, bool, error) { + return s.pending, true, nil +} + +func (s *cleanupFailingRepairStore) RepairTransportArtifact( + _ context.Context, _ Entry, trusted io.Reader, +) error { + file, ok := trusted.(*os.File) + if !ok { + return errors.New("repair spool is not a file") + } + return file.Close() +} + +func (s *cleanupFailingRepairStore) AcknowledgeTransportRepair( + context.Context, Entry, +) error { + return nil +} + +type cancelDuringOpenStore struct { + ArtifactStore + cancel context.CancelFunc + after int +} + +func (s *cancelDuringOpenStore) Open( + ctx context.Context, ref Ref, +) (Entry, VerifiedReader, error) { + entry, reader, err := s.ArtifactStore.Open(ctx, ref) + if err != nil { + return Entry{}, nil, err + } + return entry, &cancelDuringVerifiedRead{ + VerifiedReader: reader, + cancel: s.cancel, + remaining: s.after, + }, nil +} + +type cancelDuringVerifiedRead struct { + VerifiedReader + cancel context.CancelFunc + remaining int + canceled bool +} + +func (r *cancelDuringVerifiedRead) Read(p []byte) (int, error) { + if r.canceled { + return 0, context.Canceled + } + if len(p) > r.remaining { + p = p[:r.remaining] + } + n, err := r.VerifiedReader.Read(p) + r.remaining -= n + if r.remaining <= 0 { + r.cancel() + r.canceled = true + } + return n, err +} + +func (s *queuedTransportRepairStore) PendingTransportRepair( + ctx context.Context, ref Ref, +) (Entry, bool, error) { + if err := ctx.Err(); err != nil { + return Entry{}, false, err + } + return s.pending, s.pending.Ref == ref, nil +} + +func (s *queuedTransportRepairStore) RepairTransportArtifact( + ctx context.Context, entry Entry, trusted io.Reader, +) error { + if entry != s.pending { + return errors.New("unexpected repair identity") + } + if err := s.Quarantine(ctx, entry.Ref, "test repair"); err != nil { + return err + } + if _, err := s.Create(ctx, entry.Ref, entry.Identity, + canonicalArtifactMediaType(entry.Ref.Kind), trusted); err != nil { + return err + } + s.repairs++ + return nil +} + +func (s *queuedTransportRepairStore) AcknowledgeTransportRepair( + ctx context.Context, entry Entry, +) error { + if err := ctx.Err(); err != nil { + return err + } + if entry != s.pending { + return errors.New("unexpected repair acknowledgement") + } + s.pending = Entry{} + return nil +} + +func newFakeArtifactPeer() *fakeArtifactPeer { + return &fakeArtifactPeer{arts: map[string][]byte{}, supportsFinalize: true} +} + +func (p *fakeArtifactPeer) put(origin, kind, name string, data []byte) { + p.mu.Lock() + defer p.mu.Unlock() + p.arts[origin+"/"+kind+"/"+name] = append([]byte(nil), data...) +} + +func (p *fakeArtifactPeer) has(origin, kind, name string) bool { + p.mu.Lock() + defer p.mu.Unlock() + _, ok := p.arts[origin+"/"+kind+"/"+name] + return ok +} + +func (p *fakeArtifactPeer) postedKinds() []string { + p.mu.Lock() + defer p.mu.Unlock() + kinds := make([]string, 0, len(p.posts)) + for _, key := range p.posts { + parts := strings.SplitN(key, "/", 3) + if len(parts) == 3 { + kinds = append(kinds, parts[1]) + } + } + return kinds +} + +func (p *fakeArtifactPeer) batchCounts() (deferredPosts, finalizeCalls int) { + p.mu.Lock() + defer p.mu.Unlock() + return p.deferredPosts, p.finalizeCalls +} + +func (p *fakeArtifactPeer) ServeHTTP(w http.ResponseWriter, r *http.Request) { + rest := strings.TrimPrefix(r.URL.Path, artifactAPIPath+"/") + p.mu.Lock() + defer p.mu.Unlock() + if rest == "origins" { + seen := map[string]bool{} + for key := range p.arts { + seen[strings.SplitN(key, "/", 2)[0]] = true + } + origins := make([]string, 0, len(seen)) + for origin := range seen { + origins = append(origins, origin) + } + sort.Strings(origins) + _ = json.NewEncoder(w).Encode(map[string]any{"origins": origins}) + return + } + if rest == "finalize" && r.Method == http.MethodPost { + if !p.supportsFinalize { + http.NotFound(w, r) + return + } + p.finalizeCalls++ + w.WriteHeader(http.StatusOK) + return + } + parts := strings.Split(rest, "/") + if len(parts) == 2 && parts[1] == "index" { + idx := OriginArtifactIndex{Origin: parts[0]} + for key := range p.arts { + kp := strings.SplitN(key, "/", 3) + if kp[0] != parts[0] { + continue + } + switch kp[1] { + case KindCheckpoints: + idx.Checkpoints = append(idx.Checkpoints, kp[2]) + case KindManifests: + idx.Manifests = append(idx.Manifests, kp[2]) + case KindSegments: + idx.Segments = append(idx.Segments, kp[2]) + case KindMeta: + idx.Meta = append(idx.Meta, kp[2]) + case KindRaw: + idx.Raw = append(idx.Raw, kp[2]) + } + } + sort.Strings(idx.Checkpoints) + sort.Strings(idx.Manifests) + sort.Strings(idx.Segments) + sort.Strings(idx.Meta) + sort.Strings(idx.Raw) + _ = json.NewEncoder(w).Encode(idx) + return + } + if len(parts) == 3 { + switch r.Method { + case http.MethodGet: + data, ok := p.arts[rest] + if !ok { + http.NotFound(w, r) + return + } + _, _ = w.Write(data) + case http.MethodPost: + data, _ := io.ReadAll(r.Body) + p.arts[rest] = data + p.posts = append(p.posts, rest) + if r.Header.Get("X-Agentsview-Artifact-Import") == "deferred" { + p.deferredPosts++ + } + w.WriteHeader(http.StatusCreated) + } + return + } + http.NotFound(w, r) +} + +func TestHTTPTransportRequiresTLSForNonLoopbackPeers(t *testing.T) { + tests := []struct { + name string + target string + wantErr bool + }{ + {name: "public HTTP", target: "http://203.0.113.10:8080", wantErr: true}, + {name: "hostname HTTP", target: "http://peer.example.test:8080", wantErr: true}, + {name: "public HTTPS", target: "https://peer.example.test:8443"}, + {name: "localhost HTTP", target: "http://localhost:8080"}, + {name: "IPv4 loopback HTTP", target: "http://127.0.0.1:8080"}, + {name: "IPv6 loopback HTTP", target: "http://[::1]:8080"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tr, err := newHTTPTransport(tt.target, "", false) + if tt.wantErr { + require.Error(t, err) + assert.Contains(t, err.Error(), "requires HTTPS") + assert.Nil(t, tr) + return + } + require.NoError(t, err) + assert.NotNil(t, tr) + }) + } +} + +func TestHTTPTransportAllowsExplicitRemotePlaintextOptIn(t *testing.T) { + const target = "http://peer.example.test:8080" + var logs bytes.Buffer + previousOutput := log.Writer() + log.SetOutput(&logs) + t.Cleanup(func() { log.SetOutput(previousOutput) }) + + tr, err := newHTTPTransport(target, "", true) + + require.NoError(t, err) + require.NotNil(t, tr) + assert.Equal(t, "http://peer.example.test:8080"+artifactAPIPath, tr.base) + assert.Contains(t, logs.String(), "warning") + assert.Contains(t, logs.String(), "plaintext HTTP") + assert.NotContains(t, logs.String(), target) + assert.NotContains(t, logs.String(), "peer.example.test") +} + +func TestHTTPTransportPrepareHonorsCanceledSync(t *testing.T) { + var requests atomic.Int32 + peer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + _ = json.NewEncoder(w).Encode(map[string]any{"origins": []string{}}) + })) + t.Cleanup(peer.Close) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := Sync(ctx, testDB(t), SyncOptions{ + DataDir: t.TempDir(), + Target: peer.URL, + Origin: "laptop-a1b2c3", + }) + + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, requests.Load(), "canceled preparation must not contact the peer") +} + +func TestHTTPOriginIteratorDetectsArbitraryCursorCycleInConstantState(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodDelete { + w.WriteHeader(http.StatusNoContent) + return + } + requests.Add(1) + next := map[string]string{"": "a", "a": "b", "b": "c", "c": "b"}[r.URL.Query().Get("cursor")] + _ = json.NewEncoder(w).Encode(httpOriginsPage{NextCursor: next}) + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + iterator := &httpOriginIterator{transport: transport} + t.Cleanup(func() { require.NoError(t, iterator.Close(context.Background())) }) + + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.Contains(t, err.Error(), "cursor cycle") + assert.LessOrEqual(t, requests.Load(), int32(8)) +} + +func TestHTTPOriginIteratorRejectsDuplicateAcrossPageBoundary(t *testing.T) { + const origin = "peer-a1b2c3" + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodDelete { + w.WriteHeader(http.StatusNoContent) + return + } + page := httpOriginsPage{Origins: []string{origin}} + if r.URL.Query().Get("cursor") == "" { + page.NextCursor = "next" + } + _ = json.NewEncoder(w).Encode(page) + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + iterator := &httpOriginIterator{transport: transport} + t.Cleanup(func() { require.NoError(t, iterator.Close(context.Background())) }) + + got, ok, err := iterator.Next(t.Context()) + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, origin, got) + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func TestHTTPWireIteratorRejectsDuplicateAcrossPageBoundary(t *testing.T) { + const origin = "peer-a1b2c3" + name := strings64("a") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodDelete { + w.WriteHeader(http.StatusNoContent) + return + } + page := httpIndexPage{OriginArtifactIndex: OriginArtifactIndex{ + Origin: origin, + Raw: []string{name}, + }} + if r.URL.Query().Get("cursor") == "" { + page.NextCursor = "next" + } + _ = json.NewEncoder(w).Encode(page) + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + iterator := &httpWireIterator{transport: transport, origin: origin} + t.Cleanup(func() { require.NoError(t, iterator.Close()) }) + + _, ok, err := iterator.Next(t.Context()) + require.NoError(t, err) + assert.True(t, ok) + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func TestHTTPSyncReusesPreparedOriginsForFirstExchange(t *testing.T) { + peer := newFakeArtifactPeer() + var originRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if strings.HasSuffix(r.URL.Path, "/origins") { + originRequests.Add(1) + } + peer.ServeHTTP(w, r) + })) + t.Cleanup(server.Close) + + _, err := Sync(context.Background(), testDB(t), SyncOptions{ + DataDir: t.TempDir(), + Target: server.URL, + Origin: "laptop-a1b2c3", + }) + require.NoError(t, err) + assert.Equal(t, int32(1), originRequests.Load(), + "prepare and the first pull must share one peer snapshot") +} + +func TestHTTPTransportPullSkipsCorruptRemoteArtifact(t *testing.T) { + origin := "desktop-d4e5f6" + remoteStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-7", "beta") + }) + remoteIdx, err := listTransportArtifacts(context.Background(), remoteStore, origin) + require.NoError(t, err) + peer := newFakeArtifactPeer() + for _, item := range indexItems(remoteIdx) { + art, err := readTransportArtifact(context.Background(), remoteStore, origin, item.kind, item.name) + require.NoError(t, err) + peer.put(origin, item.kind, item.name, art) + } + corruptName := hashHex([]byte("corrupt")) + segmentExtension + peer.put(origin, KindSegments, corruptName, []byte("garbage")) + + srv := httptest.NewServer(peer) + t.Cleanup(srv.Close) + tr, err := newHTTPTransport(srv.URL, "", false) + require.NoError(t, err) + localStore := openTransportStore(t, filepath.Join(t.TempDir(), "artifacts")) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + gotIdx, err := listTransportArtifacts(context.Background(), localStore, origin) + require.NoError(t, err) + assert.ElementsMatch(t, indexItems(remoteIdx), indexItems(gotIdx)) + corruptRef, err := FromWireRef(origin, KindSegments, corruptName) + require.NoError(t, err) + _, err = localStore.Stat(t.Context(), corruptRef) + assert.ErrorIs(t, err, ErrArtifactNotFound) +} + +func TestHTTPTransportIgnoresUncatalogedCorruptLocalArtifact(t *testing.T) { + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + validIdx, err := listTransportArtifacts(context.Background(), localStore, origin) + require.NoError(t, err) + corruptName := hashHex([]byte("junk")) + segmentExtension + corruptRef, err := FromWireRef(origin, KindSegments, corruptName) + require.NoError(t, err) + _, err = localStore.Stat(t.Context(), corruptRef) + require.ErrorIs(t, err, ErrArtifactNotFound) + + peer := newFakeArtifactPeer() + srv := httptest.NewServer(peer) + t.Cleanup(srv.Close) + tr, err := newHTTPTransport(srv.URL, "", false) + require.NoError(t, err) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + for _, item := range indexItems(validIdx) { + assert.True(t, peer.has(origin, item.kind, item.name), + "expected %s/%s on the peer", item.kind, item.name) + } + assert.False(t, peer.has(origin, KindSegments, corruptName)) + _, err = localStore.Stat(t.Context(), corruptRef) + assert.ErrorIs(t, err, ErrArtifactNotFound, + "transport enumeration must not invent uncataloged logical artifacts") +} + +func TestHTTPTransportPushPublishesDependenciesBeforeCheckpoint(t *testing.T) { + localStore := exportStore(t, "laptop-a1b2c3", func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + peer := newFakeArtifactPeer() + server := httptest.NewServer(peer) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + + require.NoError(t, transport.Exchange(context.Background(), localStore)) + + assert.Equal(t, + []string{KindSegments, KindManifests, KindCheckpoints}, + peer.postedKinds(), + ) +} + +func TestHTTPTransportExchangeDetectsDivergentCheckpoint(t *testing.T) { + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + checkpointRef := onlyTransportRef(t, localStore, origin, KindCheckpoints) + wire, err := ToWireRef(checkpointRef) + require.NoError(t, err) + + // The peer holds a different, equally valid checkpoint under the same + // sequence name, as a rebuilt store under a reused origin id would. + divergent, err := canonicalJSON(checkpoint{ + Version: formatVersion, Origin: origin, Sequence: 1, + Sessions: map[string]string{origin + "~other": hashHex([]byte("other"))}, + }) + require.NoError(t, err) + peer := newFakeArtifactPeer() + peer.put(origin, KindCheckpoints, wire.Name, divergent) + + srv := httptest.NewServer(peer) + t.Cleanup(srv.Close) + tr, err := newHTTPTransport(srv.URL, "", false) + require.NoError(t, err) + + err = tr.Exchange(context.Background(), localStore) + require.Error(t, err) + assert.ErrorIs(t, err, errArtifactPathConflict) +} + +func TestHTTPTransportExchangeRepairsCorruptLocalCheckpoint(t *testing.T) { + origin := "laptop-a1b2c3" + baseStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + checkpointRef := onlyTransportRef(t, baseStore, origin, KindCheckpoints) + wire, valid := wireTransportArtifact(t, baseStore, checkpointRef) + localStore := &corruptUntilRepairedStore{ArtifactStore: baseStore, corruptRef: checkpointRef} + _, reader, err := localStore.Open(t.Context(), checkpointRef) + require.NoError(t, err) + require.Error(t, reader.Verify()) + require.NoError(t, reader.Close()) + + peer := newFakeArtifactPeer() + peer.put(origin, KindCheckpoints, wire.Name, valid) + + srv := httptest.NewServer(peer) + t.Cleanup(srv.Close) + tr, err := newHTTPTransport(srv.URL, "", false) + require.NoError(t, err) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + assert.True(t, localStore.repaired, "corrupt local checkpoint should be re-fetched from the peer") + _, got := wireTransportArtifact(t, localStore, checkpointRef) + assert.Equal(t, valid, got) +} + +func TestHTTPTransportPostArtifactSetsPeerOrigin(t *testing.T) { + var gotOrigin, gotImportMode string + peer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotOrigin = r.Header.Get("Origin") + gotImportMode = r.Header.Get("X-Agentsview-Artifact-Import") + w.WriteHeader(http.StatusCreated) + })) + defer peer.Close() + + tr, err := newHTTPTransport(peer.URL+"/api/v1/artifacts", "", false) + require.NoError(t, err) + + body := bytes.NewReader([]byte("artifact")) + err = tr.postArtifact(context.Background(), "peer-a1b2c3", KindSegments, strings64("a"), body, int64(body.Len())) + require.NoError(t, err) + + assert.Equal(t, peer.URL, gotOrigin) + assert.Equal(t, "deferred", gotImportMode) +} + +func TestHTTPTransportFinalizesEveryPushOnce(t *testing.T) { + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + peer := newFakeArtifactPeer() + srv := httptest.NewServer(peer) + t.Cleanup(srv.Close) + tr, err := newHTTPTransport(srv.URL, "", false) + require.NoError(t, err) + + require.NoError(t, tr.Exchange(context.Background(), localStore)) + deferred, finalized := peer.batchCounts() + assert.Positive(t, deferred, "every artifact in the batch must defer import") + assert.Equal(t, 1, finalized, "one push must trigger one import finalization") + + // A retry with no missing artifacts must still finalize a batch interrupted + // after its uploads were stored but before the earlier finalize request. + require.NoError(t, tr.Exchange(context.Background(), localStore)) + deferredAfterRetry, finalizedAfterRetry := peer.batchCounts() + assert.Equal(t, deferred, deferredAfterRetry) + assert.Equal(t, 2, finalizedAfterRetry) +} + +func TestHTTPTransportUploadCardinalityUsesBoundedStorePagesAndOneFinalize(t *testing.T) { + if testing.Short() { + t.Skip("513-object real-Docbank cardinality regression") + } + type observation struct { + uploads int + finalizes int + originPages int + rawPages int + maxRequestedEntries int + maxReturnedEntries int + } + measure := func(t *testing.T, artifacts int) observation { + t.Helper() + ctx := SuppressArtifactMaintenance(t.Context()) + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + for index := range artifacts { + body := []byte("logical artifact " + strconv.Itoa(index)) + identity := identityForBytes(t, body) + ref := requireContractRef(t, contractOrigin, KindRaw, identity.SHA256) + _, err = base.Create(ctx, ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + } + observed := &pagedArtifactStoreObservations{ + ArtifactStore: base, + listPages: make(map[Kind]int), + } + peer := newFakeArtifactPeer() + server := httptest.NewServer(peer) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + require.NoError(t, transport.Exchange(ctx, observed)) + + uploads, finalizes := peer.batchCounts() + observed.mu.Lock() + defer observed.mu.Unlock() + return observation{ + uploads: uploads, + finalizes: finalizes, + originPages: observed.originPages, + rawPages: observed.listPages[KindRaw], + maxRequestedEntries: observed.maxRequestedEntries, + maxReturnedEntries: observed.maxReturnedEntries, + } + } + + small := measure(t, 1) + large := measure(t, transportPageSize+1) + + assert.Equal(t, observation{ + uploads: 1, + finalizes: 1, + originPages: 1, + rawPages: 1, + maxRequestedEntries: transportPageSize, + maxReturnedEntries: 1, + }, small) + assert.Equal(t, transportPageSize+1, large.uploads) + assert.Equal(t, 1, large.finalizes, + "crossing the store page boundary must not split the peer import batch") + assert.Equal(t, 1, large.originPages) + assert.Equal(t, 2, large.rawPages) + assert.Equal(t, transportPageSize, large.maxRequestedEntries) + assert.Equal(t, transportPageSize, large.maxReturnedEntries) +} + +func TestHTTPTransportFetchesQueuedRepairWhenBothIndexesContainName(t *testing.T) { + body := []byte("trusted repair body") + ref := requireContractRef(t, contractOrigin, KindRaw, hashHex(body)) + base, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, base.Close()) }) + created, err := base.Create(t.Context(), ref, identityForBytes(t, body), + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + local := &queuedTransportRepairStore{ArtifactStore: base, pending: created.Entry} + peer := newFakeArtifactPeer() + peer.put(ref.Origin, string(ref.Kind), ref.Name, body) + server := httptest.NewServer(peer) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + + require.NoError(t, transport.Exchange(t.Context(), local)) + + assert.Equal(t, 1, local.repairs) + assert.Empty(t, local.pending.Ref) + assert.Equal(t, body, readContractArtifact(t, local, ref)) +} + +func TestQueuedTransportRepairReturnsSpoolCleanupFailure(t *testing.T) { + body := []byte("trusted repair body") + identity := identityForBytes(t, body) + ref := requireContractRef(t, contractOrigin, KindRaw, identity.SHA256) + wire, err := ToWireRef(ref) + require.NoError(t, err) + store := &cleanupFailingRepairStore{pending: Entry{Ref: ref, Identity: identity}} + + repaired, err := repairQueuedTransportArtifact(t.Context(), store, wire, + func(consume func(io.Reader) error) error { + return consume(bytes.NewReader(body)) + }) + + assert.False(t, repaired) + require.Error(t, err) + assert.Contains(t, err.Error(), "file already closed") +} + +func TestHTTPTransportAcceptsLegacyPeerWithoutFinalize(t *testing.T) { + localStore := exportStore(t, "laptop-a1b2c3", func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + peer := newFakeArtifactPeer() + peer.supportsFinalize = false + srv := httptest.NewServer(peer) + t.Cleanup(srv.Close) + tr, err := newHTTPTransport(srv.URL, "", false) + require.NoError(t, err) + + require.NoError(t, tr.Exchange(context.Background(), localStore), + "legacy peers import each POST and return 404 for the new finalize route") +} + +func TestHTTPTransportRejectsRedirectedArtifactPost(t *testing.T) { + type requestCapture struct { + authorization string + body []byte + readErr error + } + + var destinationReached atomic.Bool + destination := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + destinationReached.Store(true) + w.WriteHeader(http.StatusCreated) + })) + t.Cleanup(destination.Close) + + captured := make(chan requestCapture, 1) + source := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + captured <- requestCapture{ + authorization: r.Header.Get("Authorization"), + body: body, + readErr: err, + } + http.Redirect(w, r, destination.URL, http.StatusTemporaryRedirect) + })) + t.Cleanup(source.Close) + + tr, err := newHTTPTransport(source.URL, "peer-secret", false) + require.NoError(t, err) + tr.client.Transport = source.Client().Transport + + body := bytes.NewReader([]byte("artifact-secret")) + err = tr.postArtifact(context.Background(), "peer-a1b2c3", KindSegments, strings64("a"), body, int64(body.Len())) + require.Error(t, err) + assert.ErrorIs(t, err, errHTTPPeer) + var got requestCapture + select { + case got = <-captured: + case <-time.After(time.Second): + require.FailNow(t, "redirect source was not reached", "timed out waiting for the artifact POST") + } + require.NoError(t, got.readErr) + assert.Equal(t, "Bearer peer-secret", got.authorization) + assert.Equal(t, []byte("artifact-secret"), got.body) + assert.False(t, destinationReached.Load()) +} + +func TestHTTPTransportUploadMemoryRemainsBoundedByBufferSize(t *testing.T) { + var requests atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + _, err := io.Copy(io.Discard, r.Body) + assert.NoError(t, err) + w.WriteHeader(http.StatusCreated) + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + transport.origin = "sender-a1b2c3" + + measure := func(size int) uint64 { + body := bytes.Repeat([]byte{'x'}, size) + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + store := openTransportStore(t, t.TempDir()) + result, err := store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + require.NoError(t, transport.postEntry(t.Context(), store, result.Entry)) + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc + } + + small := measure(1 << 20) + large := measure(24 << 20) + assert.Less(t, large, small+(4<<20), + "upload allocation growth must not scale with artifact bytes") + assert.Equal(t, int64(2), requests.Load()) +} + +func TestHTTPTransportRejectsOversizedMalformedPageWithBoundedMemory(t *testing.T) { + measure := func(size int64) uint64 { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, `{"origins":["`) + _, _ = io.CopyN(w, repeatedByteReader('x'), size) + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + _, err = transport.getOriginsPage(t.Context(), "") + runtime.ReadMemStats(&after) + require.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), "response exceeds") + return after.TotalAlloc - before.TotalAlloc + } + + small := measure(2 << 20) + large := measure(24 << 20) + assert.Less(t, large, small+(4<<20), + "malformed peer page allocation growth must remain bounded") +} + +func TestHTTPTransportExchangeMemoryRemainsBoundedByArtifactSize(t *testing.T) { + identityForSize := func(size int64) Identity { + hasher := sha256.New() + _, err := io.CopyN(hasher, repeatedByteReader('x'), size) + require.NoError(t, err) + identity, err := NewIdentity(hex.EncodeToString(hasher.Sum(nil)), size) + require.NoError(t, err) + return identity + } + measureOutbound := func(size int64) uint64 { + identity := identityForSize(size) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + local := openTransportStore(t, t.TempDir()) + _, err = local.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), io.LimitReader(repeatedByteReader('x'), size)) + require.NoError(t, err) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch { + case strings.HasSuffix(r.URL.Path, "/origins"): + _ = json.NewEncoder(w).Encode(httpOriginsPage{}) + case strings.HasSuffix(r.URL.Path, "/index"): + parts := strings.Split(strings.Trim(r.URL.Path, "/"), "/") + _ = json.NewEncoder(w).Encode(OriginArtifactIndex{Origin: parts[len(parts)-2]}) + case strings.HasSuffix(r.URL.Path, "/finalize"): + w.WriteHeader(http.StatusOK) + case r.Method == http.MethodPost: + _, copyErr := io.Copy(io.Discard, r.Body) + assert.NoError(t, copyErr) + w.WriteHeader(http.StatusCreated) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + require.NoError(t, transport.Exchange(t.Context(), local)) + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc + } + measureInbound := func(size int64) uint64 { + identity := identityForSize(size) + const origin = "peer-a1b2c3" + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch { + case strings.HasSuffix(r.URL.Path, "/origins"): + _ = json.NewEncoder(w).Encode(httpOriginsPage{Origins: []string{origin}}) + case strings.HasSuffix(r.URL.Path, "/index"): + _ = json.NewEncoder(w).Encode(OriginArtifactIndex{ + Origin: origin, + Raw: []string{identity.SHA256}, + }) + case strings.HasSuffix(r.URL.Path, "/finalize"): + w.WriteHeader(http.StatusOK) + case r.Method == http.MethodGet: + w.Header().Set("Content-Length", strconv.FormatInt(size, 10)) + _, copyErr := io.CopyN(w, repeatedByteReader('x'), size) + assert.NoError(t, copyErr) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + local := openTransportStore(t, t.TempDir()) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + require.NoError(t, transport.Exchange(t.Context(), local)) + runtime.ReadMemStats(&after) + _, err = local.Stat(t.Context(), Ref{Origin: origin, Kind: KindRaw, Name: identity.SHA256}) + require.NoError(t, err) + return after.TotalAlloc - before.TotalAlloc + } + + smallOut, smallIn := measureOutbound(1<<20), measureInbound(1<<20) + largeOut, largeIn := measureOutbound(24<<20), measureInbound(24<<20) + assert.Less(t, largeOut, smallOut+(4<<20), + "real HTTP Exchange outbound allocation growth must remain bounded") + assert.Less(t, largeIn, smallIn+(4<<20), + "real HTTP Exchange inbound allocation growth must remain bounded") +} + +func TestHTTPTransportCancellationDuringVerificationSendsNoRequest(t *testing.T) { + var requests atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + w.WriteHeader(http.StatusCreated) + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + transport.origin = "sender-a1b2c3" + + body := bytes.Repeat([]byte{'c'}, 2<<20) + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + base := openTransportStore(t, t.TempDir()) + result, err := base.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + store := &cancelDuringOpenStore{ArtifactStore: base, cancel: cancel, after: 64 << 10} + + err = transport.postEntry(ctx, store, result.Entry) + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, requests.Load(), "verification must finish before request headers are sent") +} + +func TestHTTPTransportCancellationStopsDownloadWithoutPublishing(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write(bytes.Repeat([]byte{'d'}, 64<<10)) + if flush, ok := w.(http.Flusher); ok { + flush.Flush() + } + cancel() + })) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + ref, err := NewRef("peer-a1b2c3", KindRaw, strings64("a")) + require.NoError(t, err) + wire, err := ToWireRef(ref) + require.NoError(t, err) + store := openTransportStore(t, t.TempDir()) + + err = transport.receiveArtifact(ctx, store, wire) + require.ErrorIs(t, err, context.Canceled) + _, err = store.Stat(t.Context(), ref) + require.ErrorIs(t, err, ErrArtifactNotFound) +} diff --git a/internal/artifact/transport_s3.go b/internal/artifact/transport_s3.go new file mode 100644 index 000000000..8b782f47a --- /dev/null +++ b/internal/artifact/transport_s3.go @@ -0,0 +1,751 @@ +package artifact + +import ( + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/xml" + "errors" + "fmt" + "io" + "log" + "net" + "net/http" + "net/url" + "os" + "sort" + "strconv" + "strings" + "time" +) + +const emptyPayloadSHA256 = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + +var errObjectStore = errors.New("artifact object store request failed") + +func IsObjectTarget(target string) bool { return strings.HasPrefix(target, "s3://") } + +type ObjectStoreOptions struct { + Endpoint string + Region string + AccessKeyID string + SecretAccessKey string + SessionToken string + AllowInsecureEndpoint bool + PathStyle bool +} + +func ObjectStoreOptionsFromEnv() ObjectStoreOptions { + region := os.Getenv("AGENTSVIEW_S3_REGION") + if region == "" { + region = os.Getenv("AWS_REGION") + } + if region == "" { + region = "us-east-1" + } + endpoint := os.Getenv("AGENTSVIEW_S3_ENDPOINT") + pathStyle := os.Getenv("AGENTSVIEW_S3_PATH_STYLE") == "true" + allowInsecure := false + switch strings.ToLower(strings.TrimSpace(os.Getenv("AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT"))) { + case "1", "true", "yes": + allowInsecure = true + } + if endpoint != "" { + pathStyle = true + } + return ObjectStoreOptions{ + Endpoint: endpoint, Region: region, + AccessKeyID: os.Getenv("AWS_ACCESS_KEY_ID"), + SecretAccessKey: os.Getenv("AWS_SECRET_ACCESS_KEY"), + SessionToken: os.Getenv("AWS_SESSION_TOKEN"), + AllowInsecureEndpoint: allowInsecure, + PathStyle: pathStyle, + } +} + +type s3Transport struct { + bucket string + prefix string + endpoint *url.URL + pathStyle bool + opts ObjectStoreOptions + client *http.Client +} + +func newObjectTransport(target string, opts ObjectStoreOptions) (*s3Transport, error) { + if !IsObjectTarget(target) { + return nil, errors.New("object store target must use s3://") + } + parsed, err := url.Parse(target) + if err != nil || parsed == nil { + return nil, errors.New("object store target is invalid") + } + if parsed.Host == "" { + return nil, errors.New("object store target is missing a bucket") + } + if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { + return nil, errors.New("object store target must not contain credentials, query, or fragment") + } + rest := strings.TrimPrefix(target, "s3://") + bucket, prefix, _ := strings.Cut(rest, "/") + prefix = strings.Trim(prefix, "/") + if bucket == "" { + return nil, errors.New("object store target is missing a bucket") + } + if opts.AccessKeyID == "" || opts.SecretAccessKey == "" { + return nil, errors.New("object store target requires AWS_ACCESS_KEY_ID and AWS_SECRET_ACCESS_KEY") + } + if opts.Region == "" { + opts.Region = "us-east-1" + } + var endpoint *url.URL + if opts.Endpoint == "" { + endpoint = &url.URL{Scheme: "https", Host: "s3." + opts.Region + ".amazonaws.com"} + } else { + raw := opts.Endpoint + if !strings.Contains(raw, "://") { + raw = "https://" + raw + } + u, err := url.Parse(raw) + if err != nil { + return nil, fmt.Errorf("parsing object store endpoint %q: %w", opts.Endpoint, err) + } + if u.Host == "" { + return nil, fmt.Errorf("object store endpoint is missing a host: %q", opts.Endpoint) + } + scheme := strings.ToLower(u.Scheme) + switch scheme { + case "https": + case "http": + if !opts.AllowInsecureEndpoint && !isLoopbackEndpointHost(u.Hostname()) { + return nil, fmt.Errorf("insecure S3 endpoint %q requires HTTPS or AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT", opts.Endpoint) + } + default: + return nil, fmt.Errorf("object store endpoint %q uses unsupported scheme %q; only http and https are allowed", opts.Endpoint, u.Scheme) + } + endpoint = &url.URL{Scheme: scheme, Host: u.Host} + opts.PathStyle = true + } + return &s3Transport{ + bucket: bucket, prefix: prefix, endpoint: endpoint, + pathStyle: opts.PathStyle, opts: opts, + client: &http.Client{ + Timeout: httpTransportTimeout, + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + }, nil +} + +func isLoopbackEndpointHost(host string) bool { + if strings.EqualFold(host, "localhost") { + return true + } + if address, _, ok := strings.Cut(host, "%"); ok { + host = address + } + ip := net.ParseIP(host) + return ip != nil && ip.IsLoopback() +} + +func (t *s3Transport) Prepare(ctx context.Context, _ ArtifactStore) error { + if err := requireTransportContext(ctx); err != nil { + return err + } + if _, err := t.listPage(ctx, t.prefixWithSlash(), "", "", 1); err != nil { + return fmt.Errorf("connecting to object store: %w", err) + } + return nil +} + +func (t *s3Transport) Exchange(ctx context.Context, local ArtifactStore) (retErr error) { + if err := validateTransportStore(ctx, local); err != nil { + return err + } + remoteIt := &s3OriginIterator{transport: t} + localIt := &storeOriginIterator{store: local} + defer func() { retErr = errors.Join(retErr, localIt.Close()) }() + remoteOrigin, remoteOK, err := remoteIt.Next(ctx) + if err != nil { + return fmt.Errorf("listing object store origins: %w", err) + } + localOrigin, localOK, err := localIt.Next(ctx) + if err != nil { + return fmt.Errorf("listing local artifact origins: %w", err) + } + for remoteOK || localOK { + var origin string + switch { + case !localOK || (remoteOK && remoteOrigin < localOrigin): + origin = remoteOrigin + remoteOrigin, remoteOK, err = remoteIt.Next(ctx) + case !remoteOK || localOrigin < remoteOrigin: + origin = localOrigin + localOrigin, localOK, err = localIt.Next(ctx) + default: + origin = remoteOrigin + remoteOrigin, remoteOK, err = remoteIt.Next(ctx) + if err == nil { + localOrigin, localOK, err = localIt.Next(ctx) + } + } + if err != nil { + return fmt.Errorf("advancing object store origins: %w", err) + } + if err := t.exchangeOrigin(ctx, local, origin); err != nil { + return err + } + } + return nil +} + +func (t *s3Transport) exchangeOrigin( + ctx context.Context, local ArtifactStore, origin string, +) (retErr error) { + localIt := newStoreWireIterator(local, origin) + remoteIt := &s3AllWireIterator{transport: t, origin: origin} + defer func() { retErr = errors.Join(retErr, localIt.Close()) }() + localWire, localEntry, localOK, err := localIt.Next(ctx) + if err != nil { + return err + } + remoteWire, remoteOK, err := remoteIt.Next(ctx) + if err != nil { + return err + } + for localOK || remoteOK { + if err := ctx.Err(); err != nil { + return err + } + switch { + case !localOK || (remoteOK && compareWireRefs(remoteWire, localWire) < 0): + if err := t.receiveObject(ctx, local, remoteWire); err != nil { + return err + } + remoteWire, remoteOK, err = remoteIt.Next(ctx) + if err != nil { + return err + } + case !remoteOK || compareWireRefs(localWire, remoteWire) < 0: + if err := t.putEntry(ctx, local, localEntry); err != nil { + return err + } + localWire, localEntry, localOK, err = localIt.Next(ctx) + if err != nil { + return err + } + default: + repaired, err := repairQueuedTransportArtifact(ctx, local, remoteWire, + func(consume func(io.Reader) error) error { + return t.withObject(ctx, + t.objectKey(remoteWire.Origin, string(remoteWire.Kind), remoteWire.Name), + consume) + }) + if err != nil { + return err + } + if !repaired && localWire.Kind == KindCheckpoints { + if err := t.compareCheckpoint(ctx, local, localEntry, remoteWire); err != nil { + return err + } + } + localWire, localEntry, localOK, err = localIt.Next(ctx) + if err != nil { + return err + } + remoteWire, remoteOK, err = remoteIt.Next(ctx) + if err != nil { + return err + } + } + } + return nil +} + +func (t *s3Transport) prefixWithSlash() string { + if t.prefix == "" { + return "" + } + return t.prefix + "/" +} + +func (t *s3Transport) receiveObject( + ctx context.Context, local ArtifactStore, wire WireRef, +) error { + key := t.objectKey(wire.Origin, string(wire.Kind), wire.Name) + return t.withObject(ctx, key, func(body io.Reader) error { + _, err := createTransportArtifactFromWire(ctx, local, wire, body) + if !errors.Is(err, ErrArtifactInvalid) && !errors.Is(err, ErrArtifactCorrupt) { + return err + } + log.Printf("artifact: detected corrupt artifact %s/%s/%s in object store: %v", + wire.Origin, wire.Kind, wire.Name, err) + if deleteErr := t.deleteObject(ctx, key); deleteErr != nil { + log.Printf("artifact: deleting corrupt object %s: %v", key, deleteErr) + } + return nil + }) +} + +func (t *s3Transport) compareCheckpoint( + ctx context.Context, store ArtifactStore, local Entry, wire WireRef, +) error { + return t.withObject(ctx, t.objectKey(wire.Origin, string(wire.Kind), wire.Name), func(body io.Reader) error { + return compareOrRepairCheckpoint(ctx, store, local, wire, body, + "remote", wire.Origin, local.Ref.Name) + }) +} + +func (t *s3Transport) putEntry( + ctx context.Context, local ArtifactStore, entry Entry, +) (retErr error) { + spool, size, payloadHash, err := spoolWireArtifact(ctx, local, entry) + if err != nil { + return err + } + defer func() { retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) }() + wire, err := ToWireRef(entry.Ref) + if err != nil { + return err + } + return t.putObject(ctx, t.objectKey(wire.Origin, string(wire.Kind), wire.Name), spool, size, payloadHash) +} + +func (t *s3Transport) objectKey(origin, kind, name string) string { + parts := make([]string, 0, 4) + if t.prefix != "" { + parts = append(parts, t.prefix) + } + return strings.Join(append(parts, origin, kind, name), "/") +} + +type listBucketResult struct { + XMLName xml.Name `xml:"ListBucketResult"` + IsTruncated bool `xml:"IsTruncated"` + NextContinuationToken string `xml:"NextContinuationToken"` + Contents []struct { + Key string `xml:"Key"` + } `xml:"Contents"` + CommonPrefixes []struct { + Prefix string `xml:"Prefix"` + } `xml:"CommonPrefixes"` +} + +func (t *s3Transport) listPage( + ctx context.Context, prefix, delimiter, token string, maxKeys int, +) (listBucketResult, error) { + q := url.Values{} + q.Set("list-type", "2") + if prefix != "" { + q.Set("prefix", prefix) + } + if delimiter != "" { + q.Set("delimiter", delimiter) + } + if token != "" { + q.Set("continuation-token", token) + } + if maxKeys > 0 { + q.Set("max-keys", strconv.Itoa(maxKeys)) + } + req, err := t.newRequest(ctx, http.MethodGet, "", q, nil, 0, emptyPayloadSHA256) + if err != nil { + return listBucketResult{}, err + } + resp, err := t.client.Do(req) + if err != nil { + return listBucketResult{}, err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return listBucketResult{}, t.statusError(resp) + } + body, err := readTransportPage(resp.Body) + if err != nil { + return listBucketResult{}, err + } + var result listBucketResult + if err := xml.Unmarshal(body, &result); err != nil { + return listBucketResult{}, fmt.Errorf("decoding object store listing: %w", err) + } + if maxKeys > 0 && len(result.Contents)+len(result.CommonPrefixes) > maxKeys { + return listBucketResult{}, fmt.Errorf("%w: object store page exceeds requested max-keys", ErrArtifactInvalid) + } + if result.IsTruncated && result.NextContinuationToken == "" { + return listBucketResult{}, fmt.Errorf("%w: truncated object store page is missing a continuation token", ErrArtifactInvalid) + } + for index := 1; index < len(result.Contents); index++ { + if result.Contents[index-1].Key >= result.Contents[index].Key { + return listBucketResult{}, fmt.Errorf("%w: object store keys are not strictly increasing", ErrArtifactInvalid) + } + } + for index := 1; index < len(result.CommonPrefixes); index++ { + if result.CommonPrefixes[index-1].Prefix >= result.CommonPrefixes[index].Prefix { + return listBucketResult{}, fmt.Errorf("%w: object store prefixes are not strictly increasing", ErrArtifactInvalid) + } + } + return result, nil +} + +func (t *s3Transport) withObject( + ctx context.Context, key string, consume func(io.Reader) error, +) error { + req, err := t.newRequest(ctx, http.MethodGet, key, nil, nil, 0, emptyPayloadSHA256) + if err != nil { + return err + } + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode == http.StatusNotFound { + return ErrArtifactNotFound + } + if resp.StatusCode != http.StatusOK { + return t.statusError(resp) + } + return consume(resp.Body) +} + +func (t *s3Transport) putObject( + ctx context.Context, + key string, + body io.ReadSeeker, + size int64, + payloadHash string, +) error { + req, err := t.newRequest(ctx, http.MethodPut, key, nil, body, size, payloadHash) + if err != nil { + return err + } + req.Header.Set("If-None-Match", "*") + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + switch resp.StatusCode { + case http.StatusOK, http.StatusCreated: + _, _ = io.Copy(io.Discard, resp.Body) + return nil + case http.StatusPreconditionFailed: + _, _ = io.Copy(io.Discard, resp.Body) + return t.reconcileExistingObject(ctx, key, size, payloadHash) + default: + return t.statusError(resp) + } +} + +func (t *s3Transport) reconcileExistingObject( + ctx context.Context, key string, size int64, payloadHash string, +) error { + return t.withObject(ctx, key, func(existing io.Reader) error { + hasher := sha256.New() + read, err := io.Copy(hasher, &wireContextReader{ctx: ctx, reader: existing}) + if err != nil { + return fmt.Errorf("comparing conflicting object %s: %w", key, err) + } + if read == size && hex.EncodeToString(hasher.Sum(nil)) == payloadHash { + return nil + } + return fmt.Errorf("%w: object %s already exists with different content", errObjectStore, key) + }) +} + +func (t *s3Transport) deleteObject(ctx context.Context, key string) error { + if t.endpoint.Scheme != "https" { + return fmt.Errorf("refusing to delete object through insecure S3 endpoint %q", t.endpoint.String()) + } + req, err := t.newRequest(ctx, http.MethodDelete, key, nil, nil, 0, emptyPayloadSHA256) + if err != nil { + return err + } + resp, err := t.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + switch resp.StatusCode { + case http.StatusOK, http.StatusNoContent, http.StatusNotFound: + _, _ = io.Copy(io.Discard, resp.Body) + return nil + default: + return t.statusError(resp) + } +} + +func (t *s3Transport) newRequest( + ctx context.Context, + method, key string, + query url.Values, + body io.ReadSeeker, + size int64, + payloadHash string, +) (*http.Request, error) { + host := t.endpoint.Host + rawPath := "/" + key + if t.pathStyle { + rawPath = "/" + t.bucket + if key != "" { + rawPath += "/" + key + } + } else { + host = t.bucket + "." + t.endpoint.Host + } + u := &url.URL{Scheme: t.endpoint.Scheme, Host: host, Path: rawPath, RawPath: s3EncodePath(rawPath), RawQuery: canonicalQueryString(query)} + var reader io.Reader + var readerAt io.ReaderAt + if body != nil { + readerAt = bodyReaderAt(body) + reader = io.NewSectionReader(readerAt, 0, size) + } + req, err := http.NewRequestWithContext(ctx, method, u.String(), reader) + if err != nil { + return nil, err + } + if body != nil { + req.ContentLength = size + req.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(io.NewSectionReader(readerAt, 0, size)), nil + } + req.Header.Set("Content-Type", "application/octet-stream") + } + signRequest(req, payloadHash, t.opts, time.Now()) + return req, nil +} + +func (t *s3Transport) statusError(resp *http.Response) error { + body, _ := io.ReadAll(io.LimitReader(resp.Body, httpTransportMaxErrLen)) + msg := strings.TrimSpace(string(body)) + if msg == "" { + return fmt.Errorf("%w: %s", errObjectStore, resp.Status) + } + return fmt.Errorf("%w: %s: %s", errObjectStore, resp.Status, msg) +} + +type s3OriginIterator struct { + transport *s3Transport + token string + items []string + next int + done bool + guard boundedCursorCycleGuard + lastPrefix string +} + +func (i *s3OriginIterator) Next(ctx context.Context) (string, bool, error) { + for { + if i.next < len(i.items) { + origin := i.items[i.next] + i.next++ + return origin, true, nil + } + if i.done { + return "", false, nil + } + base := i.transport.prefixWithSlash() + page, err := i.transport.listPage(ctx, base, "/", i.token, transportPageSize) + if err != nil { + return "", false, err + } + i.items = i.items[:0] + for _, common := range page.CommonPrefixes { + if !strings.HasPrefix(common.Prefix, base) || + !strings.HasSuffix(common.Prefix, "/") || + common.Prefix <= i.lastPrefix { + return "", false, fmt.Errorf("%w: object store origin prefixes are malformed or not strictly increasing", ErrArtifactInvalid) + } + i.lastPrefix = common.Prefix + rel := strings.TrimSuffix(strings.TrimPrefix(common.Prefix, base), "/") + if strings.Contains(rel, "/") || validateOriginID(rel) != nil { + return "", false, fmt.Errorf("%w: malformed artifact origin prefix %q", ErrArtifactInvalid, common.Prefix) + } + i.items = append(i.items, rel) + } + i.next = 0 + if !page.IsTruncated { + i.done = true + } else { + if i.guard.Observe(Cursor(page.NextContinuationToken)) { + return "", false, errors.New("object store continuation token cycle") + } + i.token = page.NextContinuationToken + } + } +} + +type s3WireIterator struct { + transport *s3Transport + origin string + kind Kind + token string + items []WireRef + next int + done bool + guard boundedCursorCycleGuard + lastKey string +} + +type s3AllWireIterator struct { + transport *s3Transport + origin string + kindIndex int + iterator *s3WireIterator +} + +func (i *s3AllWireIterator) Next(ctx context.Context) (WireRef, bool, error) { + for i.kindIndex < len(transportKinds) { + if i.iterator == nil { + i.iterator = &s3WireIterator{ + transport: i.transport, origin: i.origin, kind: transportKinds[i.kindIndex], + } + } + wire, ok, err := i.iterator.Next(ctx) + if err != nil || ok { + return wire, ok, err + } + i.iterator = nil + i.kindIndex++ + } + return WireRef{}, false, nil +} + +func (i *s3WireIterator) Next(ctx context.Context) (WireRef, bool, error) { + for { + if i.next < len(i.items) { + item := i.items[i.next] + i.next++ + return item, true, nil + } + if i.done { + return WireRef{}, false, nil + } + prefix := i.transport.objectKey(i.origin, string(i.kind), "") + page, err := i.transport.listPage(ctx, prefix, "", i.token, transportPageSize) + if err != nil { + return WireRef{}, false, err + } + i.items = i.items[:0] + for _, content := range page.Contents { + if !strings.HasPrefix(content.Key, prefix) || content.Key <= i.lastKey { + return WireRef{}, false, fmt.Errorf("%w: object store keys are not strictly increasing within the requested prefix", ErrArtifactInvalid) + } + i.lastKey = content.Key + name := strings.TrimPrefix(content.Key, prefix) + if strings.Contains(name, "/") || name == "" { + return WireRef{}, false, fmt.Errorf("%w: malformed artifact object key %q", ErrArtifactInvalid, content.Key) + } + ref, err := FromWireRef(i.origin, i.kind, name) + if err != nil { + return WireRef{}, false, fmt.Errorf("%w: malformed artifact object key %q: %v", ErrArtifactInvalid, content.Key, err) + } + wire, err := ToWireRef(ref) + if err != nil { + return WireRef{}, false, err + } + i.items = append(i.items, wire) + } + i.next = 0 + if !page.IsTruncated { + i.done = true + } else { + if i.guard.Observe(Cursor(page.NextContinuationToken)) { + return WireRef{}, false, errors.New("object store continuation token cycle") + } + i.token = page.NextContinuationToken + } + } +} + +func signRequest(req *http.Request, payloadSHA256Hex string, opts ObjectStoreOptions, now time.Time) { + now = now.UTC() + amzDate := now.Format("20060102T150405Z") + dateStamp := now.Format("20060102") + req.Header.Set("X-Amz-Date", amzDate) + req.Header.Set("X-Amz-Content-Sha256", payloadSHA256Hex) + if opts.SessionToken != "" { + req.Header.Set("X-Amz-Security-Token", opts.SessionToken) + } + type header struct{ name, value string } + headers := []header{{"host", req.URL.Host}, {"x-amz-content-sha256", payloadSHA256Hex}, {"x-amz-date", amzDate}} + if opts.SessionToken != "" { + headers = append(headers, header{"x-amz-security-token", opts.SessionToken}) + } + sort.Slice(headers, func(i, j int) bool { return headers[i].name < headers[j].name }) + var canonicalHeaders strings.Builder + signedNames := make([]string, 0, len(headers)) + for _, h := range headers { + canonicalHeaders.WriteString(h.name) + canonicalHeaders.WriteByte(':') + canonicalHeaders.WriteString(strings.TrimSpace(h.value)) + canonicalHeaders.WriteByte('\n') + signedNames = append(signedNames, h.name) + } + signedHeaders := strings.Join(signedNames, ";") + canonicalRequest := strings.Join([]string{req.Method, req.URL.EscapedPath(), req.URL.RawQuery, canonicalHeaders.String(), signedHeaders, payloadSHA256Hex}, "\n") + scope := dateStamp + "/" + opts.Region + "/s3/aws4_request" + stringToSign := strings.Join([]string{"AWS4-HMAC-SHA256", amzDate, scope, hashHex([]byte(canonicalRequest))}, "\n") + signingKey := sigV4SigningKey(opts.SecretAccessKey, dateStamp, opts.Region, "s3") + signature := hex.EncodeToString(hmacSHA256(signingKey, stringToSign)) + req.Header.Set("Authorization", fmt.Sprintf("AWS4-HMAC-SHA256 Credential=%s/%s, SignedHeaders=%s, Signature=%s", opts.AccessKeyID, scope, signedHeaders, signature)) +} + +func sigV4SigningKey(secret, dateStamp, region, service string) []byte { + kDate := hmacSHA256([]byte("AWS4"+secret), dateStamp) + kRegion := hmacSHA256(kDate, region) + kService := hmacSHA256(kRegion, service) + return hmacSHA256(kService, "aws4_request") +} + +func hmacSHA256(key []byte, data string) []byte { + h := hmac.New(sha256.New, key) + _, _ = h.Write([]byte(data)) + return h.Sum(nil) +} + +func canonicalQueryString(q url.Values) string { + if len(q) == 0 { + return "" + } + keys := make([]string, 0, len(q)) + for key := range q { + keys = append(keys, key) + } + sort.Strings(keys) + parts := make([]string, 0, len(q)) + for _, key := range keys { + values := append([]string(nil), q[key]...) + sort.Strings(values) + for _, value := range values { + parts = append(parts, s3URIEncode(key)+"="+s3URIEncode(value)) + } + } + return strings.Join(parts, "&") +} + +func s3EncodePath(path string) string { + segments := strings.Split(path, "/") + for index, segment := range segments { + segments[index] = s3URIEncode(segment) + } + return strings.Join(segments, "/") +} + +func s3URIEncode(value string) string { + var builder strings.Builder + for index := 0; index < len(value); index++ { + char := value[index] + if (char >= 'A' && char <= 'Z') || (char >= 'a' && char <= 'z') || (char >= '0' && char <= '9') || char == '-' || char == '_' || char == '.' || char == '~' { + builder.WriteByte(char) + continue + } + builder.WriteByte('%') + const digits = "0123456789ABCDEF" + builder.WriteByte(digits[char>>4]) + builder.WriteByte(digits[char&0xf]) + } + return builder.String() +} diff --git a/internal/artifact/transport_s3_miniotest_test.go b/internal/artifact/transport_s3_miniotest_test.go new file mode 100644 index 000000000..fd6a6c4d4 --- /dev/null +++ b/internal/artifact/transport_s3_miniotest_test.go @@ -0,0 +1,198 @@ +//go:build miniotest + +package artifact + +import ( + "context" + "fmt" + "net/http" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/testcontainers/testcontainers-go" + "github.com/testcontainers/testcontainers-go/wait" +) + +// TestS3TransportMinIORoundTrip exercises the object-store transport against a +// real MinIO server in a container, validating that the hand-rolled SigV4 +// signing is accepted by a genuine S3 implementation end to end: create bucket, +// push a producer's artifacts, list them, and pull them into a fresh store. +func TestS3TransportMinIORoundTrip(t *testing.T) { + ctx := context.Background() + endpoint, accessKey, secretKey := startMinIO(t, ctx) + + tr, err := newObjectTransport("s3://agentsview/sync", ObjectStoreOptions{ + Endpoint: endpoint, + Region: "us-east-1", + AccessKeyID: accessKey, + SecretAccessKey: secretKey, + AllowInsecureEndpoint: true, + PathStyle: true, + }) + require.NoError(t, err) + requireCreateBucket(t, ctx, tr) + + // Producer store: one exported session plus a star metadata event. + origin := "laptop-a1b2c3" + prod := testDB(t) + seedSession(t, prod, "sess-1", "alpha") + prodDir := t.TempDir() + prodRoot := filepath.Join(prodDir, "artifacts") + prodStore := openTransportStore(t, prodRoot) + _, err = ExportToStore(ctx, prod, prodStore, ExportOptions{Origin: origin, Full: true}) + require.NoError(t, err) + recorder := NewMetadataRecorder(prod, MetadataRecorderOptions{ + Origin: origin, + Store: prodStore, + }) + _, err = recorder.Append(ctx, MetadataEventInput{SessionID: "sess-1", Op: MetadataOpStar}) + require.NoError(t, err) + + want, err := listTransportArtifacts(context.Background(), prodStore, origin) + require.NoError(t, err) + require.NotEmpty(t, want.Manifests) + require.NotEmpty(t, want.Meta) + + // Push to MinIO, then confirm the bucket lists exactly the producer's set. + require.NoError(t, tr.Prepare(context.Background(), prodStore)) + require.NoError(t, tr.Exchange(ctx, prodStore)) + + remote, err := listRemoteTransportArtifacts(ctx, tr) + require.NoError(t, err) + require.Contains(t, remote, origin) + assertIndexEqual(t, want, remote[origin]) + + // A fresh consumer store pulls every artifact back. + consDir := t.TempDir() + consRoot := filepath.Join(consDir, "artifacts") + consStore := openTransportStore(t, consRoot) + require.NoError(t, tr.Exchange(ctx, consStore)) + + got, err := listTransportArtifacts(context.Background(), consStore, origin) + require.NoError(t, err) + assertIndexEqual(t, want, got) + + // Re-running is a no-op set-union: nothing new to fetch or upload. + require.NoError(t, tr.Exchange(ctx, consStore)) + got, err = listTransportArtifacts(context.Background(), consStore, origin) + require.NoError(t, err) + assertIndexEqual(t, want, got) +} + +func listRemoteTransportArtifacts( + ctx context.Context, transport *s3Transport, +) (map[string]OriginArtifactIndex, error) { + result := make(map[string]OriginArtifactIndex) + origins := &s3OriginIterator{transport: transport} + for { + origin, ok, err := origins.Next(ctx) + if err != nil { + return nil, err + } + if !ok { + return result, nil + } + index := OriginArtifactIndex{Origin: origin} + for _, kind := range transportKinds { + iterator := &s3WireIterator{transport: transport, origin: origin, kind: kind} + var names []string + for { + wire, ok, err := iterator.Next(ctx) + if err != nil { + return nil, err + } + if !ok { + break + } + names = append(names, wire.Name) + } + switch kind { + case KindCheckpoints: + index.Checkpoints = names + case KindManifests: + index.Manifests = names + case KindSegments: + index.Segments = names + case KindMeta: + index.Meta = names + case KindRaw: + index.Raw = names + } + } + result[origin] = index + } +} + +func assertIndexEqual(t *testing.T, want, got OriginArtifactIndex) { + t.Helper() + assert.ElementsMatch(t, want.Checkpoints, got.Checkpoints, "checkpoints") + assert.ElementsMatch(t, want.Manifests, got.Manifests, "manifests") + assert.ElementsMatch(t, want.Segments, got.Segments, "segments") + assert.ElementsMatch(t, want.Meta, got.Meta, "meta") + assert.ElementsMatch(t, want.Raw, got.Raw, "raw") +} + +// requireCreateBucket issues a signed CreateBucket request through the transport, +// reusing the same SigV4 path under test. An already-owned bucket is fine. +func requireCreateBucket(t *testing.T, ctx context.Context, tr *s3Transport) { + t.Helper() + var lastStatus string + for attempt := 0; attempt < 10; attempt++ { + req, err := tr.newRequest(ctx, http.MethodPut, "", nil, nil, 0, emptyPayloadSHA256) + require.NoError(t, err) + resp, err := tr.client.Do(req) + require.NoError(t, err) + lastStatus = resp.Status + status := resp.StatusCode + resp.Body.Close() + // 200 = created, 409 = already owned. 503 can still occur briefly while + // MinIO finishes initializing; retry those. + if status == http.StatusOK || status == http.StatusConflict { + return + } + if status != http.StatusServiceUnavailable { + break + } + time.Sleep(500 * time.Millisecond) + } + require.FailNowf(t, "create bucket failed", "last status: %s", lastStatus) +} + +func startMinIO(t *testing.T, ctx context.Context) (endpoint, accessKey, secretKey string) { + t.Helper() + const user, pass = "minioadmin", "minioadmin" + container, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ + ContainerRequest: testcontainers.ContainerRequest{ + Image: "minio/minio:RELEASE.2025-07-23T15-54-02Z", + ExposedPorts: []string{"9000/tcp"}, + Env: map[string]string{ + "MINIO_ROOT_USER": user, + "MINIO_ROOT_PASSWORD": pass, + }, + Cmd: []string{"server", "/data"}, + // "ready" (not "live") signals MinIO can actually serve S3 requests; + // "live" only means the process is up and races CreateBucket to 503. + WaitingFor: wait.ForHTTP("/minio/health/ready"). + WithPort("9000/tcp"). + WithStartupTimeout(2 * time.Minute), + }, + Started: true, + }) + if err != nil { + t.Skipf("could not start MinIO container (is Docker available?): %v", err) + } + t.Cleanup(func() { + stopCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + _ = container.Terminate(stopCtx) + }) + + host, err := container.Host(ctx) + require.NoError(t, err) + port, err := container.MappedPort(ctx, "9000/tcp") + require.NoError(t, err) + return fmt.Sprintf("http://%s:%s", host, port.Port()), user, pass +} diff --git a/internal/artifact/transport_s3_test.go b/internal/artifact/transport_s3_test.go new file mode 100644 index 000000000..ec63aa900 --- /dev/null +++ b/internal/artifact/transport_s3_test.go @@ -0,0 +1,1041 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/xml" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "runtime" + "sort" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +// mockS3 is an in-memory, path-style S3-compatible server backing a single +// bucket. It implements just enough of ListObjectsV2, GetObject, PutObject, and +// DeleteObject to exercise the object-store transport, and verifies that +// requests arrive signed (Authorization plus x-amz-date) without re-validating +// the signature. +type mockS3 struct { + t *testing.T + bucket string + pageSize int + + mu sync.Mutex + objects map[string][]byte + deletes int +} + +func newMockS3(t *testing.T, bucket string, pageSize int) *mockS3 { + return &mockS3{ + t: t, + bucket: bucket, + pageSize: pageSize, + objects: map[string][]byte{}, + } +} + +func (m *mockS3) put(key string, data []byte) { + m.mu.Lock() + defer m.mu.Unlock() + m.objects[key] = append([]byte(nil), data...) +} + +func (m *mockS3) has(key string) bool { + m.mu.Lock() + defer m.mu.Unlock() + _, ok := m.objects[key] + return ok +} + +func (m *mockS3) deleteCount() int { + m.mu.Lock() + defer m.mu.Unlock() + return m.deletes +} + +func (m *mockS3) ServeHTTP(w http.ResponseWriter, r *http.Request) { + assert.True(m.t, strings.HasPrefix(r.Header.Get("Authorization"), "AWS4-HMAC-SHA256"), + "request must carry a SigV4 Authorization header") + assert.NotEmpty(m.t, r.Header.Get("X-Amz-Date"), "request must carry an x-amz-date header") + + bucketPath := "/" + m.bucket + if r.Method == http.MethodGet && r.URL.Path == bucketPath && r.URL.Query().Get("list-type") == "2" { + m.list(w, r) + return + } + if !strings.HasPrefix(r.URL.Path, bucketPath+"/") { + http.Error(w, "not found", http.StatusNotFound) + return + } + key := strings.TrimPrefix(r.URL.Path, bucketPath+"/") + switch r.Method { + case http.MethodGet: + m.mu.Lock() + data, ok := m.objects[key] + m.mu.Unlock() + if !ok { + http.Error(w, "no such key", http.StatusNotFound) + return + } + _, _ = w.Write(data) + case http.MethodPut: + body := make([]byte, 0) + buf := make([]byte, 4096) + for { + n, err := r.Body.Read(buf) + body = append(body, buf[:n]...) + if err != nil { + break + } + } + // Honor the write-once conditional: reject when the key already exists. + if r.Header.Get("If-None-Match") == "*" && m.has(key) { + http.Error(w, "precondition failed", http.StatusPreconditionFailed) + return + } + m.put(key, body) + w.WriteHeader(http.StatusOK) + case http.MethodDelete: + m.mu.Lock() + m.deletes++ + delete(m.objects, key) + m.mu.Unlock() + w.WriteHeader(http.StatusNoContent) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (m *mockS3) list(w http.ResponseWriter, r *http.Request) { + prefix := r.URL.Query().Get("prefix") + delimiter := r.URL.Query().Get("delimiter") + token := r.URL.Query().Get("continuation-token") + + m.mu.Lock() + keys := make([]string, 0, len(m.objects)) + for k := range m.objects { + if strings.HasPrefix(k, prefix) { + keys = append(keys, k) + } + } + m.mu.Unlock() + sort.Strings(keys) + if delimiter != "" { + seen := make(map[string]struct{}) + prefixes := make([]string, 0, len(keys)) + for _, key := range keys { + remainder := strings.TrimPrefix(key, prefix) + component, _, found := strings.Cut(remainder, delimiter) + if !found { + continue + } + common := prefix + component + delimiter + if _, ok := seen[common]; ok { + continue + } + seen[common] = struct{}{} + prefixes = append(prefixes, common) + } + keys = prefixes + } + + start := 0 + if token != "" { + start, _ = strconv.Atoi(token) + } + pageSize := m.pageSize + if pageSize <= 0 { + pageSize = 1000 + } + end := start + pageSize + truncated := end < len(keys) + if end > len(keys) { + end = len(keys) + } + + type contentsXML struct { + Key string `xml:"Key"` + } + type resultXML struct { + XMLName xml.Name `xml:"ListBucketResult"` + IsTruncated bool `xml:"IsTruncated"` + Contents []contentsXML `xml:"Contents"` + CommonPrefixes []struct { + Prefix string `xml:"Prefix"` + } `xml:"CommonPrefixes"` + NextContinuationToken string `xml:"NextContinuationToken,omitempty"` + } + out := resultXML{IsTruncated: truncated} + for _, k := range keys[start:end] { + if delimiter == "" { + out.Contents = append(out.Contents, contentsXML{Key: k}) + } else { + out.CommonPrefixes = append(out.CommonPrefixes, struct { + Prefix string `xml:"Prefix"` + }{Prefix: k}) + } + } + if truncated { + out.NextContinuationToken = strconv.Itoa(end) + } + w.Header().Set("Content-Type", "application/xml") + require.NoError(m.t, xml.NewEncoder(w).Encode(out)) +} + +func testObjectOptions(endpoint string) ObjectStoreOptions { + return ObjectStoreOptions{ + Endpoint: endpoint, + Region: "us-east-1", + AccessKeyID: "AKIDEXAMPLE", + SecretAccessKey: "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + PathStyle: true, + } +} + +func TestS3TransportPrepareHonorsCanceledSync(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + _, _ = w.Write([]byte("")) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err = syncWithTransport(ctx, testDB(t), SyncOptions{ + DataDir: t.TempDir(), + Target: "s3://bucket/arts", + Origin: "laptop-a1b2c3", + }, transport) + + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, requests.Load(), "canceled preparation must not contact the object store") +} + +func TestS3ListPageRejectsTruncatedPageWithoutTokenAndOversizedPage(t *testing.T) { + tests := []struct { + name string + xml string + max int + }{ + { + name: "missing continuation token", + xml: `true`, + max: 1, + }, + { + name: "more keys than requested", + xml: `false` + + `ab`, + max: 1, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, tt.xml) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + + _, err = transport.listPage(t.Context(), "arts/", "", "", tt.max) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + }) + } +} + +func TestS3OriginIteratorDetectsArbitraryContinuationCycle(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + next := map[string]string{"": "a", "a": "b", "b": "c", "c": "b"}[r.URL.Query().Get("continuation-token")] + _, _ = io.WriteString(w, `true`+ + ``+next+``) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + iterator := &s3OriginIterator{transport: transport} + + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.Contains(t, err.Error(), "continuation token cycle") + assert.LessOrEqual(t, requests.Load(), int32(8)) +} + +func TestS3WireIteratorDetectsArbitraryContinuationCycle(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + next := map[string]string{"": "a", "a": "b", "b": "c", "c": "b"}[r.URL.Query().Get("continuation-token")] + _, _ = io.WriteString(w, `true`+ + ``+next+``) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + iterator := &s3WireIterator{transport: transport, origin: "peer-a1b2c3", kind: KindRaw} + + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.Contains(t, err.Error(), "continuation token cycle") + assert.LessOrEqual(t, requests.Load(), int32(8)) +} + +func TestS3OriginIteratorRejectsDuplicateAcrossPageBoundary(t *testing.T) { + const prefix = "arts/peer-a1b2c3/" + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token := r.URL.Query().Get("continuation-token") + truncated := token == "" + _, _ = fmt.Fprintf(w, `%t`+ + `%s`, truncated, prefix) + if truncated { + _, _ = io.WriteString(w, `next`) + } + _, _ = io.WriteString(w, ``) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + iterator := &s3OriginIterator{transport: transport} + + got, ok, err := iterator.Next(t.Context()) + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, "peer-a1b2c3", got) + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func TestS3WireIteratorRejectsDuplicateAndMalformedKeys(t *testing.T) { + const origin = "peer-a1b2c3" + validKey := "arts/" + origin + "/raw/" + strings64("a") + tests := []struct { + name string + secondKey string + }{ + {name: "duplicate across pages", secondKey: validKey}, + {name: "malformed nested key", secondKey: "arts/" + origin + "/raw/nested/name"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token := r.URL.Query().Get("continuation-token") + key := validKey + truncated := token == "" + if !truncated { + key = tt.secondKey + } + _, _ = fmt.Fprintf(w, `%t`+ + `%s`, truncated, key) + if truncated { + _, _ = io.WriteString(w, `next`) + } + _, _ = io.WriteString(w, ``) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + iterator := &s3WireIterator{transport: transport, origin: origin, kind: KindRaw} + + _, ok, err := iterator.Next(t.Context()) + require.NoError(t, err) + assert.True(t, ok) + _, _, err = iterator.Next(t.Context()) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + }) + } +} + +// exportStore exports one origin's sessions into a fresh artifact store. +func exportStore(t *testing.T, origin string, seed func(*db.DB)) ArtifactStore { + t.Helper() + database := testDB(t) + seed(database) + store := openTransportStore(t, filepath.Join(t.TempDir(), "artifacts")) + _, err := ExportToStore(t.Context(), database, store, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + return store +} + +func TestS3TransportPushRoundTrip(t *testing.T) { + origin := "laptop-a1b2c3" + database := testDB(t) + seedSession(t, database, "sess-1", "alpha") + + dataDir := t.TempDir() + localRoot := filepath.Join(dataDir, "artifacts") + require.NoError(t, os.MkdirAll(localRoot, 0o755)) + localStore, err := newProtocolTestStore(localRoot) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, localStore.Close()) }) + _, err = ExportToStore(t.Context(), database, localStore, ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + + // Append a metadata event so a meta artifact is part of the push. + rec := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, + Store: localStore, + Now: func() time.Time { return fixedHLCTime() }, + }) + _, err = database.StarSession("sess-1") + require.NoError(t, err) + _, err = rec.Append(context.Background(), MetadataEventInput{ + SessionID: "sess-1", + Op: MetadataOpStar, + }) + require.NoError(t, err) + + mock := newMockS3(t, "bucket", 0) + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + require.NoError(t, tr.Prepare(context.Background(), localStore)) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + idx, err := listTransportArtifacts(context.Background(), localStore, origin) + require.NoError(t, err) + items := indexItems(idx) + require.NotEmpty(t, items) + assert.NotEmpty(t, idx.Meta, "the star event should have produced a meta artifact") + for _, item := range items { + key := "arts/" + origin + "/" + item.kind + "/" + item.name + assert.True(t, mock.has(key), "expected object %q in bucket", key) + } +} + +func TestS3TransportPullRoundTrip(t *testing.T) { + origin := "desktop-d4e5f6" + + // Produce a populated store for the origin and upload it into the bucket so + // the transport must pull it down into an empty local store. + remoteStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-7", "beta") + seedSession(t, database, "sess-8", "beta") + }) + remoteIdx, err := listTransportArtifacts(context.Background(), remoteStore, origin) + require.NoError(t, err) + uploaded := indexItems(remoteIdx) + require.NotEmpty(t, uploaded) + + // Use a small page size and more than one object to exercise the + // continuation-token pagination path. + mock := newMockS3(t, "bucket", 2) + for _, item := range uploaded { + art, err := readTransportArtifact(context.Background(), remoteStore, origin, item.kind, item.name) + require.NoError(t, err) + mock.put("arts/"+origin+"/"+item.kind+"/"+item.name, art) + } + require.Greater(t, len(uploaded), 2, "need multiple pages to test pagination") + + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + + localStore := openTransportStore(t, filepath.Join(t.TempDir(), "artifacts")) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + gotIdx, err := listTransportArtifacts(context.Background(), localStore, origin) + require.NoError(t, err) + assert.ElementsMatch(t, indexItems(remoteIdx), indexItems(gotIdx)) +} + +func TestS3TransportPullRetainsCorruptRemoteArtifactOverHTTP(t *testing.T) { + origin := "desktop-d4e5f6" + remoteStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-7", "beta") + }) + remoteIdx, err := listTransportArtifacts(context.Background(), remoteStore, origin) + require.NoError(t, err) + uploaded := indexItems(remoteIdx) + mock := newMockS3(t, "bucket", 0) + for _, item := range uploaded { + art, err := readTransportArtifact(context.Background(), remoteStore, origin, item.kind, item.name) + require.NoError(t, err) + mock.put("arts/"+origin+"/"+item.kind+"/"+item.name, art) + } + corruptName := hashHex([]byte("corrupt")) + segmentExtension + mock.put("arts/"+origin+"/segments/"+corruptName, []byte("garbage")) + + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + localStore := openTransportStore(t, filepath.Join(t.TempDir(), "artifacts")) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + gotIdx, err := listTransportArtifacts(context.Background(), localStore, origin) + require.NoError(t, err) + assert.ElementsMatch(t, uploaded, indexItems(gotIdx)) + corruptRef, err := FromWireRef(origin, KindSegments, corruptName) + require.NoError(t, err) + _, err = localStore.Stat(t.Context(), corruptRef) + assert.ErrorIs(t, err, ErrArtifactNotFound) + assert.True(t, mock.has("arts/"+origin+"/segments/"+corruptName)) + assert.Zero(t, mock.deleteCount()) +} + +func TestS3TransportPullDeletesCorruptRemoteObjectSoPushHeals(t *testing.T) { + origin := "desktop-d4e5f6" + ownerStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-7", "beta") + }) + ownerIdx, err := listTransportArtifacts(context.Background(), ownerStore, origin) + require.NoError(t, err) + require.Len(t, ownerIdx.Segments, 1) + segKey := "arts/" + origin + "/segments/" + ownerIdx.Segments[0] + + mock := newMockS3(t, "bucket", 0) + for _, item := range indexItems(ownerIdx) { + art, err := readTransportArtifact(context.Background(), ownerStore, origin, item.kind, item.name) + require.NoError(t, err) + mock.put("arts/"+origin+"/"+item.kind+"/"+item.name, art) + } + // The bucket copy is corrupted in place: its name still lists, so pushes + // from valid holders would otherwise skip it forever. + mock.put(segKey, []byte("garbage")) + + srv := httptest.NewTLSServer(mock) + t.Cleanup(srv.Close) + + // An empty peer's pull fails to validate the object and deletes it, so + // the name stops masking the valid copy. + emptyStore := openTransportStore(t, filepath.Join(t.TempDir(), "artifacts")) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + tr.client.Transport = srv.Client().Transport + require.NoError(t, tr.Exchange(context.Background(), emptyStore)) + assert.False(t, mock.has(segKey), "corrupt object deleted from the bucket") + assert.Equal(t, 1, mock.deleteCount(), "expected one DELETE for the corrupt object") + + // The owner's next exchange re-uploads its valid copy, and the empty + // peer's next pull completes its store. + ownerTr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + ownerTr.client.Transport = srv.Client().Transport + require.NoError(t, ownerTr.Exchange(context.Background(), ownerStore)) + assert.True(t, mock.has(segKey), "valid copy re-uploaded") + + require.NoError(t, tr.Exchange(context.Background(), emptyStore)) + gotIdx, err := listTransportArtifacts(context.Background(), emptyStore, origin) + require.NoError(t, err) + assert.ElementsMatch(t, indexItems(ownerIdx), indexItems(gotIdx)) +} + +func TestS3TransportRejectsRedirect(t *testing.T) { + var sourceReached atomic.Bool + var destinationReached atomic.Bool + destination := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + destinationReached.Store(true) + w.Header().Set("Content-Type", "application/xml") + _, _ = w.Write([]byte("false")) + })) + t.Cleanup(destination.Close) + + source := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + sourceReached.Store(true) + http.Redirect(w, r, destination.URL, http.StatusTemporaryRedirect) + })) + t.Cleanup(source.Close) + + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(source.URL)) + require.NoError(t, err) + tr.client.Transport = source.Client().Transport + + _, err = tr.listPage(context.Background(), tr.prefixWithSlash(), "", "", 0) + require.Error(t, err) + assert.True(t, sourceReached.Load()) + assert.False(t, destinationReached.Load()) +} + +func TestS3TransportIgnoresUncatalogedCorruptLocalArtifact(t *testing.T) { + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + validIdx, err := listTransportArtifacts(context.Background(), localStore, origin) + require.NoError(t, err) + corruptName := hashHex([]byte("junk")) + segmentExtension + corruptRef, err := FromWireRef(origin, KindSegments, corruptName) + require.NoError(t, err) + _, err = localStore.Stat(t.Context(), corruptRef) + require.ErrorIs(t, err, ErrArtifactNotFound) + + mock := newMockS3(t, "bucket", 0) + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + for _, item := range indexItems(validIdx) { + assert.True(t, mock.has("arts/"+origin+"/"+item.kind+"/"+item.name), + "expected object %s/%s in bucket", item.kind, item.name) + } + assert.False(t, mock.has("arts/"+origin+"/segments/"+corruptName)) + _, err = localStore.Stat(t.Context(), corruptRef) + assert.ErrorIs(t, err, ErrArtifactNotFound, + "transport enumeration must not invent uncataloged logical artifacts") +} + +func TestS3TransportExchangeDetectsDivergentCheckpoint(t *testing.T) { + origin := "laptop-a1b2c3" + localStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + checkpointRef := onlyTransportRef(t, localStore, origin, KindCheckpoints) + wire, err := ToWireRef(checkpointRef) + require.NoError(t, err) + + divergent, err := canonicalJSON(checkpoint{ + Version: formatVersion, Origin: origin, Sequence: 1, + Sessions: map[string]string{origin + "~other": hashHex([]byte("other"))}, + }) + require.NoError(t, err) + mock := newMockS3(t, "bucket", 0) + mock.put("arts/"+origin+"/"+KindCheckpoints+"/"+wire.Name, divergent) + + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + + err = tr.Exchange(context.Background(), localStore) + require.Error(t, err) + assert.ErrorIs(t, err, errArtifactPathConflict) +} + +func TestS3TransportExchangeRepairsCorruptLocalCheckpoint(t *testing.T) { + origin := "laptop-a1b2c3" + baseStore := exportStore(t, origin, func(database *db.DB) { + seedSession(t, database, "sess-1", "alpha") + }) + checkpointRef := onlyTransportRef(t, baseStore, origin, KindCheckpoints) + wire, valid := wireTransportArtifact(t, baseStore, checkpointRef) + localStore := &corruptUntilRepairedStore{ArtifactStore: baseStore, corruptRef: checkpointRef} + + mock := newMockS3(t, "bucket", 0) + mock.put("arts/"+origin+"/"+KindCheckpoints+"/"+wire.Name, valid) + + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + require.NoError(t, tr.Exchange(context.Background(), localStore)) + + assert.True(t, localStore.repaired, "corrupt local checkpoint should be re-fetched from the bucket") + _, got := wireTransportArtifact(t, localStore, checkpointRef) + assert.Equal(t, valid, got) +} + +func TestS3TransportWriteOnceRejectsDivergentContent(t *testing.T) { + mock := newMockS3(t, "bucket", 0) + srv := httptest.NewServer(mock) + t.Cleanup(srv.Close) + tr, err := newObjectTransport("s3://bucket/arts", testObjectOptions(srv.URL)) + require.NoError(t, err) + ctx := context.Background() + key := "arts/laptop-a1b2c3/raw/deadbeef" + + // First write creates the object. + one := []byte("one") + require.NoError(t, tr.putObject(ctx, key, bytes.NewReader(one), int64(len(one)), hashHex(one))) + // An identical re-write is an accepted duplicate, not an error. + require.NoError(t, tr.putObject(ctx, key, bytes.NewReader(one), int64(len(one)), hashHex(one))) + // Divergent content at the same key is a conflict, never a silent overwrite. + two := []byte("two") + err = tr.putObject(ctx, key, bytes.NewReader(two), int64(len(two)), hashHex(two)) + require.Error(t, err) + assert.ErrorIs(t, err, errObjectStore) + // The original content is preserved. + var got []byte + err = tr.withObject(ctx, key, func(body io.Reader) error { + var readErr error + got, readErr = io.ReadAll(body) + return readErr + }) + require.NoError(t, err) + assert.Equal(t, []byte("one"), got) +} + +func TestIsObjectTarget(t *testing.T) { + tests := []struct { + name string + target string + want bool + }{ + {"s3 url", "s3://bucket/prefix", true}, + {"s3 bucket only", "s3://bucket", true}, + {"http peer", "http://example.com", false}, + {"https peer", "https://example.com", false}, + {"folder path", "/var/data/share", false}, + {"empty", "", false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, IsObjectTarget(tt.target)) + }) + } +} + +func TestNewObjectTransport(t *testing.T) { + creds := ObjectStoreOptions{ + Region: "us-east-1", + AccessKeyID: "AK", + SecretAccessKey: "SK", + } + + t.Run("missing bucket", func(t *testing.T) { + _, err := newObjectTransport("s3://", creds) + require.Error(t, err) + assert.Contains(t, err.Error(), "bucket") + }) + + t.Run("missing credentials", func(t *testing.T) { + _, err := newObjectTransport("s3://bucket/prefix", ObjectStoreOptions{Region: "us-east-1"}) + require.Error(t, err) + assert.Contains(t, err.Error(), "AWS_ACCESS_KEY_ID") + }) + + t.Run("not an object target", func(t *testing.T) { + _, err := newObjectTransport("https://example.com", creds) + require.Error(t, err) + }) + + t.Run("parses bucket and prefix", func(t *testing.T) { + tr, err := newObjectTransport("s3://bucket/some/prefix/", creds) + require.NoError(t, err) + assert.Equal(t, "bucket", tr.bucket) + assert.Equal(t, "some/prefix", tr.prefix) + assert.Equal(t, "s3.us-east-1.amazonaws.com", tr.endpoint.Host) + assert.False(t, tr.pathStyle, "real AWS defaults to virtual-host addressing") + }) + + t.Run("custom endpoint forces path style", func(t *testing.T) { + tr, err := newObjectTransport("s3://bucket", ObjectStoreOptions{ + Endpoint: "http://localhost:9000", + Region: "us-east-1", + AccessKeyID: "AK", + SecretAccessKey: "SK", + }) + require.NoError(t, err) + assert.True(t, tr.pathStyle) + assert.Equal(t, "localhost:9000", tr.endpoint.Host) + assert.Empty(t, tr.prefix) + }) + + t.Run("rejects insecure remote endpoint by default", func(t *testing.T) { + options := creds + options.Endpoint = "http://minio.lan:9000" + + _, err := newObjectTransport("s3://bucket", options) + require.Error(t, err) + assert.Contains(t, err.Error(), "insecure S3 endpoint") + assert.Contains(t, err.Error(), "AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT") + }) + + t.Run("allows opted-in insecure remote endpoint", func(t *testing.T) { + options := creds + options.Endpoint = "http://minio.lan:9000" + options.AllowInsecureEndpoint = true + + tr, err := newObjectTransport("s3://bucket", options) + require.NoError(t, err) + assert.Equal(t, "http", tr.endpoint.Scheme) + }) + + for _, endpoint := range []string{ + "http://localhost:9000", + "http://LOCALHOST:9000", + "http://127.0.0.1:9000", + "http://[::1]:9000", + } { + t.Run("allows loopback endpoint "+endpoint, func(t *testing.T) { + options := creds + options.Endpoint = endpoint + + tr, err := newObjectTransport("s3://bucket", options) + require.NoError(t, err) + assert.Equal(t, "http", tr.endpoint.Scheme) + }) + } + + t.Run("bare host defaults to HTTPS", func(t *testing.T) { + options := creds + options.Endpoint = "minio.lan:9000" + + tr, err := newObjectTransport("s3://bucket", options) + require.NoError(t, err) + assert.Equal(t, "https", tr.endpoint.Scheme) + }) + + t.Run("rejects unsupported endpoint scheme", func(t *testing.T) { + options := creds + options.Endpoint = "ftp://minio.lan" + + _, err := newObjectTransport("s3://bucket", options) + require.Error(t, err) + assert.Contains(t, err.Error(), "ftp") + }) +} + +func TestObjectStoreOptionsFromEnvAllowsInsecureEndpoint(t *testing.T) { + for _, value := range []string{"1", "true", "yes", "YES"} { + t.Run(value, func(t *testing.T) { + t.Setenv("AWS_ACCESS_KEY_ID", "AK") + t.Setenv("AWS_SECRET_ACCESS_KEY", "SK") + t.Setenv("AGENTSVIEW_S3_ENDPOINT", "http://minio.lan:9000") + t.Setenv("AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT", value) + + tr, err := newObjectTransport("s3://bucket", ObjectStoreOptionsFromEnv()) + require.NoError(t, err) + assert.Equal(t, "http", tr.endpoint.Scheme) + }) + } +} + +func TestObjectStoreOptionsFromEnvRejectsInvalidInsecureEndpointOverride(t *testing.T) { + tests := []struct { + name string + value string + }{ + {name: "empty", value: ""}, + {name: "zero", value: "0"}, + {name: "false", value: "false"}, + {name: "no", value: "no"}, + {name: "typo", value: "treu"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv("AWS_ACCESS_KEY_ID", "AK") + t.Setenv("AWS_SECRET_ACCESS_KEY", "SK") + t.Setenv("AGENTSVIEW_S3_ENDPOINT", "http://minio.lan:9000") + t.Setenv("AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT", tt.value) + + _, err := newObjectTransport("s3://bucket", ObjectStoreOptionsFromEnv()) + require.Error(t, err) + assert.Contains(t, err.Error(), "insecure S3 endpoint") + assert.Contains(t, err.Error(), "AGENTSVIEW_ALLOW_INSECURE_S3_ENDPOINT") + }) + } +} + +func TestS3TransportUploadMemoryRemainsBoundedByBufferSize(t *testing.T) { + var requests atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + _, err := io.Copy(io.Discard, r.Body) + assert.NoError(t, err) + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + + measure := func(size int) uint64 { + body := bytes.Repeat([]byte{'s'}, size) + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + store := openTransportStore(t, t.TempDir()) + result, err := store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + require.NoError(t, transport.putEntry(t.Context(), store, result.Entry)) + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc + } + + small := measure(1 << 20) + large := measure(24 << 20) + assert.Less(t, large, small+(4<<20), + "SigV4 upload allocation growth must not scale with artifact bytes") + assert.Equal(t, int64(2), requests.Load()) +} + +func TestS3TransportRejectsOversizedMalformedPageWithBoundedMemory(t *testing.T) { + measure := func(size int64) uint64 { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, ``) + _, _ = io.CopyN(w, repeatedByteReader('x'), size) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + _, err = transport.listPage(t.Context(), "arts/", "/", "", transportPageSize) + runtime.ReadMemStats(&after) + require.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), "response exceeds") + return after.TotalAlloc - before.TotalAlloc + } + + small := measure(2 << 20) + large := measure(24 << 20) + assert.Less(t, large, small+(4<<20), + "malformed object listing allocation growth must remain bounded") +} + +func TestS3TransportExchangeMemoryRemainsBoundedByArtifactSize(t *testing.T) { + identityForSize := func(size int64) Identity { + hasher := sha256.New() + _, err := io.CopyN(hasher, repeatedByteReader('x'), size) + require.NoError(t, err) + identity, err := NewIdentity(hex.EncodeToString(hasher.Sum(nil)), size) + require.NoError(t, err) + return identity + } + newServer := func(t *testing.T, identity Identity, inbound bool) *httptest.Server { + t.Helper() + const origin = "peer-a1b2c3" + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet && r.URL.Query().Get("list-type") == "2" { + prefix := r.URL.Query().Get("prefix") + switch { + case inbound && r.URL.Query().Get("delimiter") == "/": + _, _ = io.WriteString(w, `false`+ + `arts/`+origin+`/`) + case inbound && prefix == "arts/"+origin+"/raw/": + _, _ = io.WriteString(w, `false`+ + ``+prefix+identity.SHA256+``) + default: + _, _ = io.WriteString(w, `false`) + } + return + } + if inbound && r.Method == http.MethodGet { + w.Header().Set("Content-Length", strconv.FormatInt(identity.Size, 10)) + _, err := io.CopyN(w, repeatedByteReader('x'), identity.Size) + assert.NoError(t, err) + return + } + if !inbound && r.Method == http.MethodPut { + _, err := io.Copy(io.Discard, r.Body) + assert.NoError(t, err) + w.WriteHeader(http.StatusOK) + return + } + http.NotFound(w, r) + })) + } + measure := func(size int64, inbound bool) uint64 { + identity := identityForSize(size) + server := newServer(t, identity, inbound) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + local := openTransportStore(t, t.TempDir()) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + if !inbound { + _, err = local.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), io.LimitReader(repeatedByteReader('x'), size)) + require.NoError(t, err) + } + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + require.NoError(t, transport.Exchange(t.Context(), local)) + runtime.ReadMemStats(&after) + if inbound { + _, err = local.Stat(t.Context(), ref) + require.NoError(t, err) + } + return after.TotalAlloc - before.TotalAlloc + } + + smallOut, smallIn := measure(1<<20, false), measure(1<<20, true) + largeOut, largeIn := measure(24<<20, false), measure(24<<20, true) + assert.Less(t, largeOut, smallOut+(4<<20), + "real S3 Exchange outbound allocation growth must remain bounded") + assert.Less(t, largeIn, smallIn+(4<<20), + "real S3 Exchange inbound CreateFromWire allocation growth must remain bounded") +} + +func TestS3TransportCancellationDuringVerificationSendsNoRequest(t *testing.T) { + var requests atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requests.Add(1) + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + + body := bytes.Repeat([]byte{'c'}, 2<<20) + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + base := openTransportStore(t, t.TempDir()) + result, err := base.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + store := &cancelDuringOpenStore{ArtifactStore: base, cancel: cancel, after: 64 << 10} + + err = transport.putEntry(ctx, store, result.Entry) + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, requests.Load(), "verification and hashing must finish before SigV4 request headers are sent") +} + +func TestS3TransportCancellationStopsDownloadWithoutPublishing(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write(bytes.Repeat([]byte{'d'}, 64<<10)) + if flush, ok := w.(http.Flusher); ok { + flush.Flush() + } + cancel() + })) + t.Cleanup(server.Close) + transport, err := newObjectTransport("s3://bucket/arts", testObjectOptions(server.URL)) + require.NoError(t, err) + ref, err := NewRef("peer-a1b2c3", KindRaw, strings64("a")) + require.NoError(t, err) + wire, err := ToWireRef(ref) + require.NoError(t, err) + store := openTransportStore(t, t.TempDir()) + + err = transport.receiveObject(ctx, store, wire) + require.ErrorIs(t, err, context.Canceled) + _, err = store.Stat(t.Context(), ref) + require.ErrorIs(t, err, ErrArtifactNotFound) +} diff --git a/internal/artifact/transport_test.go b/internal/artifact/transport_test.go new file mode 100644 index 000000000..b5bc9c4fa --- /dev/null +++ b/internal/artifact/transport_test.go @@ -0,0 +1,461 @@ +package artifact + +import ( + "bytes" + "context" + "errors" + "io" + "net/http/httptest" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +func openTransportStore(t *testing.T, root string) ArtifactStore { + t.Helper() + require.NoError(t, os.MkdirAll(root, 0o755)) + store, err := newProtocolTestStore(root) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + return store +} + +func listTransportArtifacts( + ctx context.Context, store ArtifactStore, origin string, +) (OriginArtifactIndex, error) { + index := OriginArtifactIndex{Origin: origin} + for _, kind := range transportKinds { + var names []string + err := visitStoreKind(ctx, store, origin, kind, func(entry Entry) error { + wire, err := ToWireRef(entry.Ref) + if err != nil { + return err + } + names = append(names, wire.Name) + return nil + }) + if err != nil { + return OriginArtifactIndex{}, err + } + switch kind { + case KindSegments: + index.Segments = names + case KindRaw: + index.Raw = names + case KindManifests: + index.Manifests = names + case KindMeta: + index.Meta = names + case KindCheckpoints: + index.Checkpoints = names + } + } + return index, nil +} + +func readTransportArtifact( + ctx context.Context, + store ArtifactStore, + origin, kind, name string, +) (_ []byte, retErr error) { + ref, err := FromWireRef(origin, Kind(kind), name) + if err != nil { + return nil, err + } + entry, found, err := findStoreEntry(ctx, store, ref) + if err != nil { + return nil, err + } + if !found { + return nil, ErrArtifactNotFound + } + spool, _, _, err := spoolWireArtifact(ctx, store, entry) + if err != nil { + return nil, err + } + defer func() { retErr = errors.Join(retErr, closeAndRemoveTransportSpool(spool)) }() + return io.ReadAll(spool) +} + +func onlyTransportRef( + t *testing.T, store ArtifactStore, origin string, kind Kind, +) Ref { + t.Helper() + page, err := firstStoreEntryPage(t.Context(), store, origin, kind, 10) + require.NoError(t, err) + require.Empty(t, page.Next) + require.Len(t, page.Items, 1) + return page.Items[0].Ref +} + +func wireTransportArtifact( + t *testing.T, store ArtifactStore, ref Ref, +) (WireRef, []byte) { + t.Helper() + wire, err := ToWireRef(ref) + require.NoError(t, err) + data, err := readTransportArtifact( + t.Context(), store, ref.Origin, string(ref.Kind), wire.Name, + ) + require.NoError(t, err) + return wire, data +} + +type storeBoundaryTransport struct { + prepared ArtifactStore + exchanged ArtifactStore +} + +type transportChangeObservingStore struct { + ArtifactStore + changed []Entry + pending Entry +} + +func (s *transportChangeObservingStore) RecordTransportChanged( + ctx context.Context, entry Entry, +) error { + if err := ctx.Err(); err != nil { + return err + } + s.changed = append(s.changed, entry) + return nil +} + +func (s *transportChangeObservingStore) PendingTransportRepair( + ctx context.Context, ref Ref, +) (Entry, bool, error) { + if err := ctx.Err(); err != nil { + return Entry{}, false, err + } + return s.pending, s.pending.Ref == ref, nil +} + +func (s *transportChangeObservingStore) RepairTransportArtifact( + ctx context.Context, entry Entry, trusted io.Reader, +) error { + if entry != s.pending { + return errors.New("unexpected transport repair") + } + if err := s.Quarantine(ctx, entry.Ref, "transport parity repair"); err != nil { + return err + } + if _, err := s.Create(ctx, entry.Ref, entry.Identity, + canonicalArtifactMediaType(entry.Ref.Kind), trusted); err != nil { + return err + } + return nil +} + +func (s *transportChangeObservingStore) AcknowledgeTransportRepair( + ctx context.Context, entry Entry, +) error { + if err := ctx.Err(); err != nil { + return err + } + if entry != s.pending { + return errors.New("unexpected transport repair acknowledgement") + } + s.pending = Entry{} + return nil +} + +func TestTransportReportsChangedArtifacts(t *testing.T) { + factories := []struct { + name string + open func(*testing.T) Transport + }{ + { + name: "folder", + open: func(t *testing.T) Transport { + return openFolderTransportForTest(t, t.TempDir()) + }, + }, + { + name: "http", + open: func(t *testing.T) Transport { + peer := newFakeArtifactPeer() + server := httptest.NewServer(peer) + t.Cleanup(server.Close) + transport, err := newHTTPTransport(server.URL, "", false) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, transport.Close()) }) + return transport + }, + }, + { + name: "s3", + open: func(t *testing.T) Transport { + mock := newMockS3(t, "bucket", 2) + server := httptest.NewServer(mock) + t.Cleanup(server.Close) + transport, err := newObjectTransport( + "s3://bucket/arts", testObjectOptions(server.URL), + ) + require.NoError(t, err) + return transport + }, + }, + } + + for _, factory := range factories { + t.Run(factory.name, func(t *testing.T) { + origin := "peer-a1b2c3" + database := testDB(t) + seedSession(t, database, "session-1", "alpha") + source := openTransportStore(t, t.TempDir()) + _, err := ExportToStore(t.Context(), database, source, ExportOptions{ + Origin: origin, Full: true, + }) + require.NoError(t, err) + recorder := NewMetadataRecorder(database, MetadataRecorderOptions{ + Origin: origin, Store: source, Now: fixedHLCTime, + }) + _, err = recorder.Append(t.Context(), MetadataEventInput{ + SessionID: "session-1", Op: MetadataOpStar, + }) + require.NoError(t, err) + + checkpoint := onlyTransportRef(t, source, origin, KindCheckpoints) + metadata := onlyTransportRef(t, source, origin, KindMeta) + checkpointEntry, err := source.Stat(t.Context(), checkpoint) + require.NoError(t, err) + metadataEntry, err := source.Stat(t.Context(), metadata) + require.NoError(t, err) + + transport := factory.open(t) + require.NoError(t, transport.Exchange(t.Context(), source)) + destination := &transportChangeObservingStore{ + ArtifactStore: openTransportStore(t, t.TempDir()), + } + require.NoError(t, transport.Exchange(t.Context(), destination)) + assert.Equal(t, 1, countTransportChange(destination.changed, checkpointEntry)) + assert.Equal(t, 1, countTransportChange(destination.changed, metadataEntry)) + + require.NoError(t, transport.Exchange(t.Context(), destination)) + assert.Equal(t, 1, countTransportChange(destination.changed, checkpointEntry), + "duplicate immutable exchange must not report a change") + assert.Equal(t, 1, countTransportChange(destination.changed, metadataEntry), + "duplicate immutable exchange must not report a change") + + destination.pending = checkpointEntry + require.NoError(t, transport.Exchange(t.Context(), destination)) + assert.Equal(t, 2, countTransportChange(destination.changed, checkpointEntry), + "verified repair must report the exact changed entry") + }) + } +} + +func countTransportChange(changed []Entry, expected Entry) int { + count := 0 + for _, entry := range changed { + if entry.Ref == expected.Ref && entry.Identity == expected.Identity { + count++ + } + } + return count +} + +type closeTrackingTransport struct { + prepareErr error + closeErr error + closed bool +} + +func (t *closeTrackingTransport) Prepare(context.Context, ArtifactStore) error { + return t.prepareErr +} + +func (*closeTrackingTransport) Exchange(context.Context, ArtifactStore) error { + return nil +} + +func (t *closeTrackingTransport) Close() error { + t.closed = true + return t.closeErr +} + +type membershipStatStore struct { + ArtifactStore + created bool + statCalls int +} + +func (s *membershipStatStore) Stat(ctx context.Context, ref Ref) (Entry, error) { + s.statCalls++ + return s.ArtifactStore.Stat(ctx, ref) +} + +func (s *membershipStatStore) Create( + ctx context.Context, ref Ref, identity Identity, mediaType string, body io.Reader, +) (CreateResult, error) { + result, err := s.ArtifactStore.Create(ctx, ref, identity, mediaType, body) + if err == nil { + s.created = true + } + return result, err +} + +func (t *storeBoundaryTransport) Prepare(ctx context.Context, store ArtifactStore) error { + if err := ctx.Err(); err != nil { + return err + } + t.prepared = store + return nil +} + +func (t *storeBoundaryTransport) Exchange(ctx context.Context, store ArtifactStore) error { + if err := ctx.Err(); err != nil { + return err + } + t.exchanged = store + return nil +} + +func TestTransportBoundaryPassesTheCanonicalStore(t *testing.T) { + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + transport := &storeBoundaryTransport{} + + require.NoError(t, transport.Prepare(t.Context(), store)) + require.NoError(t, transport.Exchange(t.Context(), store)) + + assert.Same(t, store, transport.prepared) + assert.Same(t, store, transport.exchanged) + var _ Transport = transport +} + +func TestCoordinatedTransportStoreFindsExactQueuedRepair(t *testing.T) { + database := testDB(t) + store, err := newProtocolTestStore(t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, store.Close()) }) + body := []byte("repair me") + ref := requireContractRef(t, contractOrigin, KindRaw, hashHex(body)) + entry := Entry{Ref: ref, Identity: identityForBytes(t, body)} + require.NoError(t, database.EnqueueArtifactRepair(t.Context(), db.ArtifactRepair{ + Origin: ref.Origin, Kind: string(ref.Kind), Name: ref.Name, + SHA256: entry.Identity.SHA256, Size: entry.Identity.Size, + })) + coordinator := NewStoreImportCoordinator(database, store, "local-a1b2c3") + wrapped := newCoordinatedTransportStore(database, store, coordinator) + + got, ok, err := wrapped.PendingTransportRepair(t.Context(), ref) + require.NoError(t, err) + assert.True(t, ok) + assert.Equal(t, entry, got) + + _, ok, err = wrapped.PendingTransportRepair(t.Context(), requireContractRef( + t, contractOrigin, KindRaw, strings64("f"), + )) + require.NoError(t, err) + assert.False(t, ok) +} + +func TestRepositoryConstructsFolderTransportOnlyAfterIdentityValidation(t *testing.T) { + dataDir := t.TempDir() + repository, err := OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + + transport, err := repository.NewFolderTransport(filepath.Join(dataDir, "artifacts", "share")) + require.ErrorContains(t, err, "must not overlap") + assert.Nil(t, transport) + + target := t.TempDir() + transport, err = repository.NewFolderTransport(target) + require.NoError(t, err) + folder, ok := transport.(*folderTransport) + require.True(t, ok) + assert.NotNil(t, folder.root) + t.Cleanup(func() { require.NoError(t, folder.Close()) }) +} + +func TestRepositoryFolderTransportRetainsValidatedOpenedTarget(t *testing.T) { + repository, err := OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + target := t.TempDir() + transport, err := repository.NewFolderTransport(target) + require.NoError(t, err) + moved := target + "-moved" + require.NoError(t, os.Rename(target, moved)) + require.NoError(t, os.Mkdir(target, 0o755)) + + body := []byte("retained target") + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + store := openTransportStore(t, t.TempDir()) + created, err := store.Create(t.Context(), ref, identity, + canonicalArtifactMediaType(ref.Kind), bytes.NewReader(body)) + require.NoError(t, err) + wire, err := ToWireRef(created.Entry.Ref) + require.NoError(t, err) + + require.NoError(t, transport.Exchange(t.Context(), store)) + assert.FileExists(t, filepath.Join(moved, folderWirePath(wire))) + assert.NoFileExists(t, filepath.Join(target, folderWirePath(wire))) +} + +func TestRepositoryFolderTransportCloseIsIdempotentAndStopsUse(t *testing.T) { + repository, err := OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + transport, err := repository.NewFolderTransport(t.TempDir()) + require.NoError(t, err) + folder, ok := transport.(*folderTransport) + require.True(t, ok) + require.NoError(t, folder.Close()) + require.NoError(t, folder.Close()) + + err = folder.Prepare(t.Context(), openTransportStore(t, t.TempDir())) + assert.ErrorIs(t, err, os.ErrClosed) + err = folder.Exchange(t.Context(), openTransportStore(t, t.TempDir())) + assert.ErrorIs(t, err, os.ErrClosed) +} + +func TestSyncWithTransportJoinsCloseErrorOnPrepareFailure(t *testing.T) { + prepareErr := errors.New("prepare failed") + closeErr := errors.New("close failed") + transport := &closeTrackingTransport{prepareErr: prepareErr, closeErr: closeErr} + + _, err := syncWithTransport(t.Context(), testDB(t), SyncOptions{ + DataDir: t.TempDir(), + Origin: "peer-a1b2c3", + }, transport) + + assert.True(t, transport.closed) + assert.ErrorIs(t, err, prepareErr) + assert.ErrorIs(t, err, closeErr) +} + +func TestFolderPullUsesStatForRemotePointMembership(t *testing.T) { + body := []byte("remote point membership") + identity := identityForBytes(t, body) + ref, err := NewRef("peer-a1b2c3", KindRaw, identity.SHA256) + require.NoError(t, err) + wire, err := ToWireRef(ref) + require.NoError(t, err) + targetPath := t.TempDir() + require.NoError(t, os.MkdirAll( + filepath.Join(targetPath, wire.Origin, string(wire.Kind)), 0o755, + )) + require.NoError(t, os.WriteFile( + filepath.Join(targetPath, folderWirePath(wire)), body, 0o644, + )) + transport := openFolderTransportForTest(t, targetPath) + base := openTransportStore(t, t.TempDir()) + store := &membershipStatStore{ArtifactStore: base} + + require.NoError(t, transport.Exchange(t.Context(), store)) + assert.Positive(t, store.statCalls) + assert.True(t, store.created, "uncataloged wire content must still be ingested") + _, err = store.Stat(t.Context(), ref) + require.NoError(t, err) +} diff --git a/internal/artifact/twoinstance_test.go b/internal/artifact/twoinstance_test.go new file mode 100644 index 000000000..9a9fab467 --- /dev/null +++ b/internal/artifact/twoinstance_test.go @@ -0,0 +1,349 @@ +package artifact + +import ( + "context" + "database/sql" + "encoding/json" + "slices" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/db" +) + +// syncInstance models one machine in a two-instance artifact-sync harness: its +// own database, data dir, origin, and metadata recorder, all driven through the +// public folder-sync API against a shared target folder. +type syncInstance struct { + t *testing.T + db *db.DB + dataDir string + origin string + now time.Time + repo *Repository + rec *MetadataRecorder +} + +func newSyncInstance(t *testing.T, origin string) *syncInstance { + t.Helper() + in := &syncInstance{ + t: t, + db: testDB(t), + dataDir: t.TempDir(), + origin: origin, + now: fixedHLCTime(), + } + var err error + in.repo, err = OpenRepository(t.Context(), in.dataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, in.repo.Close()) }) + in.rec = NewMetadataRecorder(in.db, MetadataRecorderOptions{ + Store: in.repo.Content(), + Origin: origin, + Now: func() time.Time { return in.now }, + }) + return in +} + +// at sets the wall clock used for the instance's next local metadata edits. +func (in *syncInstance) at(ts time.Time) *syncInstance { + in.now = ts + return in +} + +// sync exports local sessions to the shared target, exchanges the union both +// ways, and imports foreign origins, mirroring `agentsview sync `. +func (in *syncInstance) sync(target string) SyncResult { + in.t.Helper() + res, err := SyncWithRepository(context.Background(), in.db, in.repo, SyncOptions{ + DataDir: in.dataDir, + Target: target, + Origin: in.origin, + Now: func() time.Time { return in.now }, + }) + require.NoError(in.t, err) + return res +} + +// syncInit mirrors `agentsview sync --init ` for baseline metadata +// initialization. +func (in *syncInstance) syncInit(target string) SyncResult { + in.t.Helper() + res, err := SyncWithRepository(context.Background(), in.db, in.repo, SyncOptions{ + DataDir: in.dataDir, + Target: target, + Origin: in.origin, + Now: func() time.Time { return in.now }, + BaselineMetadata: true, + }) + require.NoError(in.t, err) + return res +} + +// rename mirrors the rename handler: mutate the local row, then append the +// metadata event (which also records the local LWW register entry). +func (in *syncInstance) rename(localID, name string) { + in.t.Helper() + require.NoError(in.t, in.db.RenameSession(localID, &name)) + value, err := json.Marshal(struct { + DisplayName string `json:"display_name"` + }{DisplayName: name}) + require.NoError(in.t, err) + _, err = in.rec.Append(context.Background(), MetadataEventInput{ + SessionID: localID, + Op: MetadataOpRename, + Value: value, + }) + require.NoError(in.t, err) +} + +// star mirrors the star handler: star the local row, then append the event. +func (in *syncInstance) star(localID string) { + in.t.Helper() + _, err := in.db.StarSession(localID) + require.NoError(in.t, err) + _, err = in.rec.Append(context.Background(), MetadataEventInput{ + SessionID: localID, + Op: MetadataOpStar, + }) + require.NoError(in.t, err) +} + +// purge mirrors the permanent-delete handler: soft delete, permanently delete +// from trash, then append the purge event. +func (in *syncInstance) purge(localID string) { + in.t.Helper() + require.NoError(in.t, in.db.SoftDeleteSession(localID)) + _, err := in.db.DeleteSessionIfTrashed(localID) + require.NoError(in.t, err) + _, err = in.rec.Append(context.Background(), MetadataEventInput{ + SessionID: localID, + Op: MetadataOpPurge, + }) + require.NoError(in.t, err) +} + +func (in *syncInstance) displayName(t *testing.T, id string) *string { + t.Helper() + got, err := in.db.GetSession(context.Background(), id) + require.NoError(t, err) + require.NotNil(t, got) + return got.DisplayName +} + +func (in *syncInstance) requireSession(t *testing.T, id string) *db.Session { + t.Helper() + got, err := in.db.GetSession(context.Background(), id) + require.NoError(t, err) + require.NotNil(t, got, "session %s should exist", id) + return got +} + +func (in *syncInstance) isStarred(t *testing.T, id string) bool { + t.Helper() + ids, err := in.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + return slices.Contains(ids, id) +} + +func TestTwoInstanceSessionAndRenamePropagate(t *testing.T) { + target := t.TempDir() + a := newSyncInstance(t, "laptop-a1b2c3") + b := newSyncInstance(t, "desktop-d4e5f6") + gid := a.origin + "~sess-1" + + seedSession(t, a.db, "sess-1", "alpha") + a.at(fixedHLCTime()).rename("sess-1", "Renamed on A") + + a.sync(target) + res := b.sync(target) + assert.Equal(t, 1, res.ImportedSessions) + assert.GreaterOrEqual(t, res.ImportedMessages, 1) + assert.Equal(t, 1, res.ImportedMetadata) + + got, err := b.db.GetSession(context.Background(), gid) + require.NoError(t, err) + require.NotNil(t, got) + require.NotNil(t, got.DisplayName) + assert.Equal(t, "Renamed on A", *got.DisplayName) + assert.Equal(t, a.origin, got.Machine) + + // Re-syncing is idempotent: no new rows, no new conflicts. + res = b.sync(target) + assert.Zero(t, res.ImportedSessions) + assert.Zero(t, res.ImportedMetadata) +} + +func TestTwoInstanceConcurrentRenameConverges(t *testing.T) { + target := t.TempDir() + a := newSyncInstance(t, "laptop-a1b2c3") + b := newSyncInstance(t, "desktop-d4e5f6") + gid := a.origin + "~sess-1" + bLocalID := gid // the session is foreign on B, keyed by its global id + + // Share the session from A to B first so both can edit it. + seedSession(t, a.db, "sess-1", "alpha") + a.sync(target) + b.sync(target) + b.requireSession(t, bLocalID) + + // Both rename the same session concurrently. A uses the later HLC (within the + // clock drift bound so the receiver can observe it), so A wins deterministically + // on both machines. + b.at(fixedHLCTime()).rename(bLocalID, "Renamed on B") + a.at(fixedHLCTime().Add(time.Minute)).rename("sess-1", "Renamed on A") + + // Exchange until quiescent: A publishes its edit, B pulls it (A wins on B), + // then A pulls B's losing edit and records the conflict. + a.sync(target) + b.sync(target) + a.sync(target) + + require.NotNil(t, a.displayName(t, "sess-1")) + require.NotNil(t, b.displayName(t, bLocalID)) + assert.Equal(t, "Renamed on A", *a.displayName(t, "sess-1")) + assert.Equal(t, "Renamed on A", *b.displayName(t, bLocalID)) + + // Both machines independently record exactly one losing edit for the field. + assertMetadataConflictCount(t, a.db, gid, "display_name", 1) + assertMetadataConflictCount(t, b.db, gid, "display_name", 1) +} + +func TestSyncInitDoesNotLetBaselineOutrankExistingPeerMetadata(t *testing.T) { + target := t.TempDir() + a := newSyncInstance(t, "laptop-a1b2c3") + b := newSyncInstance(t, "desktop-d4e5f6") + gid := a.origin + "~sess-1" + + seedSession(t, a.db, "sess-1", "alpha") + a.sync(target) + b.sync(target) + b.requireSession(t, gid) + + peerName := "Peer newer title" + b.at(fixedHLCTime().Add(time.Minute)).rename(gid, peerName) + b.sync(target) + + staleLocalName := "Stale local title" + require.NoError(t, a.db.RenameSession("sess-1", &staleLocalName)) + + res := a.at(fixedHLCTime().Add(2 * time.Minute)).syncInit(target) + assert.Equal(t, 1, res.ImportedMetadata) + require.NotNil(t, a.displayName(t, "sess-1")) + assert.Equal(t, peerName, *a.displayName(t, "sess-1")) + + b.sync(target) + require.NotNil(t, b.displayName(t, gid)) + assert.Equal(t, peerName, *b.displayName(t, gid)) + assertMetadataConflictCount(t, a.db, gid, "display_name", 0) + assertMetadataConflictCount(t, b.db, gid, "display_name", 0) +} + +func TestSyncInitDoesNotBaselineRowsCreatedByPreBaselineImport(t *testing.T) { + target := t.TempDir() + a := newSyncInstance(t, "laptop-a1b2c3") + b := newSyncInstance(t, "desktop-d4e5f6") + gid := b.origin + "~sess-1" + + seedSession(t, b.db, "sess-1", "alpha") + b.sync(target) + + require.NoError(t, a.db.Update(func(tx *sql.Tx) error { + _, err := tx.Exec(` + CREATE TRIGGER test_imported_display_name + AFTER INSERT ON sessions + WHEN NEW.id = 'desktop-d4e5f6~sess-1' + BEGIN + UPDATE sessions + SET display_name = 'Imported stale title' + WHERE id = NEW.id; + END`) + return err + })) + + res := a.syncInit(target) + assert.Equal(t, 1, res.ImportedSessions) + assert.Equal(t, 0, res.ImportedMetadata) + require.NotNil(t, a.displayName(t, gid)) + assert.Equal(t, "Imported stale title", *a.displayName(t, gid)) + + _, ok, err := a.db.MetadataReplayStateOp(context.Background(), gid, "display_name") + require.NoError(t, err) + assert.False(t, ok, "import-created curation must not become local baseline metadata") + assert.Empty(t, globArtifacts(t, target, a.origin, "meta", "*"+metadataEventExtension)) +} + +func TestSyncReappliesMetadataReplayStateAfterManifestRefresh(t *testing.T) { + target := t.TempDir() + a := newSyncInstance(t, "laptop-a1b2c3") + b := newSyncInstance(t, "desktop-d4e5f6") + gid := b.origin + "~sess-1" + + sourceName := "Source title" + seedSession(t, b.db, "sess-1", "alpha", func(s *db.Session) { + s.SessionName = &sourceName + }) + b.sync(target) + a.sync(target) + a.requireSession(t, gid) + + localName := "Local winning title" + a.at(fixedHLCTime().Add(time.Minute)).rename(gid, localName) + require.NotNil(t, a.displayName(t, gid)) + assert.Equal(t, localName, *a.displayName(t, gid)) + + require.NoError(t, a.db.Update(func(tx *sql.Tx) error { + _, err := tx.Exec(`UPDATE sessions SET display_name = NULL WHERE id = ?`, gid) + return err + })) + refreshedSourceName := "Refreshed source title" + seedSession(t, b.db, "sess-1", "alpha", func(s *db.Session) { + s.SessionName = &refreshedSourceName + }) + b.sync(target) + + a.sync(target) + require.NotNil(t, a.displayName(t, gid)) + assert.Equal(t, localName, *a.displayName(t, gid)) +} + +func TestTwoInstanceStarAndPurgePropagate(t *testing.T) { + target := t.TempDir() + a := newSyncInstance(t, "laptop-a1b2c3") + b := newSyncInstance(t, "desktop-d4e5f6") + gid := a.origin + "~sess-1" + + seedSession(t, a.db, "sess-1", "alpha") + a.sync(target) + b.sync(target) + b.requireSession(t, gid) + + // B stars the shared session; the star converges back to A. + b.at(fixedHLCTime()).star(gid) + b.sync(target) + a.sync(target) + assert.True(t, a.isStarred(t, "sess-1")) + assert.True(t, b.isStarred(t, gid)) + + // A purges the session; the purge tombstone propagates to B and blocks + // re-import of the now-superseded manifest. The HLC stays within the drift + // bound so B can observe it on import. + a.at(fixedHLCTime().Add(time.Minute)).purge("sess-1") + a.sync(target) + b.sync(target) + + gotA, err := a.db.GetSession(context.Background(), "sess-1") + require.NoError(t, err) + assert.Nil(t, gotA) + gotB, err := b.db.GetSession(context.Background(), gid) + require.NoError(t, err) + assert.Nil(t, gotB) + + // A later sync must not resurrect the purged session on B. + b.sync(target) + gotB, err = b.db.GetSession(context.Background(), gid) + require.NoError(t, err) + assert.Nil(t, gotB) +} diff --git a/internal/artifact/wire_codec.go b/internal/artifact/wire_codec.go new file mode 100644 index 000000000..aae01dbf0 --- /dev/null +++ b/internal/artifact/wire_codec.go @@ -0,0 +1,571 @@ +package artifact + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "hash" + "io" + "os" + "strings" + "sync" + + "github.com/klauspost/compress/zstd" +) + +const ( + manifestExtension = ".json.zst" + segmentExtension = ".ndjson.zst" + + manifestDecodedLimit = int64(16 << 20) + segmentDecodedLimit = int64(64 << 20) + + // zstd.NewWriter documents an 8 MiB maximum default window. Pinning that + // size keeps existing package-written artifacts readable without accepting + // attacker-selected large decoder windows. + zstdMaxWindowSize = uint64(8 << 20) + + wireCopyBufferSize = 32 << 10 +) + +// WireCodec identifies the transport encoding for an artifact kind. +type WireCodec string + +const ( + WireCodecIdentity WireCodec = "identity" + WireCodecZstd WireCodec = "zstd" +) + +// WireRef is a validated external protocol reference. Name may include a +// transport compression extension and is never a filesystem path. +type WireRef struct { + Origin string + Kind Kind + Name string + Codec WireCodec +} + +// WireLimits bounds the encoded bytes consumed and canonical bytes emitted by +// DecodeWire. Both limits must be positive. +type WireLimits struct { + MaxEncodedBytes int64 + MaxDecodedBytes int64 +} + +var wireCopyBufferPool = sync.Pool{ + New: func() any { + buffer := new([wireCopyBufferSize]byte) + return buffer + }, +} + +// ToWireRef maps one canonical logical reference to its external wire name. +func ToWireRef(ref Ref) (WireRef, error) { + canonical, err := NewRef(ref.Origin, ref.Kind, ref.Name) + if err != nil { + return WireRef{}, err + } + wire := WireRef{ + Origin: canonical.Origin, + Kind: canonical.Kind, + Name: canonical.Name, + Codec: WireCodecIdentity, + } + switch canonical.Kind { + case KindManifests, KindSegments: + wire.Name += ".zst" + wire.Codec = WireCodecZstd + case KindCheckpoints, KindMeta, KindRaw: + default: + return WireRef{}, fmt.Errorf( + "%w: unsupported artifact kind %q", ErrArtifactInvalid, canonical.Kind, + ) + } + return wire, nil +} + +// FromWireRef maps one external protocol name to its canonical logical +// reference. It only removes a transport extension and never joins paths. +func FromWireRef(origin string, kind Kind, name string) (Ref, error) { + if err := validateOriginID(origin); err != nil { + return Ref{}, fmt.Errorf("%w: %v", ErrArtifactInvalid, err) + } + if err := validateArtifactName(name); err != nil { + return Ref{}, err + } + canonicalName := name + switch kind { + case KindManifests: + if !strings.HasSuffix(name, manifestExtension) { + return Ref{}, fmt.Errorf( + "%w: manifest wire name must end in %s", ErrArtifactInvalid, manifestExtension, + ) + } + canonicalName = strings.TrimSuffix(name, ".zst") + case KindSegments: + if !strings.HasSuffix(name, segmentExtension) { + return Ref{}, fmt.Errorf( + "%w: segment wire name must end in %s", ErrArtifactInvalid, segmentExtension, + ) + } + canonicalName = strings.TrimSuffix(name, ".zst") + case KindCheckpoints, KindMeta, KindRaw: + default: + return Ref{}, fmt.Errorf("%w: unsupported artifact kind %q", ErrArtifactInvalid, kind) + } + return NewRef(origin, kind, canonicalName) +} + +// EncodeWire streams canonical artifact bytes to their external wire encoding. +func EncodeWire(ctx context.Context, ref Ref, src io.Reader, dst io.Writer) error { + if ctx == nil { + return fmt.Errorf("%w: context is required", ErrArtifactInvalid) + } + if src == nil { + return fmt.Errorf("%w: canonical source is required", ErrArtifactInvalid) + } + if dst == nil { + return fmt.Errorf("%w: wire destination is required", ErrArtifactInvalid) + } + wireRef, err := ToWireRef(ref) + if err != nil { + return err + } + if err := ctx.Err(); err != nil { + return err + } + switch wireRef.Codec { + case WireCodecIdentity: + return copyWireStream(ctx, dst, src) + case WireCodecZstd: + return encodeWireZstd(ctx, src, dst) + default: + return fmt.Errorf("%w: unsupported wire codec %q", ErrArtifactInvalid, wireRef.Codec) + } +} + +// DecodeWire streams an external wire object to canonical bytes while +// enforcing independent encoded and decoded size ceilings. +func DecodeWire( + ctx context.Context, + wireRef WireRef, + src io.Reader, + dst io.Writer, + limits WireLimits, +) error { + if ctx == nil { + return fmt.Errorf("%w: context is required", ErrArtifactInvalid) + } + if src == nil { + return fmt.Errorf("%w: wire source is required", ErrArtifactInvalid) + } + if dst == nil { + return fmt.Errorf("%w: canonical destination is required", ErrArtifactInvalid) + } + if limits.MaxEncodedBytes <= 0 || limits.MaxDecodedBytes <= 0 { + return fmt.Errorf("%w: wire limits must be positive", ErrArtifactInvalid) + } + if _, err := canonicalRefForWire(wireRef); err != nil { + return err + } + if err := ctx.Err(); err != nil { + return err + } + + encoded := &wireLimitedReader{ + reader: &wireContextReader{ctx: ctx, reader: src}, + limit: limits.MaxEncodedBytes, + label: "encoded wire input", + } + decoded := &wireLimitedWriter{ + writer: &wireContextWriter{ctx: ctx, writer: dst}, + limit: limits.MaxDecodedBytes, + label: "decoded canonical output", + } + var err error + switch wireRef.Codec { + case WireCodecIdentity: + err = copyWireStream(ctx, decoded, encoded) + case WireCodecZstd: + err = decodeWireZstd(ctx, encoded, decoded, limits.MaxDecodedBytes) + default: + err = fmt.Errorf("unsupported wire codec %q", wireRef.Codec) + } + if err == nil { + err = encoded.requireEOF() + } + if err == nil { + return nil + } + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + return fmt.Errorf("%w: decoding %s: %w", ErrArtifactCorrupt, wireRef.Name, err) +} + +// PeerWireLimits returns the protocol size limits for one peer artifact kind. +func PeerWireLimits(kind Kind) WireLimits { + return transportWireLimits(kind) +} + +// CanonicalWireSpool owns one private temporary file containing a decoded, +// identity-checked, and semantically validated artifact. Call Close when the +// canonical bytes have been created or repaired in a store. +type CanonicalWireSpool struct { + ref Ref + identity Identity + file *os.File + path string +} + +// Ref returns the canonical logical reference decoded into the spool. +func (s *CanonicalWireSpool) Ref() Ref { return s.ref } + +// Identity returns the SHA-256 and size of the canonical decoded bytes. +func (s *CanonicalWireSpool) Identity() Identity { return s.identity } + +// Rewind positions the canonical stream at its beginning for one store +// operation. Callers must serialize uses of a spool. +func (s *CanonicalWireSpool) Rewind() (io.Reader, error) { + if s == nil || s.file == nil { + return nil, fmt.Errorf("%w: canonical artifact spool is closed", ErrArtifactInvalid) + } + if _, err := s.file.Seek(0, io.SeekStart); err != nil { + return nil, fmt.Errorf("rewinding canonical artifact spool: %w", err) + } + return s.file, nil +} + +// Create performs immutable creation from the already validated canonical +// bytes without decoding the peer stream a second time. +func (s *CanonicalWireSpool) Create( + ctx context.Context, store ArtifactStore, +) (CreateResult, error) { + if store == nil { + return CreateResult{}, fmt.Errorf("%w: artifact store is required", ErrArtifactInvalid) + } + reader, err := s.Rewind() + if err != nil { + return CreateResult{}, err + } + return store.Create(ctx, s.ref, s.identity, + canonicalArtifactMediaType(s.ref.Kind), reader) +} + +// Close closes and removes the private canonical spool. +func (s *CanonicalWireSpool) Close() error { + if s == nil || s.file == nil { + return nil + } + file := s.file + s.file = nil + return errors.Join(file.Close(), removeCanonicalWireSpool(s.path)) +} + +func removeCanonicalWireSpool(path string) error { + err := os.Remove(path) + if err != nil && !errors.Is(err, os.ErrNotExist) { + return fmt.Errorf("removing canonical artifact spool: %w", err) + } + return nil +} + +// DecodeWireToCanonicalSpool decodes one inbound wire object exactly once, +// validates its protocol identity and bounded semantics, and returns a private +// canonical stream suitable for immutable creation or an authorized repair. +func DecodeWireToCanonicalSpool( + ctx context.Context, + wireRef WireRef, + src io.Reader, + limits WireLimits, +) (_ *CanonicalWireSpool, retErr error) { + ref, err := canonicalRefForWire(wireRef) + if err != nil { + return nil, err + } + file, err := os.CreateTemp("", "agentsview-artifact-spool-*") + if err != nil { + return nil, fmt.Errorf("creating canonical artifact spool: %w", err) + } + spool := &CanonicalWireSpool{ref: ref, file: file, path: file.Name()} + defer func() { + if retErr != nil { + retErr = errors.Join(retErr, spool.Close()) + } + }() + if err := file.Chmod(0o600); err != nil { + return nil, fmt.Errorf("securing canonical artifact spool: %w", err) + } + + hasher := sha256.New() + canonical := &wireSpoolWriter{file: file, hash: hasher} + if err := DecodeWire(ctx, wireRef, src, canonical, limits); err != nil { + return nil, err + } + identity, err := NewIdentity(hex.EncodeToString(hasher.Sum(nil)), canonical.size) + if err != nil { + return nil, err + } + if err := validateRefIdentity(ref, identity); err != nil { + return nil, err + } + if err := ctx.Err(); err != nil { + return nil, err + } + if err := validateCanonicalArtifactSpool(ctx, ref, file, limits.MaxDecodedBytes); err != nil { + return nil, err + } + spool.identity = identity + return spool, nil +} + +// CreateFromWire decodes one inbound wire object to a private canonical spool, +// validates its content identity and bounded semantics, and creates it in the +// store through CreateImmutable. +func CreateFromWire( + ctx context.Context, + store ArtifactStore, + wireRef WireRef, + src io.Reader, + limits WireLimits, +) (result CreateResult, retErr error) { + if store == nil { + return CreateResult{}, fmt.Errorf("%w: artifact store is required", ErrArtifactInvalid) + } + spool, err := DecodeWireToCanonicalSpool(ctx, wireRef, src, limits) + if err != nil { + return CreateResult{}, err + } + defer func() { + retErr = errors.Join(retErr, spool.Close()) + }() + return spool.Create(ctx, store) +} + +func validateCanonicalArtifactSpool( + ctx context.Context, ref Ref, spool *os.File, limit int64, +) error { + if ref.Kind == KindRaw { + return ctx.Err() + } + if _, err := spool.Seek(0, io.SeekStart); err != nil { + return err + } + data, err := io.ReadAll(io.LimitReader( + &wireContextReader{ctx: ctx, reader: spool}, limit+1, + )) + if err != nil { + return err + } + if int64(len(data)) > limit { + return fmt.Errorf("%w: decoded artifact exceeds limit", ErrArtifactInvalid) + } + switch ref.Kind { + case KindCheckpoints: + return validateCheckpointData(data, ref.Origin, ref.Name) + case KindManifests: + return validateCanonicalManifestArtifactData(data, ref.Origin) + case KindSegments: + return validateCanonicalSegmentArtifactData(data) + case KindMeta: + _, hash, err := normalizeMetadataName(ref.Name) + if err != nil { + return err + } + return validateMetadataArtifactData(data, ref.Origin, ref.Name, hash) + default: + return fmt.Errorf("%w: unsupported artifact kind %q", ErrArtifactInvalid, ref.Kind) + } +} + +func canonicalRefForWire(wireRef WireRef) (Ref, error) { + ref, err := FromWireRef(wireRef.Origin, wireRef.Kind, wireRef.Name) + if err != nil { + return Ref{}, err + } + want, err := ToWireRef(ref) + if err != nil { + return Ref{}, err + } + if wireRef != want { + return Ref{}, fmt.Errorf( + "%w: wire codec %q does not match %s artifacts", + ErrArtifactInvalid, wireRef.Codec, wireRef.Kind, + ) + } + return ref, nil +} + +func encodeWireZstd(ctx context.Context, src io.Reader, dst io.Writer) error { + writer, err := zstd.NewWriter( + &wireContextWriter{ctx: ctx, writer: dst}, + zstd.WithEncoderConcurrency(1), + zstd.WithWindowSize(int(zstdMaxWindowSize)), + zstd.WithEncoderCRC(true), + ) + if err != nil { + return err + } + copyErr := copyWireStream(ctx, writer, src) + closeErr := writer.Close() + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + if copyErr != nil { + return copyErr + } + return closeErr +} + +func decodeWireZstd( + ctx context.Context, + src io.Reader, + dst io.Writer, + maxDecoded int64, +) error { + maxMemory := max(uint64(maxDecoded), zstdMaxWindowSize) + reader, err := zstd.NewReader( + src, + zstd.WithDecoderConcurrency(1), + zstd.WithDecoderLowmem(true), + zstd.WithDecoderMaxMemory(maxMemory), + zstd.WithDecoderMaxWindow(zstdMaxWindowSize), + ) + if err != nil { + return err + } + defer reader.Close() + return copyWireStream(ctx, dst, struct{ io.Reader }{Reader: reader}) +} + +func copyWireStream(ctx context.Context, dst io.Writer, src io.Reader) error { + if err := ctx.Err(); err != nil { + return err + } + pooled := wireCopyBufferPool.Get().(*[wireCopyBufferSize]byte) + defer wireCopyBufferPool.Put(pooled) + _, err := io.CopyBuffer( + &wireContextWriter{ctx: ctx, writer: dst}, + &wireContextReader{ctx: ctx, reader: src}, + pooled[:], + ) + if err != nil { + return err + } + // Cancellation is cooperative at I/O boundaries: the wrappers check before + // every underlying Read and Write, and this final check catches cancellation + // during the last successful call. An arbitrary Reader or Writer that is + // already blocked must provide its own context-aware unblocking mechanism; + // this codec does not spawn per-I/O goroutines that could leak behind it. + return ctx.Err() +} + +type wireContextReader struct { + ctx context.Context + reader io.Reader +} + +func (r *wireContextReader) Read(p []byte) (int, error) { + if err := r.ctx.Err(); err != nil { + return 0, err + } + return r.reader.Read(p) +} + +type wireContextWriter struct { + ctx context.Context + writer io.Writer +} + +func (w *wireContextWriter) Write(p []byte) (int, error) { + if err := w.ctx.Err(); err != nil { + return 0, err + } + return w.writer.Write(p) +} + +type wireLimitedReader struct { + reader io.Reader + limit int64 + read int64 + label string +} + +func (r *wireLimitedReader) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + remaining := r.limit - r.read + if remaining <= 0 { + var probe [1]byte + n, err := r.reader.Read(probe[:]) + if n > 0 { + return 0, fmt.Errorf("%s exceeds %d-byte limit", r.label, r.limit) + } + return 0, err + } + if int64(len(p)) > remaining { + p = p[:int(remaining)] + } + n, err := r.reader.Read(p) + r.read += int64(n) + return n, err +} + +func (r *wireLimitedReader) requireEOF() error { + var probe [1]byte + n, err := r.Read(probe[:]) + if n > 0 { + return fmt.Errorf("%s exceeds %d-byte limit", r.label, r.limit) + } + if err == nil { + return io.ErrNoProgress + } + if errors.Is(err, io.EOF) { + return nil + } + return err +} + +type wireLimitedWriter struct { + writer io.Writer + limit int64 + wrote int64 + label string +} + +func (w *wireLimitedWriter) Write(p []byte) (int, error) { + remaining := w.limit - w.wrote + if remaining <= 0 && len(p) > 0 { + return 0, fmt.Errorf("%s exceeds %d-byte limit", w.label, w.limit) + } + if int64(len(p)) > remaining { + p = p[:int(remaining)] + n, err := w.writer.Write(p) + w.wrote += int64(n) + if err != nil { + return n, err + } + return n, fmt.Errorf("%s exceeds %d-byte limit", w.label, w.limit) + } + n, err := w.writer.Write(p) + w.wrote += int64(n) + return n, err +} + +type wireSpoolWriter struct { + file *os.File + hash hash.Hash + size int64 +} + +func (w *wireSpoolWriter) Write(p []byte) (int, error) { + n, err := w.file.Write(p) + if n > 0 { + _, _ = w.hash.Write(p[:n]) + w.size += int64(n) + } + return n, err +} diff --git a/internal/artifact/wire_codec_test.go b/internal/artifact/wire_codec_test.go new file mode 100644 index 000000000..b40799cb9 --- /dev/null +++ b/internal/artifact/wire_codec_test.go @@ -0,0 +1,940 @@ +package artifact + +import ( + "bytes" + "context" + "crypto/sha256" + "errors" + "fmt" + "io" + "os" + "strings" + "testing" + + "github.com/klauspost/compress/zstd" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestWireRefMappings(t *testing.T) { + hash := strings.Repeat("a", 64) + tests := []struct { + name string + kind Kind + canonicalName string + wireName string + codec WireCodec + }{ + { + name: "checkpoint", kind: KindCheckpoints, + canonicalName: "cp-0000000001.json", + wireName: "cp-0000000001.json", + codec: WireCodecIdentity, + }, + { + name: "manifest", kind: KindManifests, + canonicalName: hash + ".json", + wireName: hash + ".json.zst", + codec: WireCodecZstd, + }, + { + name: "segment", kind: KindSegments, + canonicalName: hash + ".ndjson", + wireName: hash + ".ndjson.zst", + codec: WireCodecZstd, + }, + { + name: "metadata", kind: KindMeta, + canonicalName: "20260721T010203.000000000Z-0-" + hash + ".json", + wireName: "20260721T010203.000000000Z-0-" + hash + ".json", + codec: WireCodecIdentity, + }, + { + name: "raw", kind: KindRaw, + canonicalName: hash, + wireName: hash, + codec: WireCodecIdentity, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + canonical, err := NewRef(contractOrigin, tt.kind, tt.canonicalName) + require.NoError(t, err) + + wire, err := ToWireRef(canonical) + require.NoError(t, err) + assert.Equal(t, WireRef{ + Origin: contractOrigin, + Kind: tt.kind, + Name: tt.wireName, + Codec: tt.codec, + }, wire) + + roundTrip, err := FromWireRef(contractOrigin, tt.kind, tt.wireName) + require.NoError(t, err) + assert.Equal(t, canonical, roundTrip) + }) + } +} + +func TestWireRefRejectsInvalidOrNonWireNames(t *testing.T) { + hash := strings.Repeat("a", 64) + tests := []struct { + name string + origin string + kind Kind + wire string + }{ + {name: "invalid origin", origin: "../peer", kind: KindRaw, wire: hash}, + {name: "unknown kind", origin: contractOrigin, kind: "future", wire: hash}, + {name: "manifest missing zstd extension", origin: contractOrigin, kind: KindManifests, wire: hash + ".json"}, + {name: "segment missing zstd extension", origin: contractOrigin, kind: KindSegments, wire: hash + ".ndjson"}, + {name: "raw gains extension", origin: contractOrigin, kind: KindRaw, wire: hash + ".zst"}, + {name: "path separator", origin: contractOrigin, kind: KindManifests, wire: "../" + hash + ".json.zst"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := FromWireRef(tt.origin, tt.kind, tt.wire) + assert.ErrorIs(t, err, ErrArtifactInvalid) + }) + } + + _, err := ToWireRef(Ref{Origin: contractOrigin, Kind: KindManifests, Name: hash + ".json.zst"}) + assert.ErrorIs(t, err, ErrArtifactInvalid) +} + +func TestWireCodecRoundTripsIdentityAndZstd(t *testing.T) { + hash := strings.Repeat("a", 64) + tests := []struct { + name string + ref Ref + body []byte + }{ + { + name: "checkpoint", + ref: Ref{Origin: contractOrigin, Kind: KindCheckpoints, Name: "cp-0000000001.json"}, + body: []byte("canonical checkpoint bytes\n"), + }, + { + name: "manifest", + ref: Ref{Origin: contractOrigin, Kind: KindManifests, Name: hash + ".json"}, + body: []byte("{\"canonical\":\"manifest\"}\n"), + }, + { + name: "segment", + ref: Ref{Origin: contractOrigin, Kind: KindSegments, Name: hash + ".ndjson"}, + body: []byte("{\"canonical\":\"segment\"}\n"), + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + wireRef, err := ToWireRef(tt.ref) + require.NoError(t, err) + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), tt.ref, bytes.NewReader(tt.body), &encoded)) + + var decoded bytes.Buffer + err = DecodeWire(t.Context(), wireRef, bytes.NewReader(encoded.Bytes()), &decoded, WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(tt.body)), + }) + require.NoError(t, err) + assert.Equal(t, tt.body, decoded.Bytes()) + }) + } +} + +func TestWireZstdEncodingIsDeterministic(t *testing.T) { + ref := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("a", 64) + ".ndjson", + } + body := bytes.Repeat([]byte("deterministic canonical record\n"), 1024) + var first, second bytes.Buffer + + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &first)) + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &second)) + assert.Equal(t, first.Bytes(), second.Bytes()) + assert.NotEqual(t, body, first.Bytes()) +} + +func TestWireDecodeEnforcesEncodedAndDecodedLimits(t *testing.T) { + rawRef := Ref{Origin: contractOrigin, Kind: KindRaw, Name: strings.Repeat("a", 64)} + rawWire, err := ToWireRef(rawRef) + require.NoError(t, err) + + t.Run("encoded exact ceiling", func(t *testing.T) { + var decoded bytes.Buffer + err := DecodeWire(t.Context(), rawWire, strings.NewReader("12345"), &decoded, WireLimits{ + MaxEncodedBytes: 5, + MaxDecodedBytes: 5, + }) + require.NoError(t, err) + assert.Equal(t, "12345", decoded.String()) + }) + + t.Run("encoded ceiling exceeded", func(t *testing.T) { + var decoded bytes.Buffer + err := DecodeWire(t.Context(), rawWire, strings.NewReader("123456"), &decoded, WireLimits{ + MaxEncodedBytes: 5, + MaxDecodedBytes: 100, + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactCorrupt) + assert.Contains(t, err.Error(), "encoded wire input exceeds 5-byte limit") + }) + + t.Run("decoded ceiling exceeded", func(t *testing.T) { + ref := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("b", 64) + ".ndjson", + } + wireRef, err := ToWireRef(ref) + require.NoError(t, err) + body := bytes.Repeat([]byte("x"), 1024) + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &encoded)) + + var decoded bytes.Buffer + err = DecodeWire(t.Context(), wireRef, bytes.NewReader(encoded.Bytes()), &decoded, WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body) - 1), + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactCorrupt) + assert.Contains(t, err.Error(), "decoded canonical output exceeds 1023-byte limit") + }) +} + +func TestWireDecodeRejectsLargeZstdWindow(t *testing.T) { + body := bytes.Repeat([]byte("windowed-record\n"), 700_000) + var encoded bytes.Buffer + enc, err := zstd.NewWriter( + &encoded, + zstd.WithWindowSize(16<<20), + zstd.WithEncoderConcurrency(1), + ) + require.NoError(t, err) + _, err = enc.Write(body) + require.NoError(t, err) + require.NoError(t, enc.Close()) + wireRef := requireWireRef(t, Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("a", 64) + ".ndjson", + }) + + err = DecodeWire(t.Context(), wireRef, bytes.NewReader(encoded.Bytes()), io.Discard, WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body)), + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactCorrupt) + assert.Contains(t, err.Error(), "window size exceeded") +} + +func TestWireDecodeRejectsCorruptAndTruncatedZstd(t *testing.T) { + ref := Ref{ + Origin: contractOrigin, + Kind: KindManifests, + Name: strings.Repeat("a", 64) + ".json", + } + wireRef := requireWireRef(t, ref) + body := bytes.Repeat([]byte("canonical manifest bytes\n"), 128) + var valid bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &valid)) + require.Greater(t, valid.Len(), 4) + + tests := []struct { + name string + data []byte + }{ + {name: "corrupt", data: []byte("this is not a zstd frame")}, + {name: "truncated", data: append([]byte(nil), valid.Bytes()[:valid.Len()-1]...)}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := DecodeWire(t.Context(), wireRef, bytes.NewReader(tt.data), io.Discard, WireLimits{ + MaxEncodedBytes: int64(len(tt.data)), + MaxDecodedBytes: int64(len(body)), + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactCorrupt) + }) + } +} + +func TestWireDecodeRejectsNonPositiveLimits(t *testing.T) { + wireRef := requireWireRef(t, Ref{ + Origin: contractOrigin, + Kind: KindRaw, + Name: strings.Repeat("a", 64), + }) + tests := []struct { + name string + limits WireLimits + }{ + { + name: "zero encoded", + limits: WireLimits{ + MaxEncodedBytes: 0, + MaxDecodedBytes: 1, + }, + }, + { + name: "negative encoded", + limits: WireLimits{ + MaxEncodedBytes: -1, + MaxDecodedBytes: 1, + }, + }, + { + name: "zero decoded", + limits: WireLimits{ + MaxEncodedBytes: 1, + MaxDecodedBytes: 0, + }, + }, + { + name: "negative decoded", + limits: WireLimits{ + MaxEncodedBytes: 1, + MaxDecodedBytes: -1, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := DecodeWire( + t.Context(), wireRef, strings.NewReader("x"), io.Discard, tt.limits, + ) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.Contains(t, err.Error(), "wire limits must be positive") + }) + } +} + +func TestWireCodecHonorsCancellation(t *testing.T) { + ref := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("a", 64) + ".ndjson", + } + + t.Run("encode identity", func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + src := &cancelAfterReader{ + cancel: cancel, + remaining: 1 << 20, + perRead: 1024, + } + rawRef := Ref{ + Origin: contractOrigin, + Kind: KindRaw, + Name: strings.Repeat("b", 64), + } + err := EncodeWire(ctx, rawRef, src, io.Discard) + assert.ErrorIs(t, err, context.Canceled) + assert.Less(t, src.read, int64(1<<20)) + }) + + t.Run("encode identity canceled by final read", func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + src := &cancelOnFinalRead{ + cancel: cancel, + body: []byte("complete but canceled input"), + } + rawRef := Ref{ + Origin: contractOrigin, + Kind: KindRaw, + Name: strings.Repeat("c", 64), + } + err := EncodeWire(ctx, rawRef, src, io.Discard) + assert.ErrorIs(t, err, context.Canceled) + }) + + t.Run("encode", func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + src := &cancelAfterReader{ + cancel: cancel, + remaining: 1 << 20, + perRead: 1024, + } + err := EncodeWire(ctx, ref, src, io.Discard) + assert.ErrorIs(t, err, context.Canceled) + assert.Less(t, src.read, int64(1<<20)) + }) + + t.Run("decode", func(t *testing.T) { + body := bytes.Repeat([]byte("incompressible-ish-0123456789abcdef\n"), 4096) + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &encoded)) + ctx, cancel := context.WithCancel(t.Context()) + src := &cancelAfterReader{ + reader: bytes.NewReader(encoded.Bytes()), + cancel: cancel, + perRead: 1, + remaining: int64(encoded.Len()), + } + err := DecodeWire(ctx, requireWireRef(t, ref), src, io.Discard, WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body)), + }) + assert.ErrorIs(t, err, context.Canceled) + assert.Less(t, src.read, int64(encoded.Len())) + }) +} + +func TestWireCodecDetectsCancellationDuringFinalSuccessfulWrite(t *testing.T) { + body := []byte("one final successful destination write") + rawRef := Ref{ + Origin: contractOrigin, + Kind: KindRaw, + Name: strings.Repeat("a", 64), + } + zstdRef := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("b", 64) + ".ndjson", + } + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), zstdRef, bytes.NewReader(body), &encoded)) + zstdWire := requireWireRef(t, zstdRef) + rawWire := requireWireRef(t, rawRef) + + tests := []struct { + name string + run func(context.Context, io.Writer) error + }{ + { + name: "encode identity", + run: func(ctx context.Context, dst io.Writer) error { + return EncodeWire(ctx, rawRef, &singleReadEOF{body: body}, dst) + }, + }, + { + name: "encode zstd", + run: func(ctx context.Context, dst io.Writer) error { + return EncodeWire(ctx, zstdRef, &singleReadEOF{body: body}, dst) + }, + }, + { + name: "decode identity", + run: func(ctx context.Context, dst io.Writer) error { + return DecodeWire(ctx, rawWire, &singleReadEOF{body: body}, dst, WireLimits{ + MaxEncodedBytes: int64(len(body)), + MaxDecodedBytes: int64(len(body)), + }) + }, + }, + { + name: "decode zstd", + run: func(ctx context.Context, dst io.Writer) error { + return DecodeWire(ctx, zstdWire, bytes.NewReader(encoded.Bytes()), dst, WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body)), + }) + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + count := &cancelOnSuccessfulWrite{writer: io.Discard} + require.NoError(t, tt.run(t.Context(), count)) + require.Positive(t, count.writes) + + ctx, cancel := context.WithCancel(t.Context()) + dst := &cancelOnSuccessfulWrite{ + writer: io.Discard, + cancel: cancel, + cancelOn: count.writes, + } + err := tt.run(ctx, dst) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, count.writes, dst.writes, + "cancellation must occur during the operation's final full write") + }) + } +} + +func TestWireCodecStreamsMultiMegabyteArtifactsWithBoundedBuffers(t *testing.T) { + const size = int64(6 << 20) + ref := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("a", 64) + ".ndjson", + } + canonicalHash := sha256.New() + src := &boundedPatternReader{ + remaining: size, + maxRequest: 128 << 10, + } + encoded, err := os.CreateTemp(t.TempDir(), "wire-*.zst") + require.NoError(t, err) + t.Cleanup(func() { _ = encoded.Close() }) + + err = EncodeWire(t.Context(), ref, io.TeeReader(src, canonicalHash), encoded) + require.NoError(t, err) + assert.LessOrEqual(t, src.largestRequest, 128<<10) + encodedSize, err := encoded.Seek(0, io.SeekCurrent) + require.NoError(t, err) + require.Positive(t, encodedSize) + _, err = encoded.Seek(0, io.SeekStart) + require.NoError(t, err) + + decodedHash := sha256.New() + dst := &boundedWriteObserver{writer: decodedHash, maxWrite: 128 << 10} + err = DecodeWire(t.Context(), requireWireRef(t, ref), encoded, dst, WireLimits{ + MaxEncodedBytes: encodedSize, + MaxDecodedBytes: size, + }) + require.NoError(t, err) + assert.Equal(t, size, dst.written) + assert.LessOrEqual(t, dst.largestWrite, 128<<10) + assert.Equal(t, canonicalHash.Sum(nil), decodedHash.Sum(nil)) +} + +func TestWireCodecAllocatedBytesStayBoundedAsArtifactsGrow(t *testing.T) { + const ( + smallSize = int64(32 << 10) + largeSize = int64(12 << 20) + // Identity has no codec window, so growth above this allowance is an + // artifact-sized buffer rather than fixed streaming overhead. + identityMaxAllocationGrowth = int64(1 << 20) + // The zstd encoder may grow to two fixed 8 MiB working windows plus + // bookkeeping. Deterministically incompressible input makes either a + // canonical or compressed-wire artifact buffer cross this allowance. + zstdMaxAllocationGrowth = int64(20 << 20) + ) + rawRef := Ref{ + Origin: contractOrigin, + Kind: KindRaw, + Name: strings.Repeat("a", 64), + } + zstdRef := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("b", 64) + ".ndjson", + } + rawWire := requireWireRef(t, rawRef) + zstdWire := requireWireRef(t, zstdRef) + + tests := []struct { + name string + maxGrowth int64 + factory func(t *testing.T, size int64) func() error + }{ + { + name: "encode identity", + maxGrowth: identityMaxAllocationGrowth, + factory: func(_ *testing.T, size int64) func() error { + return func() error { + return EncodeWire(context.Background(), rawRef, newWireBenchmarkReader(size), io.Discard) + } + }, + }, + { + name: "encode zstd", + maxGrowth: zstdMaxAllocationGrowth, + factory: func(t *testing.T, size int64) func() error { + requireIncompressibleWireFixture(t, zstdRef, size) + return func() error { + return EncodeWire(context.Background(), zstdRef, newWireBenchmarkReader(size), io.Discard) + } + }, + }, + { + name: "decode identity", + maxGrowth: identityMaxAllocationGrowth, + factory: func(_ *testing.T, size int64) func() error { + return func() error { + return DecodeWire( + context.Background(), rawWire, newWireBenchmarkReader(size), io.Discard, + WireLimits{MaxEncodedBytes: size, MaxDecodedBytes: size}, + ) + } + }, + }, + { + name: "decode zstd", + maxGrowth: zstdMaxAllocationGrowth, + factory: func(t *testing.T, size int64) func() error { + encoded := requireIncompressibleWireFixture(t, zstdRef, size) + return func() error { + return DecodeWire( + context.Background(), zstdWire, bytes.NewReader(encoded), io.Discard, + WireLimits{ + MaxEncodedBytes: int64(len(encoded)), + MaxDecodedBytes: size, + }, + ) + } + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + small := wireCodecAllocatedBytes(t, tt.factory(t, smallSize)) + large := wireCodecAllocatedBytes(t, tt.factory(t, largeSize)) + t.Logf("allocated bytes/op: small=%d large=%d max_growth=%d", + small, large, tt.maxGrowth) + assert.LessOrEqual(t, large, small+tt.maxGrowth, + "allocation bytes must stay bounded as the artifact grows from %d to %d bytes; small=%d large=%d max_growth=%d", + smallSize, largeSize, small, large, tt.maxGrowth) + }) + } +} + +func TestWireCreateFromWireUsesPrivateCanonicalSpoolAndExactIdentity(t *testing.T) { + spoolDir := t.TempDir() + t.Setenv("TMPDIR", spoolDir) + body := fmt.Appendf(nil, `{"origin":%q,"seq":1,"sessions":{},"v":1}`+"\n", contractOrigin) + identity := identityForBytes(t, body) + ref := Ref{ + Origin: contractOrigin, + Kind: KindCheckpoints, + Name: "cp-0000000001.json", + } + wireRef := requireWireRef(t, ref) + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &encoded)) + store := &recordingWireStore{} + + result, err := CreateFromWire(t.Context(), store, wireRef, bytes.NewReader(encoded.Bytes()), WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body)), + }) + require.NoError(t, err) + assert.True(t, result.Created) + assert.Equal(t, ref, store.ref) + assert.Equal(t, identity, store.identity) + assert.Equal(t, "application/json", store.mediaType) + assert.Equal(t, body, store.body) + assert.Equal(t, os.FileMode(0o600), store.spoolMode.Perm()) + assert.NoFileExists(t, store.spoolName) + entries, err := os.ReadDir(spoolDir) + require.NoError(t, err) + assert.Empty(t, entries) +} + +func TestWireCreateFromWireRejectsNameHashMismatchBeforeStoreCreate(t *testing.T) { + spoolDir := t.TempDir() + t.Setenv("TMPDIR", spoolDir) + body := []byte("canonical segment bytes\n") + ref := Ref{ + Origin: contractOrigin, + Kind: KindSegments, + Name: strings.Repeat("f", 64) + ".ndjson", + } + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &encoded)) + store := &recordingWireStore{} + + _, err := CreateFromWire(t.Context(), store, requireWireRef(t, ref), bytes.NewReader(encoded.Bytes()), WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), + MaxDecodedBytes: int64(len(body)), + }) + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.False(t, store.called) + entries, readErr := os.ReadDir(spoolDir) + require.NoError(t, readErr) + assert.Empty(t, entries) +} + +func TestWireCreateFromWireRejectsSemanticallyInvalidCanonicalContent(t *testing.T) { + body := []byte(`{"not":"a manifest"}`) + identity := identityForBytes(t, body) + ref := Ref{ + Origin: contractOrigin, + Kind: KindManifests, + Name: identity.SHA256 + ".json", + } + var encoded bytes.Buffer + require.NoError(t, EncodeWire(t.Context(), ref, bytes.NewReader(body), &encoded)) + store := &recordingWireStore{} + + _, err := CreateFromWire(t.Context(), store, requireWireRef(t, ref), + bytes.NewReader(encoded.Bytes()), WireLimits{ + MaxEncodedBytes: int64(encoded.Len()), MaxDecodedBytes: int64(len(body)), + }) + + require.Error(t, err) + assert.ErrorIs(t, err, ErrArtifactInvalid) + assert.False(t, store.called) +} + +func TestWireCreateFromWireRemovesSpoolWhenStoreFails(t *testing.T) { + spoolDir := t.TempDir() + t.Setenv("TMPDIR", spoolDir) + body := []byte("raw canonical bytes") + identity := identityForBytes(t, body) + ref := Ref{Origin: contractOrigin, Kind: KindRaw, Name: identity.SHA256} + storeErr := errors.New("store unavailable") + store := &recordingWireStore{createErr: storeErr} + + _, err := CreateFromWire(t.Context(), store, requireWireRef(t, ref), bytes.NewReader(body), WireLimits{ + MaxEncodedBytes: int64(len(body)), + MaxDecodedBytes: int64(len(body)), + }) + assert.ErrorIs(t, err, storeErr) + assert.True(t, store.called) + assert.NoFileExists(t, store.spoolName) + entries, readErr := os.ReadDir(spoolDir) + require.NoError(t, readErr) + assert.Empty(t, entries) +} + +func requireWireRef(t *testing.T, ref Ref) WireRef { + t.Helper() + wireRef, err := ToWireRef(ref) + require.NoError(t, err) + return wireRef +} + +type cancelAfterReader struct { + reader io.Reader + cancel context.CancelFunc + remaining int64 + perRead int + read int64 + canceled bool +} + +type cancelOnFinalRead struct { + cancel context.CancelFunc + body []byte +} + +type singleReadEOF struct { + body []byte +} + +func (r *singleReadEOF) Read(p []byte) (int, error) { + if len(r.body) == 0 { + return 0, io.EOF + } + n := copy(p, r.body) + r.body = r.body[n:] + if len(r.body) == 0 { + return n, io.EOF + } + return n, nil +} + +type cancelOnSuccessfulWrite struct { + writer io.Writer + cancel context.CancelFunc + cancelOn int + writes int +} + +func (w *cancelOnSuccessfulWrite) Write(p []byte) (int, error) { + n, err := w.writer.Write(p) + if err == nil && n == len(p) { + w.writes++ + if w.cancel != nil && w.writes == w.cancelOn { + w.cancel() + } + } + return n, err +} + +func (r *cancelOnFinalRead) Read(p []byte) (int, error) { + if len(r.body) == 0 { + return 0, io.EOF + } + n := copy(p, r.body) + r.body = r.body[n:] + if len(r.body) == 0 { + r.cancel() + return n, io.EOF + } + return n, nil +} + +func (r *cancelAfterReader) Read(p []byte) (int, error) { + if r.remaining == 0 { + return 0, io.EOF + } + if len(p) > r.perRead { + p = p[:r.perRead] + } + if int64(len(p)) > r.remaining { + p = p[:r.remaining] + } + var n int + var err error + if r.reader != nil { + n, err = r.reader.Read(p) + } else { + for i := range p { + p[i] = byte(i) + } + n = len(p) + } + r.remaining -= int64(n) + r.read += int64(n) + if !r.canceled && n > 0 { + r.canceled = true + r.cancel() + } + return n, err +} + +type boundedPatternReader struct { + remaining int64 + offset int64 + maxRequest int + largestRequest int +} + +type deterministicNoiseReader struct { + remaining int64 + state uint64 +} + +func newWireBenchmarkReader(size int64) *deterministicNoiseReader { + return &deterministicNoiseReader{ + remaining: size, + state: 0x9e3779b97f4a7c15, + } +} + +func (r *deterministicNoiseReader) Read(p []byte) (int, error) { + if r.remaining == 0 { + return 0, io.EOF + } + if int64(len(p)) > r.remaining { + p = p[:r.remaining] + } + state := r.state + for i := range p { + state ^= state << 13 + state ^= state >> 7 + state ^= state << 17 + p[i] = byte(state >> 56) + } + r.state = state + r.remaining -= int64(len(p)) + return len(p), nil +} + +func requireIncompressibleWireFixture(t *testing.T, ref Ref, size int64) []byte { + t.Helper() + var wire bytes.Buffer + require.NoError(t, EncodeWire( + t.Context(), ref, newWireBenchmarkReader(size), &wire, + )) + encoded := append([]byte(nil), wire.Bytes()...) + require.Greater(t, int64(len(encoded)), size*9/10, + "fixture wire bytes must remain proportional to canonical size") + return encoded +} + +func wireCodecAllocatedBytes(t *testing.T, run func() error) int64 { + t.Helper() + require.NoError(t, run()) + var runErr error + result := testing.Benchmark(func(b *testing.B) { + b.ReportAllocs() + for range b.N { + if err := run(); err != nil { + runErr = err + b.StopTimer() + return + } + } + }) + require.NoError(t, runErr) + return result.AllocedBytesPerOp() +} + +func (r *boundedPatternReader) Read(p []byte) (int, error) { + if len(p) > r.largestRequest { + r.largestRequest = len(p) + } + if len(p) > r.maxRequest { + return 0, fmt.Errorf("artifact-sized read buffer: %d", len(p)) + } + if r.remaining == 0 { + return 0, io.EOF + } + if int64(len(p)) > r.remaining { + p = p[:r.remaining] + } + for i := range p { + p[i] = byte((r.offset + int64(i)*31) % 251) + } + n := len(p) + r.offset += int64(n) + r.remaining -= int64(n) + return n, nil +} + +type boundedWriteObserver struct { + writer io.Writer + maxWrite int + largestWrite int + written int64 +} + +func (w *boundedWriteObserver) Write(p []byte) (int, error) { + if len(p) > w.largestWrite { + w.largestWrite = len(p) + } + if len(p) > w.maxWrite { + return 0, fmt.Errorf("artifact-sized write buffer: %d", len(p)) + } + n, err := w.writer.Write(p) + w.written += int64(n) + return n, err +} + +type recordingWireStore struct { + ArtifactStore + called bool + ref Ref + identity Identity + mediaType string + body []byte + spoolName string + spoolMode os.FileMode + createErr error +} + +func (s *recordingWireStore) Create( + _ context.Context, + ref Ref, + identity Identity, + mediaType string, + body io.Reader, +) (CreateResult, error) { + s.called = true + s.ref = ref + s.identity = identity + s.mediaType = mediaType + if file, ok := body.(*os.File); ok { + s.spoolName = file.Name() + info, err := file.Stat() + if err != nil { + return CreateResult{}, err + } + s.spoolMode = info.Mode() + } + var err error + s.body, err = io.ReadAll(body) + if err != nil { + return CreateResult{}, err + } + if s.createErr != nil { + return CreateResult{}, s.createErr + } + return CreateResult{ + Entry: Entry{Ref: ref, Identity: identity}, + Created: true, + }, nil +} diff --git a/internal/config/config.go b/internal/config/config.go index 8c146a25c..8b6dcab2c 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -4,7 +4,9 @@ import ( "bytes" "crypto/rand" "encoding/base64" + "encoding/hex" "encoding/json" + "errors" "flag" "fmt" "log" @@ -455,6 +457,7 @@ type Config struct { GithubToken string `json:"github_token,omitempty" toml:"github_token"` Terminal TerminalConfig `json:"terminal,omitempty" toml:"terminal"` AuthToken string `json:"auth_token,omitempty" toml:"auth_token"` + ArtifactOriginID string `json:"artifact_origin_id,omitempty" toml:"artifact_origin_id"` RequireAuth bool `json:"require_auth" toml:"require_auth"` NoBrowser bool `json:"no_browser" toml:"no_browser"` DisableUpdateCheck bool `json:"disable_update_check" toml:"disable_update_check"` @@ -1016,6 +1019,7 @@ func (c *Config) applyConfigTOML(data string) error { ResultContentBlockedCategories []string `toml:"result_content_blocked_categories"` Terminal TerminalConfig `toml:"terminal"` AuthToken string `toml:"auth_token"` + ArtifactOriginID string `toml:"artifact_origin_id"` RequireAuth bool `toml:"require_auth"` RemoteAccess bool `toml:"remote_access"` DisableUpdateCheck bool `toml:"disable_update_check"` @@ -1087,6 +1091,9 @@ func (c *Config) applyConfigTOML(data string) error { if file.AuthToken != "" && c.AuthToken == "" { c.AuthToken = file.AuthToken } + if file.ArtifactOriginID != "" { + c.ArtifactOriginID = file.ArtifactOriginID + } c.RequireAuth = file.RequireAuth || file.RemoteAccess c.DisableUpdateCheck = file.DisableUpdateCheck if meta.IsDefined("default_pg") { @@ -1706,6 +1713,11 @@ func finalize(cfg *Config) error { if err := cfg.Recall.Extract.Validate(); err != nil { return err } + if cfg.ArtifactOriginID != "" { + if err := ValidateArtifactOriginID(cfg.ArtifactOriginID); err != nil { + return fmt.Errorf("invalid artifact origin id: %w", err) + } + } return nil } @@ -2487,6 +2499,11 @@ func (c *Config) SaveSettings(patch map[string]any) error { c.AuthToken = s } } + if v, ok := patch["artifact_origin_id"]; ok { + if s, ok := v.(string); ok { + c.ArtifactOriginID = s + } + } if v, ok := patch["require_auth"]; ok { if b, ok := v.(bool); ok { c.RequireAuth = b @@ -2526,6 +2543,176 @@ func (c *Config) EnsureAuthToken() error { }) } +// EnsureArtifactOriginID generates and persists a stable artifact sync origin +// ID if one does not already exist. +func (c *Config) EnsureArtifactOriginID() (string, error) { + if c.ArtifactOriginID != "" { + if err := ValidateArtifactOriginID(c.ArtifactOriginID); err != nil { + return "", fmt.Errorf("stored artifact origin: %w", err) + } + return c.ArtifactOriginID, nil + } + + var origin string + if err := c.withConfigLock(func() error { + existing, err := c.readConfigMap() + if err != nil { + return err + } + if stored, ok := existing["artifact_origin_id"].(string); ok && stored != "" { + if err := ValidateArtifactOriginID(stored); err != nil { + return fmt.Errorf("stored artifact origin: %w", err) + } + c.ArtifactOriginID = stored + origin = stored + return nil + } + + machineName, err := c.artifactOriginMachineName() + if err != nil { + return err + } + generated, err := newArtifactOriginID(machineName) + if err != nil { + return err + } + if err := ValidateArtifactOriginID(generated); err != nil { + return fmt.Errorf("generated artifact origin: %w", err) + } + + existing["artifact_origin_id"] = generated + if err := c.writeConfigMap(existing); err != nil { + return err + } + c.ArtifactOriginID = generated + origin = generated + return nil + }); err != nil { + return "", err + } + return origin, nil +} + +// AdoptArtifactOriginID persists origin as this machine's artifact sync +// origin unless the config file already records one, in which case the +// recorded origin wins and is returned. Artifact sync uses this to promote an +// origin that exists only in database sync state -- for example one minted by +// an incoming peer exchange before the config ever initialized an origin -- +// so the machine keeps publishing under a single origin instead of generating +// a competing config origin that would strand earlier metadata events. +func (c *Config) AdoptArtifactOriginID(origin string) (string, error) { + if err := ValidateArtifactOriginID(origin); err != nil { + return "", fmt.Errorf("adopting artifact origin: %w", err) + } + if c.ArtifactOriginID != "" { + if err := ValidateArtifactOriginID(c.ArtifactOriginID); err != nil { + return "", fmt.Errorf("stored artifact origin: %w", err) + } + return c.ArtifactOriginID, nil + } + + adopted := origin + if err := c.withConfigLock(func() error { + existing, err := c.readConfigMap() + if err != nil { + return err + } + if stored, ok := existing["artifact_origin_id"].(string); ok && stored != "" { + if err := ValidateArtifactOriginID(stored); err != nil { + return fmt.Errorf("stored artifact origin: %w", err) + } + c.ArtifactOriginID = stored + adopted = stored + return nil + } + + existing["artifact_origin_id"] = origin + if err := c.writeConfigMap(existing); err != nil { + return err + } + c.ArtifactOriginID = origin + return nil + }); err != nil { + return "", err + } + return adopted, nil +} + +func (c *Config) artifactOriginMachineName() (string, error) { + pgMachine := strings.TrimSpace(c.PG.MachineName) + if pgMachine != "" { + if pgMachine == "local" { + return "", fmt.Errorf( + "machine name %q is reserved; choose a different pg.machine_name", + pgMachine, + ) + } + return c.PG.MachineName, nil + } + host, err := os.Hostname() + if err != nil || strings.TrimSpace(host) == "" { + return "machine", nil + } + return host, nil +} + +func newArtifactOriginID(machine string) (string, error) { + base := sanitizeArtifactOriginPart(machine) + if base == "" || base == "local" { + base = "machine" + } + var suffix [3]byte + if _, err := rand.Read(suffix[:]); err != nil { + return "", fmt.Errorf("generating artifact origin suffix: %w", err) + } + return fmt.Sprintf("%s-%s", base, hex.EncodeToString(suffix[:])), nil +} + +func sanitizeArtifactOriginPart(s string) string { + s = strings.ToLower(strings.TrimSpace(s)) + var b strings.Builder + lastDash := false + for _, r := range s { + ok := r >= 'a' && r <= 'z' || r >= '0' && r <= '9' + if ok { + b.WriteRune(r) + lastDash = false + continue + } + if !lastDash { + b.WriteByte('-') + lastDash = true + } + } + return strings.Trim(b.String(), "-") +} + +// ValidateArtifactOriginID checks the persisted single-writer origin prefix. +func ValidateArtifactOriginID(origin string) error { + if origin == "" { + return errors.New("artifact origin is required") + } + if origin != strings.TrimSpace(origin) { + return fmt.Errorf("invalid artifact origin %q", origin) + } + if origin == "local" { + return fmt.Errorf("invalid artifact origin %q", origin) + } + if strings.ContainsAny(origin, `/\`) || filepath.Base(origin) != origin { + return fmt.Errorf("invalid artifact origin %q", origin) + } + if strings.HasPrefix(origin, "-") || strings.HasSuffix(origin, "-") { + return fmt.Errorf("invalid artifact origin %q", origin) + } + for _, r := range origin { + if r >= 'a' && r <= 'z' || r >= '0' && r <= '9' || r == '-' { + continue + } + return fmt.Errorf("invalid artifact origin %q", origin) + } + return nil +} + // SaveGithubToken persists the GitHub token to the config file. func (c *Config) SaveGithubToken(token string) error { return c.withConfigLock(func() error { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index f3110704e..99f46204d 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -13,6 +13,7 @@ import ( "time" "github.com/BurntSushi/toml" + "github.com/gofrs/flock" "github.com/spf13/pflag" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -546,6 +547,221 @@ func TestLoad_PublicURLMergedIntoOrigins(t *testing.T) { assert.Equal(t, "https://viewer.example.test", strings.Join(cfg.PublicOrigins, ",")) } +func TestLoad_ArtifactOriginIDFromConfigFile(t *testing.T) { + tmp := setupTestEnv(t) + writeConfig(t, tmp, map[string]any{ + "artifact_origin_id": "desk-abcdef", + }) + + cfg, err := LoadMinimal() + require.NoError(t, err) + + assert.Equal(t, "desk-abcdef", cfg.ArtifactOriginID) +} + +func TestLoad_ArtifactOriginIDRejectsInvalid(t *testing.T) { + tmp := setupTestEnv(t) + writeConfig(t, tmp, map[string]any{ + "artifact_origin_id": "local", + }) + + _, err := LoadMinimal() + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid artifact origin id") +} + +func TestEnsureArtifactOriginIDPersists(t *testing.T) { + tmp := setupTestEnv(t) + cfg, err := LoadMinimal() + require.NoError(t, err) + + origin, err := cfg.EnsureArtifactOriginID() + require.NoError(t, err) + assert.Regexp(t, `^[a-z0-9]+(?:-[a-z0-9]+)*-[0-9a-f]{6}$`, origin) + + data, err := os.ReadFile(filepath.Join(tmp, configFileName)) + require.NoError(t, err) + assert.Contains(t, string(data), `artifact_origin_id = "`+origin+`"`) + + reloaded, err := LoadMinimal() + require.NoError(t, err) + assert.Equal(t, origin, reloaded.ArtifactOriginID) + again, err := reloaded.EnsureArtifactOriginID() + require.NoError(t, err) + assert.Equal(t, origin, again) +} + +func TestEnsureArtifactOriginIDUsesConfiguredMachineName(t *testing.T) { + tmp := setupTestEnv(t) + writeConfig(t, tmp, map[string]any{ + "pg": map[string]any{ + "machine_name": "Desk Box", + }, + }) + cfg, err := LoadMinimal() + require.NoError(t, err) + + origin, err := cfg.EnsureArtifactOriginID() + require.NoError(t, err) + + assert.Regexp(t, `^desk-box-[0-9a-f]{6}$`, origin) +} + +func TestEnsureArtifactOriginIDRejectsReservedMachineName(t *testing.T) { + tmp := setupTestEnv(t) + writeConfig(t, tmp, map[string]any{ + "pg": map[string]any{ + "machine_name": "local", + }, + }) + cfg, err := LoadMinimal() + require.NoError(t, err) + + origin, err := cfg.EnsureArtifactOriginID() + require.Error(t, err) + assert.Empty(t, origin) + assert.Contains(t, err.Error(), "reserved") +} + +func TestEnsureArtifactOriginIDDoesNotRewriteExistingOrigin(t *testing.T) { + tmp := setupTestEnv(t) + writeConfig(t, tmp, map[string]any{ + "artifact_origin_id": "original-a1b2c3", + "pg": map[string]any{ + "machine_name": "new-machine", + }, + }) + cfg, err := LoadMinimal() + require.NoError(t, err) + + origin, err := cfg.EnsureArtifactOriginID() + require.NoError(t, err) + + assert.Equal(t, "original-a1b2c3", origin) +} + +func TestAdoptArtifactOriginIDPersists(t *testing.T) { + tmp := setupTestEnv(t) + cfg, err := LoadMinimal() + require.NoError(t, err) + + adopted, err := cfg.AdoptArtifactOriginID("laptop-a1b2c3") + require.NoError(t, err) + assert.Equal(t, "laptop-a1b2c3", adopted) + + data, err := os.ReadFile(filepath.Join(tmp, configFileName)) + require.NoError(t, err) + assert.Contains(t, string(data), `artifact_origin_id = "laptop-a1b2c3"`) + + reloaded, err := LoadMinimal() + require.NoError(t, err) + origin, err := reloaded.EnsureArtifactOriginID() + require.NoError(t, err) + assert.Equal(t, "laptop-a1b2c3", origin, + "later ensure must reuse the adopted origin instead of generating") +} + +func TestAdoptArtifactOriginIDKeepsExistingFileOrigin(t *testing.T) { + tmp := setupTestEnv(t) + cfg, err := LoadMinimal() + require.NoError(t, err) + + // The config file gains an origin after load but before adoption; the + // recorded origin wins. + writeConfig(t, tmp, map[string]any{ + "artifact_origin_id": "desktop-d4e5f6", + }) + + adopted, err := cfg.AdoptArtifactOriginID("laptop-a1b2c3") + require.NoError(t, err) + assert.Equal(t, "desktop-d4e5f6", adopted) + assert.Equal(t, "desktop-d4e5f6", cfg.ArtifactOriginID) +} + +func TestAdoptArtifactOriginIDKeepsInMemoryOrigin(t *testing.T) { + setupTestEnv(t) + cfg, err := LoadMinimal() + require.NoError(t, err) + cfg.ArtifactOriginID = "desktop-d4e5f6" + + adopted, err := cfg.AdoptArtifactOriginID("laptop-a1b2c3") + require.NoError(t, err) + assert.Equal(t, "desktop-d4e5f6", adopted) +} + +func TestAdoptArtifactOriginIDRejectsInvalidOrigin(t *testing.T) { + setupTestEnv(t) + cfg, err := LoadMinimal() + require.NoError(t, err) + + _, err = cfg.AdoptArtifactOriginID("local") + require.Error(t, err) + assert.Empty(t, cfg.ArtifactOriginID) +} + +func TestEnsureArtifactOriginIDReusesOriginWrittenBeforeLock(t *testing.T) { + tmp := setupTestEnv(t) + cfg, err := LoadMinimal() + require.NoError(t, err) + + lock := flock.New(cfg.configPath() + ".lock") + require.NoError(t, lock.Lock()) + t.Cleanup(func() { + _ = lock.Unlock() + }) + + type ensureResult struct { + origin string + err error + } + done := make(chan ensureResult, 1) + go func() { + origin, err := cfg.EnsureArtifactOriginID() + done <- ensureResult{origin: origin, err: err} + }() + + select { + case res := <-done: + require.Failf(t, "EnsureArtifactOriginID ignored config lock", + "origin=%q err=%v", res.origin, res.err) + case <-time.After(250 * time.Millisecond): + } + + writeConfig(t, tmp, map[string]any{ + "artifact_origin_id": "winner-a1b2c3", + }) + require.NoError(t, lock.Unlock()) + + select { + case res := <-done: + require.NoError(t, res.err) + assert.Equal(t, "winner-a1b2c3", res.origin) + assert.Equal(t, "winner-a1b2c3", cfg.ArtifactOriginID) + case <-time.After(2 * time.Second): + require.Fail(t, "EnsureArtifactOriginID did not finish after lock release") + } +} + +func TestMigrateJSONToTOMLPreservesArtifactOriginID(t *testing.T) { + tmp := setupTestEnv(t) + jsonPath := filepath.Join(tmp, "config.json") + require.NoError(t, os.WriteFile( + jsonPath, + []byte(`{"artifact_origin_id":"desk-abcdef"}`), + 0o600, + )) + + cfg, err := LoadMinimal() + require.NoError(t, err) + + assert.Equal(t, "desk-abcdef", cfg.ArtifactOriginID) + _, err = os.Stat(jsonPath + ".bak") + require.NoError(t, err) + data, err := os.ReadFile(filepath.Join(tmp, configFileName)) + require.NoError(t, err) + assert.Contains(t, string(data), `artifact_origin_id = "desk-abcdef"`) +} + func TestLoad_ProxyConfigFromFile(t *testing.T) { cfg := loadMinimalWithConfig(t, map[string]any{ "public_url": "https://viewer.example.test", diff --git a/internal/db/artifact_publication.go b/internal/db/artifact_publication.go new file mode 100644 index 000000000..d61767879 --- /dev/null +++ b/internal/db/artifact_publication.go @@ -0,0 +1,1365 @@ +package db + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "strconv" + "strings" +) + +const maxArtifactQueuePageSize = 1024 + +const artifactResetRepublishPendingKey = "artifact_reset_republish_pending" + +var ( + // ErrArtifactExportClaimStale tells callers to discard computed export + // output and retry from a fresh queue claim. + ErrArtifactExportClaimStale = errors.New("artifact export claim is stale") + // ErrArtifactRepairClaimStale tells callers that the expected repair + // identity changed while repair work was in flight. + ErrArtifactRepairClaimStale = errors.New("artifact repair claim is stale") +) + +// ArtifactExportQueueItem identifies one locally-owned session whose artifact +// publication may no longer match the archive. +type ArtifactExportQueueItem struct { + SessionID string + EnqueuedAt string + Generation int64 +} + +// ArtifactImportWork identifies one exact immutable artifact whose import must +// be retried after the current bounded transfer pass. +type ArtifactImportWork struct { + Origin string + Kind string + Name string + SHA256 string + Size int64 + Reason string + RequiredFormatVersion int + EnqueuedAt string +} + +func validateArtifactImportWork(work ArtifactImportWork, requireClaim bool) error { + if strings.TrimSpace(work.Origin) == "" || work.Origin != strings.TrimSpace(work.Origin) { + return errors.New("artifact import origin is required") + } + if work.Kind != "checkpoints" && work.Kind != "meta" { + return errors.New("artifact import kind must be checkpoints or meta") + } + if strings.TrimSpace(work.Name) == "" || strings.ContainsAny(work.Name, `/\\`) { + return errors.New("artifact import name is required") + } + if len(work.SHA256) != 64 { + return errors.New("complete artifact import identity is required") + } + for _, c := range work.SHA256 { + if (c < '0' || c > '9') && (c < 'a' || c > 'f') { + return errors.New("artifact import identity must be lowercase hexadecimal") + } + } + if work.Size < 0 { + return errors.New("artifact import size must not be negative") + } + if strings.TrimSpace(work.Reason) == "" || work.Reason != strings.TrimSpace(work.Reason) { + return errors.New("artifact import reason is required") + } + if work.RequiredFormatVersion < 1 { + return errors.New("artifact import required format version must be positive") + } + if work.Kind == "meta" { + base := strings.TrimSuffix(work.Name, ".json") + separator := strings.LastIndexByte(base, '-') + if base == work.Name || separator < 1 || base[separator+1:] != work.SHA256 { + return errors.New("artifact import metadata name must match its identity") + } + } + if work.Kind == "checkpoints" { + if _, err := artifactImportCheckpointSequence(work.Name); err != nil { + return err + } + } + if requireClaim && strings.TrimSpace(work.EnqueuedAt) == "" { + return errors.New("artifact import enqueue time is required") + } + return nil +} + +func artifactImportCheckpointSequence(name string) (int64, error) { + if len(name) != len("cp-0000000000.json") || !strings.HasPrefix(name, "cp-") || + !strings.HasSuffix(name, ".json") { + return 0, errors.New("canonical artifact checkpoint name is required") + } + digits := strings.TrimSuffix(strings.TrimPrefix(name, "cp-"), ".json") + for _, digit := range digits { + if digit < '0' || digit > '9' { + return 0, errors.New("canonical artifact checkpoint name is required") + } + } + sequence, err := strconv.ParseInt(digits, 10, 64) + if err != nil || sequence < 1 { + return 0, errors.New("positive artifact checkpoint sequence is required") + } + return sequence, nil +} + +// EnqueueArtifactImport retains one exact retry claim. Repeated observations +// of the same immutable identity are idempotent; conflicting identities fail +// closed. A newer checkpoint supersedes older queued checkpoints for its +// origin, while metadata events are retained independently. +func (db *DB) EnqueueArtifactImport(ctx context.Context, work ArtifactImportWork) error { + if err := validateArtifactImportWork(work, false); err != nil { + return err + } + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("beginning artifact import enqueue: %w", err) + } + defer func() { _ = tx.Rollback() }() + + var existingSHA string + var existingSize int64 + err = tx.QueryRowContext(ctx, ` + SELECT sha256, size FROM artifact_import_queue + WHERE origin = ? AND kind = ? AND name = ?`, + work.Origin, work.Kind, work.Name, + ).Scan(&existingSHA, &existingSize) + if err == nil { + if existingSHA != work.SHA256 || existingSize != work.Size { + return errors.New("artifact import reference has a conflicting identity") + } + if _, err := tx.ExecContext(ctx, ` + UPDATE artifact_import_queue SET + reason = ?, + required_format_version = max(required_format_version, ?) + WHERE origin = ? AND kind = ? AND name = ?`, + work.Reason, work.RequiredFormatVersion, work.Origin, work.Kind, work.Name, + ); err != nil { + return fmt.Errorf("refreshing artifact import work: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing artifact import refresh: %w", err) + } + return nil + } + if !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("reading artifact import identity: %w", err) + } + + if work.Kind == "checkpoints" { + var newer bool + if err := tx.QueryRowContext(ctx, ` + SELECT EXISTS( + SELECT 1 FROM artifact_import_queue + WHERE origin = ? AND kind = 'checkpoints' AND name > ? + )`, work.Origin, work.Name).Scan(&newer); err != nil { + return fmt.Errorf("reading newer artifact checkpoint work: %w", err) + } + if newer { + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing superseded artifact checkpoint: %w", err) + } + return nil + } + if _, err := tx.ExecContext(ctx, ` + DELETE FROM artifact_import_queue + WHERE origin = ? AND kind = 'checkpoints' AND name < ?`, + work.Origin, work.Name, + ); err != nil { + return fmt.Errorf("retiring older artifact checkpoint work: %w", err) + } + } + if _, err := tx.ExecContext(ctx, ` + INSERT INTO artifact_import_queue( + origin, kind, name, sha256, size, reason, required_format_version + ) VALUES (?, ?, ?, ?, ?, ?, ?)`, + work.Origin, work.Kind, work.Name, work.SHA256, work.Size, + work.Reason, work.RequiredFormatVersion, + ); err != nil { + return fmt.Errorf("enqueueing artifact import work: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing artifact import enqueue: %w", err) + } + return nil +} + +// PendingArtifactImports returns one bounded FIFO page whose required format +// is understood by readerFormatVersion. +func (db *DB) PendingArtifactImports( + ctx context.Context, readerFormatVersion, limit int, +) ([]ArtifactImportWork, error) { + if readerFormatVersion < 1 { + return nil, errors.New("artifact import reader format version must be positive") + } + if limit < 1 || limit > maxArtifactQueuePageSize { + return nil, fmt.Errorf("artifact import page size must be between 1 and %d", + maxArtifactQueuePageSize) + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT origin, kind, name, sha256, size, reason, + required_format_version, enqueued_at + FROM artifact_import_queue + WHERE required_format_version <= ? + ORDER BY enqueued_at, origin, kind, name + LIMIT ?`, readerFormatVersion, limit) + if err != nil { + return nil, fmt.Errorf("reading pending artifact imports: %w", err) + } + defer rows.Close() + work := make([]ArtifactImportWork, 0, min(limit, 64)) + for rows.Next() { + var item ArtifactImportWork + if err := rows.Scan( + &item.Origin, &item.Kind, &item.Name, &item.SHA256, &item.Size, + &item.Reason, &item.RequiredFormatVersion, &item.EnqueuedAt, + ); err != nil { + return nil, fmt.Errorf("scanning pending artifact import: %w", err) + } + work = append(work, item) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating pending artifact imports: %w", err) + } + return work, nil +} + +// AcknowledgeArtifactImport compare-and-deletes one exact queue claim. +func (db *DB) AcknowledgeArtifactImport( + ctx context.Context, work ArtifactImportWork, +) (bool, error) { + if err := validateArtifactImportWork(work, true); err != nil { + return false, err + } + db.mu.Lock() + defer db.mu.Unlock() + result, err := db.getWriter().ExecContext(ctx, ` + DELETE FROM artifact_import_queue + WHERE origin = ? AND kind = ? AND name = ? AND sha256 = ? AND size = ? + AND reason = ? AND required_format_version = ? AND enqueued_at = ?`, + work.Origin, work.Kind, work.Name, work.SHA256, work.Size, work.Reason, + work.RequiredFormatVersion, work.EnqueuedAt, + ) + if err != nil { + return false, fmt.Errorf("acknowledging artifact import: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return false, fmt.Errorf("reading artifact import acknowledgement: %w", err) + } + return rows == 1, nil +} + +// ArtifactImportQueueStats reports all durable unfinished work, including +// future-format rows that are not yet eligible to drain. +func (db *DB) ArtifactImportQueueStats(ctx context.Context) (int, string, error) { + var count int + var oldest string + if err := db.getReader().QueryRowContext(ctx, ` + SELECT count(*), coalesce(min(enqueued_at), '') + FROM artifact_import_queue`).Scan(&count, &oldest); err != nil { + return 0, "", fmt.Errorf("reading artifact import queue statistics: %w", err) + } + return count, oldest, nil +} + +// ArtifactResetRepublishPending is the durable authority for reconstructing a +// repository after its previous vault has been moved aside. RootFingerprint +// binds the intent to one canonical repository without persisting its path. +type ArtifactResetRepublishPending struct { + Version int `json:"v"` + RootFingerprint string `json:"root_fingerprint"` + Origin string `json:"origin"` + Token string `json:"token"` + BaselineHLC string `json:"baseline_hlc"` +} + +func validateArtifactResetRepublishPending(state ArtifactResetRepublishPending) error { + if state.Version != 1 || len(state.RootFingerprint) != 64 || + strings.TrimSpace(state.Origin) == "" || len(state.Token) != 64 || + strings.TrimSpace(state.BaselineHLC) == "" { + return errors.New("complete artifact reset republish state is required") + } + for _, value := range []string{state.RootFingerprint, state.Token} { + for _, c := range value { + if (c < '0' || c > '9') && (c < 'a' || c > 'f') { + return errors.New("artifact reset republish identity must be lowercase hexadecimal") + } + } + } + return nil +} + +// SetArtifactResetRepublishPending stores the singleton reset intent. +func (db *DB) SetArtifactResetRepublishPending( + ctx context.Context, state ArtifactResetRepublishPending, +) error { + if err := validateArtifactResetRepublishPending(state); err != nil { + return err + } + encoded, err := json.Marshal(state) + if err != nil { + return fmt.Errorf("encoding artifact reset republish state: %w", err) + } + db.mu.Lock() + defer db.mu.Unlock() + _, err = db.getWriter().ExecContext(ctx, ` + INSERT INTO pg_sync_state(key, value) VALUES (?, ?) + ON CONFLICT(key) DO UPDATE SET value = excluded.value`, + artifactResetRepublishPendingKey, string(encoded), + ) + if err != nil { + return fmt.Errorf("storing artifact reset republish state: %w", err) + } + return nil +} + +// ArtifactResetRepublishPending returns the durable reset intent, if present. +func (db *DB) ArtifactResetRepublishPending( + ctx context.Context, +) (ArtifactResetRepublishPending, bool, error) { + var encoded string + err := db.getReader().QueryRowContext(ctx, + `SELECT value FROM pg_sync_state WHERE key = ?`, + artifactResetRepublishPendingKey, + ).Scan(&encoded) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactResetRepublishPending{}, false, nil + } + if err != nil { + return ArtifactResetRepublishPending{}, false, + fmt.Errorf("reading artifact reset republish state: %w", err) + } + var state ArtifactResetRepublishPending + if err := json.Unmarshal([]byte(encoded), &state); err != nil { + return ArtifactResetRepublishPending{}, false, + fmt.Errorf("decoding artifact reset republish state: %w", err) + } + if err := validateArtifactResetRepublishPending(state); err != nil { + return ArtifactResetRepublishPending{}, false, err + } + return state, true, nil +} + +// ClearArtifactResetRepublishPending compare-and-swap clears one completed +// intent without allowing a stale recovery to erase a newer reset. +func (db *DB) ClearArtifactResetRepublishPending( + ctx context.Context, state ArtifactResetRepublishPending, +) (bool, error) { + if err := validateArtifactResetRepublishPending(state); err != nil { + return false, err + } + encoded, err := json.Marshal(state) + if err != nil { + return false, fmt.Errorf("encoding artifact reset republish state: %w", err) + } + db.mu.Lock() + defer db.mu.Unlock() + result, err := db.getWriter().ExecContext(ctx, + `DELETE FROM pg_sync_state WHERE key = ? AND value = ?`, + artifactResetRepublishPendingKey, string(encoded), + ) + if err != nil { + return false, fmt.Errorf("clearing artifact reset republish state: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return false, fmt.Errorf("reading artifact reset republish clear result: %w", err) + } + return rows == 1, nil +} + +// ArtifactPublication is the last manifest selected for one locally-owned +// session. Rows are the authority used to stream a full checkpoint map. +type ArtifactPublication struct { + Origin string + SessionID string + ManifestHash string + SourceFingerprint string +} + +// ArtifactPublicationChange changes or removes one publication row. Delete is +// used when a queued session is no longer locally owned or no longer exists. +type ArtifactPublicationChange struct { + SessionID string + Generation int64 + ManifestHash string + SourceFingerprint string + Delete bool +} + +// ArtifactCheckpointHead records the last successfully created checkpoint, +// the exact publication revision it represents, and its catalog identity. +type ArtifactCheckpointHead struct { + Origin string + Sequence int + PublicationRevision int64 + SessionMapSHA256 string + CheckpointSHA256 string + CheckpointSize int64 +} + +// ArtifactCheckpointLanding identifies the exact foreign checkpoint whose +// complete GID-to-manifest map has durable local provenance. +type ArtifactCheckpointLanding struct { + Origin string + Sequence int +} + +// ArtifactPeerCheckpointHead records the highest immutable checkpoint +// identity received for one foreign origin, including checkpoints whose +// dependency closure has not landed yet. +type ArtifactPeerCheckpointHead struct { + Origin string + Sequence int + CheckpointSHA256 string + CheckpointSize int64 +} + +// ArtifactRepair identifies canonical content whose physical representation +// must be repaired from a trusted peer. +type ArtifactRepair struct { + Origin string + Kind string + Name string + SHA256 string + Size int64 + DetectedAt string +} + +// PendingArtifactExports returns the oldest bounded page of dirty local +// sessions. Reading work never acknowledges it. +func (db *DB) PendingArtifactExports( + ctx context.Context, limit int, +) ([]ArtifactExportQueueItem, error) { + if err := validateArtifactQueueLimit(limit); err != nil { + return nil, err + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT session_id, enqueued_at, generation + FROM artifact_export_queue + WHERE pending = 1 + ORDER BY enqueued_at, session_id + LIMIT ?`, limit) + if err != nil { + return nil, fmt.Errorf("reading artifact export queue: %w", err) + } + defer rows.Close() + items := make([]ArtifactExportQueueItem, 0, min(limit, 64)) + for rows.Next() { + var item ArtifactExportQueueItem + if err := rows.Scan(&item.SessionID, &item.EnqueuedAt, &item.Generation); err != nil { + return nil, fmt.Errorf("scanning artifact export queue: %w", err) + } + items = append(items, item) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating artifact export queue: %w", err) + } + return items, nil +} + +// ArtifactExportClaims returns pending generation claims for the exact bounded +// session set requested by a watcher batch. Missing or already-clean IDs are +// omitted; mutation APIs revalidate every returned generation under a writer +// reservation before changing publication authority. +func (db *DB) ArtifactExportClaims( + ctx context.Context, sessionIDs []string, +) ([]ArtifactExportQueueItem, error) { + if len(sessionIDs) == 0 { + return []ArtifactExportQueueItem{}, nil + } + if len(sessionIDs) > maxArtifactQueuePageSize { + return nil, fmt.Errorf("artifact export claim batch exceeds %d rows", maxArtifactQueuePageSize) + } + unique := make([]string, 0, len(sessionIDs)) + seen := make(map[string]struct{}, len(sessionIDs)) + for _, id := range sessionIDs { + if id == "" { + return nil, errors.New("artifact export claim session id is required") + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + unique = append(unique, id) + } + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(unique)), ",") + args := make([]any, len(unique)) + for i, id := range unique { + args[i] = id + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT session_id, enqueued_at, generation + FROM artifact_export_queue + WHERE pending = 1 AND session_id IN (`+placeholders+`) + ORDER BY session_id`, args...) + if err != nil { + return nil, fmt.Errorf("reading exact artifact export claims: %w", err) + } + defer rows.Close() + items := make([]ArtifactExportQueueItem, 0, len(unique)) + for rows.Next() { + var item ArtifactExportQueueItem + if err := rows.Scan(&item.SessionID, &item.EnqueuedAt, &item.Generation); err != nil { + return nil, fmt.Errorf("scanning exact artifact export claim: %w", err) + } + items = append(items, item) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating exact artifact export claims: %w", err) + } + return items, nil +} + +// ApplyArtifactPublicationChanges atomically applies a bounded export batch and +// returns the resulting per-origin publication revision. The revision advances +// only when a publication row changes. Queue rows remain pending until the +// resulting checkpoint has been created. +func (db *DB) ApplyArtifactPublicationChanges( + ctx context.Context, origin string, changes []ArtifactPublicationChange, +) (int64, bool, error) { + if origin == "" { + return 0, false, errors.New("artifact publication origin is required") + } + if len(changes) > maxArtifactQueuePageSize { + return 0, false, fmt.Errorf( + "artifact publication batch exceeds %d rows", maxArtifactQueuePageSize, + ) + } + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return 0, false, fmt.Errorf("beginning artifact publication changes: %w", err) + } + defer func() { _ = tx.Rollback() }() + if err := lockArtifactPublicationTx(ctx, tx); err != nil { + return 0, false, err + } + claims := make([]ArtifactExportQueueItem, 0, len(changes)) + for _, change := range changes { + claims = append(claims, ArtifactExportQueueItem{ + SessionID: change.SessionID, Generation: change.Generation, + }) + } + if _, err := validateArtifactExportClaimsTx(ctx, tx, claims); err != nil { + return 0, false, err + } + changed := false + for _, change := range changes { + if change.SessionID == "" { + return 0, false, errors.New("artifact publication session id is required") + } + var result sql.Result + if change.Delete { + result, err = tx.ExecContext(ctx, ` + DELETE FROM artifact_publications + WHERE origin = ? AND session_id = ?`, origin, change.SessionID) + } else { + if change.ManifestHash == "" || change.SourceFingerprint == "" { + return 0, false, fmt.Errorf( + "artifact publication %s requires manifest hash and source fingerprint", + change.SessionID, + ) + } + result, err = tx.ExecContext(ctx, ` + INSERT INTO artifact_publications ( + origin, session_id, manifest_hash, source_fingerprint + ) VALUES (?, ?, ?, ?) + ON CONFLICT(origin, session_id) DO UPDATE SET + manifest_hash = excluded.manifest_hash, + source_fingerprint = excluded.source_fingerprint + WHERE artifact_publications.manifest_hash <> excluded.manifest_hash + OR artifact_publications.source_fingerprint <> excluded.source_fingerprint`, + origin, change.SessionID, change.ManifestHash, change.SourceFingerprint) + } + if err != nil { + return 0, false, fmt.Errorf("applying artifact publication %s: %w", change.SessionID, err) + } + if result == nil { + return 0, false, fmt.Errorf("applying artifact publication %s returned no result", change.SessionID) + } + rows, rowsErr := result.RowsAffected() + if rowsErr != nil { + return 0, false, fmt.Errorf("reading artifact publication result %s: %w", change.SessionID, rowsErr) + } + changed = changed || rows > 0 + } + revision, err := artifactPublicationRevisionTx(ctx, tx, origin, changed) + if err != nil { + return 0, false, err + } + if err := tx.Commit(); err != nil { + return 0, false, fmt.Errorf("committing artifact publication changes: %w", err) + } + return revision, changed, nil +} + +func artifactPublicationRevisionTx( + ctx context.Context, tx *sql.Tx, origin string, increment bool, +) (int64, error) { + var revision int64 + if increment { + err := tx.QueryRowContext(ctx, ` + INSERT INTO artifact_publication_revisions(origin, revision) VALUES (?, 1) + ON CONFLICT(origin) DO UPDATE SET revision = revision + 1 + RETURNING revision`, origin).Scan(&revision) + if err != nil { + return 0, fmt.Errorf("advancing artifact publication revision: %w", err) + } + return revision, nil + } + err := tx.QueryRowContext(ctx, ` + SELECT revision FROM artifact_publication_revisions WHERE origin = ?`, origin, + ).Scan(&revision) + if errors.Is(err, sql.ErrNoRows) { + return 0, nil + } + if err != nil { + return 0, fmt.Errorf("reading artifact publication revision: %w", err) + } + return revision, nil +} + +// AcknowledgeArtifactExports marks successfully processed work clean while +// retaining its generation authority when no checkpoint-head update is needed. +func (db *DB) AcknowledgeArtifactExports( + ctx context.Context, items []ArtifactExportQueueItem, +) error { + if len(items) > maxArtifactQueuePageSize { + return fmt.Errorf("artifact export acknowledgement exceeds %d rows", maxArtifactQueuePageSize) + } + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("beginning artifact export acknowledgement: %w", err) + } + defer func() { _ = tx.Rollback() }() + if err := lockArtifactPublicationTx(ctx, tx); err != nil { + return err + } + if err := acknowledgeArtifactExportsTx(ctx, tx, items); err != nil { + return err + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing artifact export acknowledgement: %w", err) + } + return nil +} + +// RecordArtifactCheckpointHead atomically records a successfully created head +// and acknowledges exactly the queue rows represented by that export batch. +func (db *DB) RecordArtifactCheckpointHead( + ctx context.Context, head ArtifactCheckpointHead, acknowledgedItems []ArtifactExportQueueItem, +) error { + if head.Origin == "" || head.Sequence < 1 || head.PublicationRevision < 0 || + head.SessionMapSHA256 == "" || head.CheckpointSHA256 == "" || head.CheckpointSize < 0 { + return errors.New("complete artifact checkpoint head is required") + } + if len(acknowledgedItems) > maxArtifactQueuePageSize { + return fmt.Errorf("artifact checkpoint acknowledgement exceeds %d rows", maxArtifactQueuePageSize) + } + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("beginning artifact checkpoint head: %w", err) + } + defer func() { _ = tx.Rollback() }() + if err := lockArtifactPublicationTx(ctx, tx); err != nil { + return err + } + currentRevision, err := artifactPublicationRevisionTx(ctx, tx, head.Origin, false) + if err != nil { + return err + } + if currentRevision != head.PublicationRevision { + return fmt.Errorf("%w: artifact publication revision %d is now %d", + ErrArtifactExportClaimStale, head.PublicationRevision, currentRevision) + } + uniqueClaims, err := validateArtifactExportClaimsTx(ctx, tx, acknowledgedItems) + if err != nil { + return err + } + result, err := tx.ExecContext(ctx, ` + INSERT INTO artifact_checkpoint_heads ( + origin, sequence, publication_revision, session_map_sha256, + checkpoint_sha256, checkpoint_size + ) VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT(origin) DO UPDATE SET + sequence = excluded.sequence, + publication_revision = excluded.publication_revision, + session_map_sha256 = excluded.session_map_sha256, + checkpoint_sha256 = excluded.checkpoint_sha256, + checkpoint_size = excluded.checkpoint_size + WHERE excluded.sequence > artifact_checkpoint_heads.sequence + OR ( + excluded.sequence = artifact_checkpoint_heads.sequence + AND excluded.session_map_sha256 = artifact_checkpoint_heads.session_map_sha256 + AND excluded.checkpoint_sha256 = artifact_checkpoint_heads.checkpoint_sha256 + AND excluded.checkpoint_size = artifact_checkpoint_heads.checkpoint_size + )`, + head.Origin, head.Sequence, head.PublicationRevision, head.SessionMapSHA256, + head.CheckpointSHA256, head.CheckpointSize, + ) + if err != nil { + return fmt.Errorf("recording artifact checkpoint head: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("reading artifact checkpoint head result: %w", err) + } + if rows != 1 { + return fmt.Errorf( + "artifact checkpoint head %s sequence %d conflicts with a newer or different head", + head.Origin, head.Sequence, + ) + } + if err := markArtifactExportClaimsCleanTx(ctx, tx, uniqueClaims); err != nil { + return err + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing artifact checkpoint head: %w", err) + } + return nil +} + +func acknowledgeArtifactExportsTx( + ctx context.Context, tx *sql.Tx, items []ArtifactExportQueueItem, +) error { + unique, err := validateArtifactExportClaimsTx(ctx, tx, items) + if err != nil { + return err + } + return markArtifactExportClaimsCleanTx(ctx, tx, unique) +} + +func validateArtifactExportClaimsTx( + ctx context.Context, tx *sql.Tx, items []ArtifactExportQueueItem, +) ([]ArtifactExportQueueItem, error) { + unique := make([]ArtifactExportQueueItem, 0, len(items)) + seen := make(map[string]int64, len(items)) + for _, item := range items { + if item.SessionID == "" || item.Generation < 1 { + return nil, errors.New("complete artifact export acknowledgement item is required") + } + if generation, ok := seen[item.SessionID]; ok { + if generation != item.Generation { + return nil, fmt.Errorf("%w: conflicting generations for session %s", + ErrArtifactExportClaimStale, item.SessionID) + } + continue + } + seen[item.SessionID] = item.Generation + unique = append(unique, item) + } + for _, item := range unique { + var generation int64 + var pending bool + err := tx.QueryRowContext(ctx, ` + SELECT generation, pending FROM artifact_export_queue WHERE session_id = ?`, + item.SessionID, + ).Scan(&generation, &pending) + if errors.Is(err, sql.ErrNoRows) || err == nil && (!pending || generation != item.Generation) { + return nil, fmt.Errorf("%w: session %s generation %d", + ErrArtifactExportClaimStale, item.SessionID, item.Generation) + } + if err != nil { + return nil, fmt.Errorf("validating artifact export claim %s: %w", item.SessionID, err) + } + } + return unique, nil +} + +func markArtifactExportClaimsCleanTx( + ctx context.Context, tx *sql.Tx, items []ArtifactExportQueueItem, +) error { + stmt, err := tx.PrepareContext(ctx, ` + UPDATE artifact_export_queue SET pending = 0 + WHERE session_id = ? AND generation = ? AND pending = 1`) + if err != nil { + return fmt.Errorf("preparing artifact export acknowledgement: %w", err) + } + defer stmt.Close() + for _, item := range items { + result, err := stmt.ExecContext(ctx, item.SessionID, item.Generation) + if err != nil { + return fmt.Errorf("acknowledging artifact export %s: %w", item.SessionID, err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("reading artifact export acknowledgement %s: %w", item.SessionID, err) + } + if rows != 1 { + return fmt.Errorf("%w: session %s generation %d", + ErrArtifactExportClaimStale, item.SessionID, item.Generation) + } + } + return nil +} + +// lockArtifactPublicationTx obtains SQLite's writer reservation before claim +// validation, closing the check-to-mutate race with other database handles. +func lockArtifactPublicationTx(ctx context.Context, tx *sql.Tx) error { + if _, err := tx.ExecContext(ctx, ` + UPDATE artifact_export_queue SET generation = generation WHERE 0`); err != nil { + return fmt.Errorf("locking artifact publication transaction: %w", err) + } + return nil +} + +// GetArtifactCheckpointHead returns the current recorded head for origin. +func (db *DB) GetArtifactCheckpointHead( + ctx context.Context, origin string, +) (ArtifactCheckpointHead, bool, error) { + var head ArtifactCheckpointHead + err := db.getReader().QueryRowContext(ctx, ` + SELECT origin, sequence, publication_revision, session_map_sha256, + checkpoint_sha256, checkpoint_size + FROM artifact_checkpoint_heads WHERE origin = ?`, origin).Scan( + &head.Origin, &head.Sequence, &head.PublicationRevision, + &head.SessionMapSHA256, &head.CheckpointSHA256, &head.CheckpointSize, + ) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactCheckpointHead{}, false, nil + } + if err != nil { + return ArtifactCheckpointHead{}, false, fmt.Errorf("reading artifact checkpoint head: %w", err) + } + return head, true, nil +} + +// RecordArtifactPeerCheckpointHead advances one foreign origin's received +// checkpoint head. Replaying the same immutable identity is idempotent; a +// different identity at the same sequence is rejected. +func (db *DB) RecordArtifactPeerCheckpointHead( + ctx context.Context, head ArtifactPeerCheckpointHead, +) error { + if head.Origin == "" || head.Sequence < 1 || head.CheckpointSHA256 == "" || head.CheckpointSize < 0 { + return errors.New("complete artifact peer checkpoint head is required") + } + db.mu.Lock() + defer db.mu.Unlock() + result, err := db.getWriter().ExecContext(ctx, ` + INSERT INTO artifact_peer_checkpoint_heads( + origin, sequence, checkpoint_sha256, checkpoint_size + ) VALUES (?, ?, ?, ?) + ON CONFLICT(origin) DO UPDATE SET + sequence = excluded.sequence, + checkpoint_sha256 = excluded.checkpoint_sha256, + checkpoint_size = excluded.checkpoint_size + WHERE excluded.sequence > artifact_peer_checkpoint_heads.sequence + OR (excluded.sequence = artifact_peer_checkpoint_heads.sequence + AND excluded.checkpoint_sha256 = artifact_peer_checkpoint_heads.checkpoint_sha256 + AND excluded.checkpoint_size = artifact_peer_checkpoint_heads.checkpoint_size)`, + head.Origin, head.Sequence, head.CheckpointSHA256, head.CheckpointSize, + ) + if err != nil { + return fmt.Errorf("recording artifact peer checkpoint head: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("reading artifact peer checkpoint head result: %w", err) + } + if rows != 1 { + return fmt.Errorf("artifact peer checkpoint head %s sequence %d conflicts with a newer or different head", + head.Origin, head.Sequence) + } + return nil +} + +// GetArtifactPeerCheckpointHead returns the exact highest received checkpoint +// identity for one foreign origin. +func (db *DB) GetArtifactPeerCheckpointHead( + ctx context.Context, origin string, +) (ArtifactPeerCheckpointHead, bool, error) { + var head ArtifactPeerCheckpointHead + err := db.getReader().QueryRowContext(ctx, ` + SELECT origin, sequence, checkpoint_sha256, checkpoint_size + FROM artifact_peer_checkpoint_heads WHERE origin = ?`, origin).Scan( + &head.Origin, &head.Sequence, &head.CheckpointSHA256, &head.CheckpointSize, + ) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactPeerCheckpointHead{}, false, nil + } + if err != nil { + return ArtifactPeerCheckpointHead{}, false, + fmt.Errorf("reading artifact peer checkpoint head: %w", err) + } + return head, true, nil +} + +// GetArtifactCheckpointLandingHead returns one foreign origin's landed checkpoint +// sequence without materializing its session map. +func (db *DB) GetArtifactCheckpointLandingHead( + ctx context.Context, origin string, +) (ArtifactCheckpointLanding, bool, error) { + var landing ArtifactCheckpointLanding + err := db.getReader().QueryRowContext(ctx, ` + SELECT origin, sequence FROM artifact_checkpoint_landings WHERE origin = ?`, origin, + ).Scan(&landing.Origin, &landing.Sequence) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactCheckpointLanding{}, false, nil + } + if err != nil { + return ArtifactCheckpointLanding{}, false, + fmt.Errorf("reading artifact checkpoint landing: %w", err) + } + return landing, true, nil +} + +// RecordArtifactCheckpointLanding atomically replaces one origin's exact +// landed checkpoint map. Older observations cannot regress durable provenance. +func (db *DB) RecordArtifactCheckpointLanding( + ctx context.Context, + landing ArtifactCheckpointLanding, + manifestByGID map[string]string, +) error { + if landing.Origin == "" || landing.Sequence < 1 { + return errors.New("complete artifact checkpoint landing is required") + } + for gid, manifestHash := range manifestByGID { + if gid == "" || manifestHash == "" { + return errors.New("complete artifact checkpoint landing session is required") + } + } + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("beginning artifact checkpoint landing: %w", err) + } + defer func() { _ = tx.Rollback() }() + var current int + err = tx.QueryRowContext(ctx, ` + SELECT sequence FROM artifact_checkpoint_landings WHERE origin = ?`, + landing.Origin, + ).Scan(¤t) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("reading artifact checkpoint landing: %w", err) + } + if err == nil && landing.Sequence < current { + return fmt.Errorf("artifact checkpoint landing %s sequence %d is older than %d", + landing.Origin, landing.Sequence, current) + } + if _, err := tx.ExecContext(ctx, ` + INSERT INTO artifact_checkpoint_landings(origin, sequence) VALUES (?, ?) + ON CONFLICT(origin) DO UPDATE SET sequence = excluded.sequence`, + landing.Origin, landing.Sequence, + ); err != nil { + return fmt.Errorf("recording artifact checkpoint landing: %w", err) + } + if _, err := tx.ExecContext(ctx, ` + DELETE FROM artifact_checkpoint_landing_sessions WHERE origin = ?`, + landing.Origin, + ); err != nil { + return fmt.Errorf("replacing artifact checkpoint landing sessions: %w", err) + } + stmt, err := tx.PrepareContext(ctx, ` + INSERT INTO artifact_checkpoint_landing_sessions(origin, gid, manifest_hash) + VALUES (?, ?, ?)`) + if err != nil { + return fmt.Errorf("preparing artifact checkpoint landing sessions: %w", err) + } + defer stmt.Close() + for gid, manifestHash := range manifestByGID { + if _, err := stmt.ExecContext(ctx, landing.Origin, gid, manifestHash); err != nil { + return fmt.Errorf("recording artifact checkpoint landing session %s: %w", gid, err) + } + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing artifact checkpoint landing: %w", err) + } + return nil +} + +// GetArtifactCheckpointLanding returns one origin's landed sequence and exact +// GID-to-manifest map from a single read snapshot. +func (db *DB) GetArtifactCheckpointLanding( + ctx context.Context, origin string, +) (ArtifactCheckpointLanding, map[string]string, bool, error) { + if origin == "" { + return ArtifactCheckpointLanding{}, nil, false, + errors.New("artifact checkpoint landing origin is required") + } + tx, err := db.getReader().BeginTx(ctx, &sql.TxOptions{ReadOnly: true}) + if err != nil { + return ArtifactCheckpointLanding{}, nil, false, err + } + defer func() { _ = tx.Rollback() }() + landing := ArtifactCheckpointLanding{Origin: origin} + err = tx.QueryRowContext(ctx, ` + SELECT sequence FROM artifact_checkpoint_landings WHERE origin = ?`, origin, + ).Scan(&landing.Sequence) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactCheckpointLanding{}, map[string]string{}, false, nil + } + if err != nil { + return ArtifactCheckpointLanding{}, nil, false, + fmt.Errorf("reading artifact checkpoint landing: %w", err) + } + rows, err := tx.QueryContext(ctx, ` + SELECT gid, manifest_hash FROM artifact_checkpoint_landing_sessions + WHERE origin = ? ORDER BY gid`, origin) + if err != nil { + return ArtifactCheckpointLanding{}, nil, false, + fmt.Errorf("reading artifact checkpoint landing sessions: %w", err) + } + defer rows.Close() + manifestByGID := make(map[string]string) + for rows.Next() { + var gid, manifestHash string + if err := rows.Scan(&gid, &manifestHash); err != nil { + return ArtifactCheckpointLanding{}, nil, false, + fmt.Errorf("scanning artifact checkpoint landing session: %w", err) + } + manifestByGID[gid] = manifestHash + } + if err := rows.Err(); err != nil { + return ArtifactCheckpointLanding{}, nil, false, + fmt.Errorf("iterating artifact checkpoint landing sessions: %w", err) + } + return landing, manifestByGID, true, nil +} + +// StreamArtifactCheckpointLanding visits one origin's exact landing map from +// the same read snapshot as its checkpoint sequence. +func (db *DB) StreamArtifactCheckpointLanding( + ctx context.Context, + origin string, + visit func(gid, manifestHash string) error, +) (ArtifactCheckpointLanding, bool, error) { + if origin == "" || visit == nil { + return ArtifactCheckpointLanding{}, false, + errors.New("artifact checkpoint landing origin and visitor are required") + } + tx, err := db.getReader().BeginTx(ctx, &sql.TxOptions{ReadOnly: true}) + if err != nil { + return ArtifactCheckpointLanding{}, false, err + } + defer func() { _ = tx.Rollback() }() + landing := ArtifactCheckpointLanding{Origin: origin} + err = tx.QueryRowContext(ctx, ` + SELECT sequence FROM artifact_checkpoint_landings WHERE origin = ?`, origin, + ).Scan(&landing.Sequence) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactCheckpointLanding{}, false, nil + } + if err != nil { + return ArtifactCheckpointLanding{}, false, + fmt.Errorf("reading artifact checkpoint landing: %w", err) + } + rows, err := tx.QueryContext(ctx, ` + SELECT gid, manifest_hash FROM artifact_checkpoint_landing_sessions + WHERE origin = ? ORDER BY gid`, origin) + if err != nil { + return ArtifactCheckpointLanding{}, false, + fmt.Errorf("streaming artifact checkpoint landing sessions: %w", err) + } + defer rows.Close() + for rows.Next() { + var gid, manifestHash string + if err := rows.Scan(&gid, &manifestHash); err != nil { + return ArtifactCheckpointLanding{}, false, + fmt.Errorf("scanning artifact checkpoint landing session: %w", err) + } + if err := visit(gid, manifestHash); err != nil { + return ArtifactCheckpointLanding{}, false, err + } + } + if err := rows.Err(); err != nil { + return ArtifactCheckpointLanding{}, false, + fmt.Errorf("iterating artifact checkpoint landing sessions: %w", err) + } + return landing, true, nil +} + +// StreamArtifactPublications visits publication rows in canonical session-id +// order without materializing the full checkpoint map. The returned revision +// and every visited row come from the same SQLite read snapshot. +func (db *DB) StreamArtifactPublications( + ctx context.Context, origin string, visit func(ArtifactPublication) error, +) (int64, error) { + if visit == nil { + return 0, errors.New("artifact publication visitor is required") + } + db.connMu.RLock() + reader := db.reader.Load() + if reader == nil { + db.connMu.RUnlock() + return 0, errors.New("database is closed") + } + tx, err := reader.BeginTx(ctx, &sql.TxOptions{ReadOnly: true}) + db.connMu.RUnlock() + if err != nil { + return 0, fmt.Errorf("beginning artifact publication snapshot: %w", err) + } + defer func() { _ = tx.Rollback() }() + revision, err := artifactPublicationRevisionTx(ctx, tx, origin, false) + if err != nil { + return 0, err + } + rows, err := tx.QueryContext(ctx, ` + SELECT origin, session_id, manifest_hash, source_fingerprint + FROM artifact_publications + WHERE origin = ? + ORDER BY session_id`, origin) + if err != nil { + return 0, fmt.Errorf("streaming artifact publications: %w", err) + } + defer rows.Close() + for rows.Next() { + var publication ArtifactPublication + if err := rows.Scan( + &publication.Origin, &publication.SessionID, + &publication.ManifestHash, &publication.SourceFingerprint, + ); err != nil { + return 0, fmt.Errorf("scanning artifact publication: %w", err) + } + if err := visit(publication); err != nil { + return 0, err + } + } + if err := rows.Err(); err != nil { + return 0, fmt.Errorf("iterating artifact publications: %w", err) + } + if err := rows.Close(); err != nil { + return 0, fmt.Errorf("closing artifact publications: %w", err) + } + if err := tx.Commit(); err != nil { + return 0, fmt.Errorf("committing artifact publication snapshot: %w", err) + } + return revision, nil +} + +// ReserveArtifactCheckpointSequence commits the next sequence in an immediate +// transaction before returning. observedFloor is a stable vault traversal's +// maximum sequence and can only raise, never lower, the retained authority. +func (db *DB) ReserveArtifactCheckpointSequence( + ctx context.Context, origin string, observedFloor int, +) (_ int, retErr error) { + if origin == "" { + return 0, errors.New("artifact checkpoint origin is required") + } + if observedFloor < 0 { + return 0, errors.New("artifact checkpoint observed floor must not be negative") + } + db.mu.Lock() + defer db.mu.Unlock() + conn, err := db.getWriter().Conn(ctx) + if err != nil { + return 0, fmt.Errorf("acquiring artifact checkpoint connection: %w", err) + } + defer conn.Close() + if _, err := conn.ExecContext(ctx, "BEGIN IMMEDIATE"); err != nil { + return 0, fmt.Errorf("beginning artifact checkpoint reservation: %w", err) + } + committed := false + defer func() { + if !committed { + _, rollbackErr := conn.ExecContext(context.WithoutCancel(ctx), "ROLLBACK") + retErr = errors.Join(retErr, rollbackErr) + } + }() + var sequence int + err = conn.QueryRowContext(ctx, ` + INSERT INTO artifact_checkpoint_floors(origin, sequence) + VALUES (?, ? + 1) + ON CONFLICT(origin) DO UPDATE SET + sequence = max(artifact_checkpoint_floors.sequence, ?)+1 + RETURNING sequence`, origin, observedFloor, observedFloor).Scan(&sequence) + if err != nil { + return 0, fmt.Errorf("reserving artifact checkpoint sequence: %w", err) + } + if _, err := conn.ExecContext(context.WithoutCancel(ctx), "COMMIT"); err != nil { + return 0, fmt.Errorf("committing artifact checkpoint reservation: %w", err) + } + committed = true + if err := ctx.Err(); err != nil { + return 0, err + } + return sequence, nil +} + +// GetArtifactCheckpointFloor reports the durable sequence authority, if this +// origin has already been bootstrapped. +func (db *DB) GetArtifactCheckpointFloor( + ctx context.Context, origin string, +) (int, bool, error) { + if origin == "" { + return 0, false, errors.New("artifact checkpoint origin is required") + } + var sequence int + err := db.getReader().QueryRowContext(ctx, ` + SELECT sequence FROM artifact_checkpoint_floors WHERE origin = ?`, origin, + ).Scan(&sequence) + if errors.Is(err, sql.ErrNoRows) { + return 0, false, nil + } + if err != nil { + return 0, false, fmt.Errorf("reading artifact checkpoint floor: %w", err) + } + return sequence, true, nil +} + +// EnqueueArtifactRepair records corrupt physical content for trusted-peer +// repair. A repeated detection refreshes the expected identity and timestamp. +func (db *DB) EnqueueArtifactRepair(ctx context.Context, repair ArtifactRepair) error { + if repair.Origin == "" || repair.Kind == "" || repair.Name == "" || repair.SHA256 == "" || repair.Size < 0 { + return errors.New("complete artifact repair identity is required") + } + db.mu.Lock() + defer db.mu.Unlock() + if _, err := db.getWriter().ExecContext(ctx, ` + INSERT INTO artifact_repair_queue( + origin, kind, name, sha256, size, detected_at + ) VALUES (?, ?, ?, ?, ?, COALESCE( + NULLIF(?, ''), strftime('%Y-%m-%dT%H:%M:%fZ','now') + )) + ON CONFLICT(origin, kind, name) DO UPDATE SET + sha256 = excluded.sha256, + size = excluded.size, + detected_at = excluded.detected_at`, + repair.Origin, repair.Kind, repair.Name, repair.SHA256, repair.Size, + repair.DetectedAt, + ); err != nil { + return fmt.Errorf("enqueueing artifact repair: %w", err) + } + return nil +} + +// PendingArtifactRepairs returns the oldest bounded repair page. +func (db *DB) PendingArtifactRepairs(ctx context.Context, limit int) ([]ArtifactRepair, error) { + if err := validateArtifactQueueLimit(limit); err != nil { + return nil, err + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT origin, kind, name, sha256, size, detected_at + FROM artifact_repair_queue + ORDER BY detected_at, origin, kind, name + LIMIT ?`, limit) + if err != nil { + return nil, fmt.Errorf("reading artifact repair queue: %w", err) + } + defer rows.Close() + items := make([]ArtifactRepair, 0, min(limit, 64)) + for rows.Next() { + var repair ArtifactRepair + if err := rows.Scan( + &repair.Origin, &repair.Kind, &repair.Name, &repair.SHA256, + &repair.Size, &repair.DetectedAt, + ); err != nil { + return nil, fmt.Errorf("scanning artifact repair queue: %w", err) + } + items = append(items, repair) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating artifact repair queue: %w", err) + } + return items, nil +} + +// ArtifactRepairForRef returns the exact queued repair for one logical +// reference without materializing or scanning the queue. +func (db *DB) ArtifactRepairForRef( + ctx context.Context, origin, kind, name string, +) (ArtifactRepair, bool, error) { + if origin == "" || kind == "" || name == "" { + return ArtifactRepair{}, false, errors.New("complete artifact repair reference is required") + } + var repair ArtifactRepair + err := db.getReader().QueryRowContext(ctx, ` + SELECT origin, kind, name, sha256, size, detected_at + FROM artifact_repair_queue + WHERE origin = ? AND kind = ? AND name = ?`, + origin, kind, name, + ).Scan( + &repair.Origin, &repair.Kind, &repair.Name, &repair.SHA256, + &repair.Size, &repair.DetectedAt, + ) + if errors.Is(err, sql.ErrNoRows) { + return ArtifactRepair{}, false, nil + } + if err != nil { + return ArtifactRepair{}, false, fmt.Errorf("reading artifact repair: %w", err) + } + return repair, true, nil +} + +// AcknowledgeArtifactRepair removes only the exact expected identity that was +// repaired. Re-detection of an identical identity may refresh DetectedAt and +// remains the same claim; a different hash or size must be retried. +func (db *DB) AcknowledgeArtifactRepair( + ctx context.Context, repair ArtifactRepair, +) error { + if repair.Origin == "" || repair.Kind == "" || repair.Name == "" || repair.SHA256 == "" || repair.Size < 0 { + return errors.New("complete artifact repair identity is required") + } + db.mu.Lock() + defer db.mu.Unlock() + result, err := db.getWriter().ExecContext(ctx, ` + DELETE FROM artifact_repair_queue + WHERE origin = ? AND kind = ? AND name = ? AND sha256 = ? AND size = ?`, + repair.Origin, repair.Kind, repair.Name, repair.SHA256, repair.Size) + if err != nil { + return fmt.Errorf("acknowledging artifact repair: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("reading artifact repair acknowledgement: %w", err) + } + if rows != 1 { + return fmt.Errorf("%w: %s/%s/%s", ErrArtifactRepairClaimStale, + repair.Origin, repair.Kind, repair.Name) + } + return nil +} + +func validateArtifactQueueLimit(limit int) error { + if limit < 1 || limit > maxArtifactQueuePageSize { + return fmt.Errorf("artifact queue limit must be between 1 and %d", maxArtifactQueuePageSize) + } + return nil +} + +// enqueueArtifactExportTx advances one locally-owned session exactly once for +// a production transaction whose child-row mutation does not also change an +// export-relevant sessions column. +func enqueueArtifactExportTx(tx *sql.Tx, sessionID string) error { + _, err := tx.Exec(` + INSERT INTO artifact_export_queue(session_id) + SELECT id FROM sessions WHERE id = ? AND machine = 'local' + ON CONFLICT(session_id) DO UPDATE SET + enqueued_at = CASE WHEN pending = 0 + THEN strftime('%Y-%m-%dT%H:%M:%fZ','now') ELSE enqueued_at END, + generation = generation + 1, + pending = 1`, sessionID) + if err != nil { + return fmt.Errorf("enqueueing artifact export for %s: %w", sessionID, err) + } + return nil +} + +func artifactExportGenerationTx( + tx *sql.Tx, sessionID string, +) (int64, bool, error) { + var generation int64 + err := tx.QueryRow(` + SELECT generation FROM artifact_export_queue WHERE session_id = ?`, sessionID, + ).Scan(&generation) + if errors.Is(err, sql.ErrNoRows) { + return 0, false, nil + } + if err != nil { + return 0, false, fmt.Errorf("reading artifact export generation for %s: %w", sessionID, err) + } + return generation, true, nil +} diff --git a/internal/db/artifact_publication_test.go b/internal/db/artifact_publication_test.go new file mode 100644 index 000000000..b635bf20e --- /dev/null +++ b/internal/db/artifact_publication_test.go @@ -0,0 +1,1340 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "fmt" + "path/filepath" + "sort" + "strconv" + "strings" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestArtifactPublicationAtomicLifecycle(t *testing.T) { + database := testDB(t) + ctx := t.Context() + require.NoError(t, database.UpsertSession(Session{ + ID: "session-a", Project: "project", Machine: "local", Agent: "claude", + })) + + pending, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Len(t, pending, 1) + require.Equal(t, []ArtifactExportQueueItem{{ + SessionID: "session-a", EnqueuedAt: pending[0].EnqueuedAt, + Generation: pending[0].Generation, + }}, pending) + + // Merely reading work models a failed export: nothing is acknowledged until + // a checkpoint is durably created and its head is recorded. + pendingAgain, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Equal(t, pending, pendingAgain) + + revision, changed, err := database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "session-a", Generation: pending[0].Generation, + ManifestHash: "manifest-a", SourceFingerprint: "source-a", + }}) + require.NoError(t, err) + assert.True(t, changed) + assert.Equal(t, int64(1), revision) + + revision, changed, err = database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "session-a", Generation: pending[0].Generation, + ManifestHash: "manifest-a", SourceFingerprint: "source-a", + }}) + require.NoError(t, err) + assert.False(t, changed, "identical publication state must not force a checkpoint") + + head := ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 7, PublicationRevision: revision, + SessionMapSHA256: "map-hash", CheckpointSHA256: "checkpoint-hash", + } + require.NoError(t, database.RecordArtifactCheckpointHead(ctx, head, pending)) + gotHead, ok, err := database.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, head, gotHead) + pending, err = database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + assert.Empty(t, pending) + + require.NoError(t, database.UpsertSession(Session{ + ID: "session-a", Project: "project-2", Machine: "local", Agent: "claude", + })) + pending, err = database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Len(t, pending, 1) + _, changed, err = database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "session-a", Generation: pending[0].Generation, Delete: true, + }}) + require.NoError(t, err) + assert.True(t, changed) + staleHead, ok, err := database.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, revision, staleHead.PublicationRevision, + "the retained head revision identifies it as stale") + pending, err = database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.NoError(t, database.AcknowledgeArtifactExports(ctx, pending)) + var publications []ArtifactPublication + _, err = database.StreamArtifactPublications(ctx, "desktop-a1b2c3", func(row ArtifactPublication) error { + publications = append(publications, row) + return nil + }) + require.NoError(t, err) + assert.Empty(t, publications) +} + +func TestArtifactCheckpointHeadRejectsStalePublicationRevision(t *testing.T) { + path := filepath.Join(t.TempDir(), "archive.db") + first, err := Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, first.Close()) }) + second, err := Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, second.Close()) }) + + ctx := t.Context() + require.NoError(t, first.UpsertSession(Session{ + ID: "session-a", Project: "project", Machine: "local", Agent: "claude", + })) + claimA, err := first.ArtifactExportClaims(ctx, []string{"session-a"}) + require.NoError(t, err) + require.Len(t, claimA, 1) + revisionA, changed, err := first.ApplyArtifactPublicationChanges( + ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "session-a", Generation: claimA[0].Generation, + ManifestHash: "manifest-a", SourceFingerprint: "source-a", + }}, + ) + require.NoError(t, err) + require.True(t, changed) + streamedRevisionA, err := first.StreamArtifactPublications( + ctx, "desktop-a1b2c3", func(ArtifactPublication) error { return nil }, + ) + require.NoError(t, err) + require.Equal(t, revisionA, streamedRevisionA) + + require.NoError(t, second.UpsertSession(Session{ + ID: "session-b", Project: "project", Machine: "local", Agent: "claude", + })) + claimB, err := second.ArtifactExportClaims(ctx, []string{"session-b"}) + require.NoError(t, err) + require.Len(t, claimB, 1) + revisionB, changed, err := second.ApplyArtifactPublicationChanges( + ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "session-b", Generation: claimB[0].Generation, + ManifestHash: "manifest-b", SourceFingerprint: "source-b", + }}, + ) + require.NoError(t, err) + require.True(t, changed) + require.Greater(t, revisionB, revisionA) + streamedRevisionB, err := second.StreamArtifactPublications( + ctx, "desktop-a1b2c3", func(ArtifactPublication) error { return nil }, + ) + require.NoError(t, err) + require.Equal(t, revisionB, streamedRevisionB) + require.NoError(t, second.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 1, PublicationRevision: revisionB, + SessionMapSHA256: "map-b", CheckpointSHA256: "checkpoint-b", CheckpointSize: 12, + }, claimB)) + + err = first.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 2, PublicationRevision: streamedRevisionA, + SessionMapSHA256: "map-a", CheckpointSHA256: "checkpoint-a", CheckpointSize: 12, + }, claimA) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) + head, ok, err := first.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, 1, head.Sequence) + require.Equal(t, revisionB, head.PublicationRevision) + pending, err := first.ArtifactExportClaims(ctx, []string{"session-a"}) + require.NoError(t, err) + require.Equal(t, claimA, pending, "stale recording cannot consume pending work") + + retryRevision, changed, err := first.ApplyArtifactPublicationChanges( + ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "session-a", Generation: claimA[0].Generation, + ManifestHash: "manifest-a", SourceFingerprint: "source-a", + }}, + ) + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, revisionB, retryRevision) + require.NoError(t, first.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 3, PublicationRevision: retryRevision, + SessionMapSHA256: "map-b", CheckpointSHA256: "checkpoint-c", CheckpointSize: 12, + }, claimA)) +} + +func TestArtifactPublicationQueueBootstrapsExistingLocalSessionsOnOpen(t *testing.T) { + path := filepath.Join(t.TempDir(), "archive.db") + database, err := Open(path) + require.NoError(t, err) + require.NoError(t, database.UpsertSession(Session{ + ID: "existing-local", Project: "project", Machine: "local", Agent: "claude", + })) + require.NoError(t, database.UpsertSession(Session{ + ID: "existing-peer", Project: "project", Machine: "peer-a1b2c3", Agent: "claude", + })) + require.NoError(t, database.UpsertSession(Session{ + ID: "existing-trash", Project: "project", Machine: "local", Agent: "claude", + })) + require.NoError(t, database.SoftDeleteSession("existing-trash")) + require.NoError(t, database.UpsertSession(Session{ + ID: "existing-clean", Project: "project", Machine: "local", Agent: "claude", + })) + cleanClaim, err := database.ArtifactExportClaims(t.Context(), []string{"existing-clean"}) + require.NoError(t, err) + require.Len(t, cleanClaim, 1) + require.NoError(t, database.AcknowledgeArtifactExports(t.Context(), cleanClaim)) + _, err = database.getWriter().Exec( + `DELETE FROM artifact_export_queue WHERE session_id <> 'existing-clean'`, + ) + require.NoError(t, err) + require.NoError(t, database.Close()) + + database, err = Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, database.Close()) }) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "existing-local", pending[0].SessionID) + assert.Equal(t, int64(1), pending[0].Generation) + var cleanPending bool + var cleanGeneration int64 + require.NoError(t, database.getReader().QueryRow(` + SELECT pending, generation FROM artifact_export_queue + WHERE session_id = 'existing-clean'`).Scan(&cleanPending, &cleanGeneration)) + assert.False(t, cleanPending, "idempotent bootstrap cannot requeue clean authority") + assert.Equal(t, cleanClaim[0].Generation, cleanGeneration) + + require.NoError(t, database.Close()) + database, err = Open(path) + require.NoError(t, err) + pending, err = database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "existing-local", pending[0].SessionID) +} + +func TestArtifactExportClaimsSelectsExactPendingSessionSet(t *testing.T) { + database := testDB(t) + for _, id := range []string{"alpha", "bravo", "charlie"} { + require.NoError(t, database.UpsertSession(Session{ + ID: id, Project: "project", Machine: "local", Agent: "claude", + })) + } + + claims, err := database.ArtifactExportClaims(t.Context(), + []string{"charlie", "missing", "alpha", "charlie"}) + require.NoError(t, err) + require.Len(t, claims, 2) + assert.Equal(t, "alpha", claims[0].SessionID) + assert.Equal(t, "charlie", claims[1].SessionID) + assert.Positive(t, claims[0].Generation) + assert.Positive(t, claims[1].Generation) +} + +func TestArtifactPublicationStaleClaimRollsBackEntireBatch(t *testing.T) { + database := testDB(t) + ctx := t.Context() + for _, id := range []string{"alpha", "bravo"} { + require.NoError(t, database.UpsertSession(Session{ + ID: id, Project: "project", Machine: "local", Agent: "claude", + })) + } + claimed, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Len(t, claimed, 2) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'newer' WHERE id = 'bravo'`, + ) + require.NoError(t, err) + _, _, err = database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{ + {SessionID: "alpha", Generation: claimed[0].Generation, ManifestHash: "alpha", SourceFingerprint: "alpha"}, + {SessionID: "bravo", Generation: claimed[1].Generation, ManifestHash: "bravo", SourceFingerprint: "bravo"}, + }) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) + + var publications []ArtifactPublication + _, err = database.StreamArtifactPublications(ctx, "desktop-a1b2c3", func(row ArtifactPublication) error { + publications = append(publications, row) + return nil + }) + require.NoError(t, err) + assert.Empty(t, publications, "a stale claim must not partially mutate publication state") + + fresh, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Len(t, fresh, 2) + _, _, err = database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{ + {SessionID: fresh[0].SessionID, Generation: fresh[0].Generation, ManifestHash: "fresh-0", SourceFingerprint: "fresh-0"}, + {SessionID: fresh[1].SessionID, Generation: fresh[1].Generation, ManifestHash: "fresh-1", SourceFingerprint: "fresh-1"}, + }) + require.NoError(t, err) + publications = nil + _, err = database.StreamArtifactPublications(ctx, "desktop-a1b2c3", func(row ArtifactPublication) error { + publications = append(publications, row) + return nil + }) + require.NoError(t, err) + assert.Len(t, publications, 2, "fresh claims remain applicable after a stale retry") +} + +func TestArtifactCheckpointHeadStaleClaimRollsBackHead(t *testing.T) { + database := testDB(t) + ctx := t.Context() + require.NoError(t, database.UpsertSession(Session{ + ID: "racing", Project: "project", Machine: "local", Agent: "claude", + })) + claimed, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, claimed, 1) + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'newer' WHERE id = 'racing'`, + ) + require.NoError(t, err) + + err = database.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 1, + SessionMapSHA256: "map-1", CheckpointSHA256: "checkpoint-1", + }, claimed) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) + _, ok, err := database.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + assert.False(t, ok, "stale acknowledgement must roll back the checkpoint head") + pending, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Greater(t, pending[0].Generation, claimed[0].Generation) +} + +func TestArtifactPublicationRowsStreamInCanonicalOrder(t *testing.T) { + database := testDB(t) + ctx := t.Context() + for _, id := range []string{"zulu", "alpha"} { + require.NoError(t, database.UpsertSession(Session{ + ID: id, Project: "project", Machine: "local", Agent: "claude", + })) + } + claimed, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + claimGeneration := map[string]int64{} + for _, item := range claimed { + claimGeneration[item.SessionID] = item.Generation + } + _, changed, err := database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{ + {SessionID: "zulu", Generation: claimGeneration["zulu"], ManifestHash: "hash-z", SourceFingerprint: "source-z"}, + {SessionID: "alpha", Generation: claimGeneration["alpha"], ManifestHash: "hash-a", SourceFingerprint: "source-a"}, + }) + require.NoError(t, err) + require.True(t, changed) + + var got []string + _, err = database.StreamArtifactPublications(ctx, "desktop-a1b2c3", func(row ArtifactPublication) error { + got = append(got, row.SessionID+"="+row.ManifestHash) + return nil + }) + require.NoError(t, err) + assert.Equal(t, []string{"alpha=hash-a", "zulu=hash-z"}, got) +} + +func TestArtifactCheckpointHeadRejectsRegressionWithoutAcknowledgingWork(t *testing.T) { + database := testDB(t) + ctx := t.Context() + require.NoError(t, database.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 7, + SessionMapSHA256: "map-7", CheckpointSHA256: "checkpoint-7", + }, nil)) + require.NoError(t, database.UpsertSession(Session{ + ID: "still-pending", Project: "project", Machine: "local", Agent: "claude", + })) + + pendingBefore, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + err = database.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 6, + SessionMapSHA256: "map-6", CheckpointSHA256: "checkpoint-6", + }, pendingBefore) + require.Error(t, err) + + pending, err := database.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "still-pending", pending[0].SessionID) + head, ok, err := database.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 7, head.Sequence) +} + +func TestArtifactPublicationAcknowledgementCannotConsumeNewerMutation(t *testing.T) { + database := testDB(t) + ctx := t.Context() + require.NoError(t, database.UpsertSession(Session{ + ID: "racing", Project: "project", Machine: "local", Agent: "claude", + })) + claimed, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, claimed, 1) + + // Both writes can share the same SQLite millisecond. Generation, not wall + // time, is the compare-and-ack token. + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'newer' WHERE id = 'racing'`, + ) + require.NoError(t, err) + newer, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, newer, 1) + assert.Equal(t, claimed[0].EnqueuedAt, newer[0].EnqueuedAt) + assert.Greater(t, newer[0].Generation, claimed[0].Generation) + + err = database.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 1, + SessionMapSHA256: "map-1", CheckpointSHA256: "checkpoint-1", + }, claimed) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) + pending, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Equal(t, newer, pending, "stale acknowledgment must leave newer work pending") +} + +func TestArtifactPublicationQueueGenerationSurvivesAcknowledgementABA(t *testing.T) { + database := testDB(t) + ctx := t.Context() + require.NoError(t, database.UpsertSession(Session{ + ID: "aba", Project: "project", Machine: "local", Agent: "claude", + })) + oldClaim, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, oldClaim, 1) + _, _, err = database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "aba", Generation: oldClaim[0].Generation, + ManifestHash: "old-manifest", SourceFingerprint: "old-source", + }}) + require.NoError(t, err) + require.NoError(t, database.AcknowledgeArtifactExports(ctx, oldClaim)) + pending, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + assert.Empty(t, pending) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'new mutation' WHERE id = 'aba'`, + ) + require.NoError(t, err) + newClaim, err := database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, newClaim, 1) + _, err = database.getWriter().Exec( + `UPDATE artifact_export_queue SET enqueued_at = ? WHERE session_id = 'aba'`, + oldClaim[0].EnqueuedAt, + ) + require.NoError(t, err) + newClaim, err = database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + require.Len(t, newClaim, 1) + assert.Equal(t, oldClaim[0].EnqueuedAt, newClaim[0].EnqueuedAt, + "the ABA guard cannot depend on timestamp precision") + assert.Greater(t, newClaim[0].Generation, oldClaim[0].Generation) + + revision, _, err := database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "aba", Generation: newClaim[0].Generation, + ManifestHash: "new-manifest", SourceFingerprint: "new-source", + }}) + require.NoError(t, err) + require.NoError(t, database.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 1, PublicationRevision: revision, + SessionMapSHA256: "map-1", CheckpointSHA256: "checkpoint-1", + }, nil)) + + _, _, err = database.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "aba", Generation: oldClaim[0].Generation, + ManifestHash: "stale-manifest", SourceFingerprint: "stale-source", + }}) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) + err = database.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 2, + SessionMapSHA256: "stale-map", CheckpointSHA256: "stale-checkpoint", + }, oldClaim) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) + require.ErrorIs(t, database.AcknowledgeArtifactExports(ctx, oldClaim), ErrArtifactExportClaimStale) + + var publications []ArtifactPublication + _, err = database.StreamArtifactPublications(ctx, "desktop-a1b2c3", func(row ArtifactPublication) error { + publications = append(publications, row) + return nil + }) + require.NoError(t, err) + require.Len(t, publications, 1) + assert.Equal(t, "new-manifest", publications[0].ManifestHash) + head, ok, err := database.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 1, head.Sequence) + pending, err = database.PendingArtifactExports(ctx, 1) + require.NoError(t, err) + assert.Equal(t, newClaim, pending) +} + +func TestArtifactPublicationUsageOnlyAtomicBatchEnqueuesExactlyOnce(t *testing.T) { + for _, machine := range []string{"local", "peer-a1b2c3"} { + t.Run(machine, func(t *testing.T) { + database := testDB(t) + session := Session{ + ID: "usage-batch", Project: "project", Machine: machine, Agent: "claude", + } + require.NoError(t, database.UpsertSession(session)) + if machine == "local" { + claim, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.NoError(t, database.AcknowledgeArtifactExports(t.Context(), claim)) + } + stored, err := database.GetSessionFull(t.Context(), session.ID) + require.NoError(t, err) + require.NotNil(t, stored) + + _, err = database.WriteSessionBatchAtomic([]SessionBatchWrite{{ + Session: *stored, + UsageEvents: []UsageEvent{{ + SessionID: session.ID, Source: "event", Model: "model", + InputTokens: 10, OutputTokens: 2, DedupKey: "changed", + }}, + ReplaceMessages: false, + }}) + require.NoError(t, err) + pending, err := database.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + if machine != "local" { + assert.Empty(t, pending) + return + } + require.Len(t, pending, 1) + assert.Equal(t, int64(2), pending[0].Generation, + "usage-only batch advances the clean authority exactly once") + }) + } +} + +func TestArtifactPublicationRepairQueueIsBoundedAndAcknowledged(t *testing.T) { + database := testDB(t) + ctx := t.Context() + for _, repair := range []ArtifactRepair{ + {Origin: "desktop-a1b2c3", Kind: "manifests", Name: "b.json", SHA256: "hash-b", Size: 20}, + {Origin: "desktop-a1b2c3", Kind: "segments", Name: "a.ndjson", SHA256: "hash-a", Size: 10}, + } { + require.NoError(t, database.EnqueueArtifactRepair(ctx, repair)) + } + + pending, err := database.PendingArtifactRepairs(ctx, 1) + require.NoError(t, err) + require.Len(t, pending, 1) + require.NoError(t, database.AcknowledgeArtifactRepair(ctx, pending[0])) + pending, err = database.PendingArtifactRepairs(ctx, 10) + require.NoError(t, err) + assert.Len(t, pending, 1) +} + +func TestArtifactRepairAcknowledgementUsesFullClaimIdentity(t *testing.T) { + database := testDB(t) + ctx := t.Context() + original := ArtifactRepair{ + Origin: "desktop-a1b2c3", Kind: "segments", Name: "segment.ndjson", + SHA256: "old-hash", Size: 10, + } + require.NoError(t, database.EnqueueArtifactRepair(ctx, original)) + claimed, err := database.PendingArtifactRepairs(ctx, 1) + require.NoError(t, err) + require.Len(t, claimed, 1) + require.NoError(t, database.EnqueueArtifactRepair(ctx, ArtifactRepair{ + Origin: original.Origin, Kind: original.Kind, Name: original.Name, + SHA256: "new-hash", Size: 11, + })) + + err = database.AcknowledgeArtifactRepair(ctx, claimed[0]) + require.ErrorIs(t, err, ErrArtifactRepairClaimStale) + pending, err := database.PendingArtifactRepairs(ctx, 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "new-hash", pending[0].SHA256) + sameIdentityClaim := pending[0] + + // Re-detecting the identical expected identity only refreshes its timestamp; + // the original claim still describes the same repair and may acknowledge it. + require.NoError(t, database.EnqueueArtifactRepair(ctx, ArtifactRepair{ + Origin: original.Origin, Kind: original.Kind, Name: original.Name, + SHA256: "new-hash", Size: 11, DetectedAt: "2030-01-01T00:00:00.000Z", + })) + require.NoError(t, database.AcknowledgeArtifactRepair(ctx, sameIdentityClaim)) +} + +func TestArtifactPeerCheckpointHeadIsMonotonicAndImmutable(t *testing.T) { + database := testDB(t) + ctx := t.Context() + head := ArtifactPeerCheckpointHead{ + Origin: "peer-a1b2c3", Sequence: 2, + CheckpointSHA256: strings.Repeat("a", 64), CheckpointSize: 123, + } + require.NoError(t, database.RecordArtifactPeerCheckpointHead(ctx, head)) + require.NoError(t, database.RecordArtifactPeerCheckpointHead(ctx, head), + "an exact replay is idempotent") + + stale := head + stale.Sequence = 1 + stale.CheckpointSHA256 = strings.Repeat("b", 64) + require.Error(t, database.RecordArtifactPeerCheckpointHead(ctx, stale)) + conflict := head + conflict.CheckpointSHA256 = strings.Repeat("c", 64) + require.Error(t, database.RecordArtifactPeerCheckpointHead(ctx, conflict)) + + got, found, err := database.GetArtifactPeerCheckpointHead(ctx, head.Origin) + require.NoError(t, err) + require.True(t, found) + assert.Equal(t, head, got) +} + +func TestArtifactCheckpointLandingReplacesExactManifestMap(t *testing.T) { + database := testDB(t) + ctx := t.Context() + origin := "peer-a1b2c3" + + recorded := ArtifactCheckpointLanding{Origin: origin, Sequence: 7} + recordedMap := map[string]string{ + origin + "~alpha": "manifest-alpha", + origin + "~bravo": "manifest-bravo", + } + require.NoError(t, database.RecordArtifactCheckpointLanding(ctx, recorded, recordedMap)) + + got, gotMap, ok, err := database.GetArtifactCheckpointLanding(ctx, origin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, recorded, got) + assert.Equal(t, recordedMap, gotMap) + + replacement := ArtifactCheckpointLanding{Origin: origin, Sequence: 8} + replacementMap := map[string]string{ + origin + "~alpha": "manifest-alpha-v2", + } + require.NoError(t, database.RecordArtifactCheckpointLanding(ctx, replacement, replacementMap)) + + got, gotMap, ok, err = database.GetArtifactCheckpointLanding(ctx, origin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, replacement, got) + assert.Equal(t, replacementMap, gotMap, + "removed checkpoint entries must not remain in exact landing provenance") + + stale := ArtifactCheckpointLanding{Origin: origin, Sequence: 7} + err = database.RecordArtifactCheckpointLanding(ctx, stale, map[string]string{ + origin + "~stale": "manifest-stale", + }) + require.Error(t, err) + + got, gotMap, ok, err = database.GetArtifactCheckpointLanding(ctx, origin) + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, replacement, got) + assert.Equal(t, replacementMap, gotMap, + "a lower sequence must preserve the newer exact landing snapshot") +} + +func TestArtifactClaimErrorsRemainDiscoverableWhenWrapped(t *testing.T) { + assert.True(t, errors.Is(fmt.Errorf("retry: %w", ErrArtifactExportClaimStale), ErrArtifactExportClaimStale)) + assert.True(t, errors.Is(fmt.Errorf("retry: %w", ErrArtifactRepairClaimStale), ErrArtifactRepairClaimStale)) +} + +func TestArtifactCheckpointFloorReservationIsConcurrentAndNeverLowers(t *testing.T) { + path := filepath.Join(t.TempDir(), "floor.db") + first, err := Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, first.Close()) }) + second, err := Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, second.Close()) }) + + const reservations = 24 + sequences := make(chan int, reservations) + errs := make(chan error, reservations) + var wg sync.WaitGroup + for i := range reservations { + wg.Add(1) + go func(i int) { + defer wg.Done() + database := first + if i%2 == 1 { + database = second + } + sequence, reserveErr := database.ReserveArtifactCheckpointSequence( + context.Background(), "desktop-a1b2c3", 10, + ) + if reserveErr != nil { + errs <- reserveErr + return + } + sequences <- sequence + }(i) + } + wg.Wait() + close(errs) + for reserveErr := range errs { + require.NoError(t, reserveErr) + } + close(sequences) + var got []int + for sequence := range sequences { + got = append(got, sequence) + } + sort.Ints(got) + want := make([]int, reservations) + for i := range want { + want[i] = 11 + i + } + assert.Equal(t, want, got) + + // Simulate a crash after the committed floor reservation but before the + // checkpoint node is created, followed by a vault reset reporting no live + // sequence. The durable floor consumes the missing sequence permanently. + next, err := first.ReserveArtifactCheckpointSequence(t.Context(), "desktop-a1b2c3", 0) + require.NoError(t, err) + assert.Equal(t, 11+reservations, next) +} + +func TestArtifactPublicationStateSurvivesFullResyncCopy(t *testing.T) { + dir := t.TempDir() + sourcePath := filepath.Join(dir, "source.db") + source, err := Open(sourcePath) + require.NoError(t, err) + + ctx := t.Context() + require.NoError(t, source.UpsertSession(Session{ + ID: "queued", Project: "project", Machine: "local", Agent: "claude", + })) + require.NoError(t, source.UpsertSession(Session{ + ID: "published", Project: "project", Machine: "local", Agent: "claude", + })) + claims, err := source.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + var publishedClaim ArtifactExportQueueItem + for _, claim := range claims { + if claim.SessionID == "published" { + publishedClaim = claim + } + } + require.NotZero(t, publishedClaim.Generation) + revision, _, err := source.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "published", Generation: publishedClaim.Generation, + ManifestHash: "manifest", SourceFingerprint: "source", + }}) + require.NoError(t, err) + require.NoError(t, source.AcknowledgeArtifactExports(ctx, []ArtifactExportQueueItem{publishedClaim})) + require.NoError(t, source.RecordArtifactCheckpointHead(ctx, ArtifactCheckpointHead{ + Origin: "desktop-a1b2c3", Sequence: 4, PublicationRevision: revision, + SessionMapSHA256: "map", CheckpointSHA256: "checkpoint", + }, nil)) + sequence, err := source.ReserveArtifactCheckpointSequence(ctx, "desktop-a1b2c3", 8) + require.NoError(t, err) + require.Equal(t, 9, sequence) + require.NoError(t, source.EnqueueArtifactRepair(ctx, ArtifactRepair{ + Origin: "peer-d4e5f6", Kind: "segments", Name: "segment.ndjson", + SHA256: "segment", Size: 123, + })) + require.NoError(t, source.RecordArtifactCheckpointLanding(ctx, + ArtifactCheckpointLanding{Origin: "peer-d4e5f6", Sequence: 6}, + map[string]string{"peer-d4e5f6~session": "manifest"}, + )) + require.NoError(t, source.RecordArtifactPeerCheckpointHead(ctx, + ArtifactPeerCheckpointHead{ + Origin: "peer-d4e5f6", Sequence: 7, + CheckpointSHA256: strings.Repeat("d", 64), CheckpointSize: 456, + })) + require.NoError(t, source.Close()) + + target, err := Open(filepath.Join(dir, "target.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, target.Close()) }) + require.NoError(t, target.CopySyncStateFrom(sourcePath)) + + pending, err := target.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "queued", pending[0].SessionID) + var publications []ArtifactPublication + streamedRevision, err := target.StreamArtifactPublications(ctx, "desktop-a1b2c3", func(row ArtifactPublication) error { + publications = append(publications, row) + return nil + }) + require.NoError(t, err) + assert.Equal(t, revision, streamedRevision) + require.Len(t, publications, 1) + assert.Equal(t, "published", publications[0].SessionID) + head, ok, err := target.GetArtifactCheckpointHead(ctx, "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 4, head.Sequence) + next, err := target.ReserveArtifactCheckpointSequence(ctx, "desktop-a1b2c3", 0) + require.NoError(t, err) + assert.Equal(t, 10, next) + repairs, err := target.PendingArtifactRepairs(ctx, 10) + require.NoError(t, err) + assert.Len(t, repairs, 1) + landing, manifests, ok, err := target.GetArtifactCheckpointLanding(ctx, "peer-d4e5f6") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 6, landing.Sequence) + assert.Equal(t, map[string]string{"peer-d4e5f6~session": "manifest"}, manifests) + peerHead, ok, err := target.GetArtifactPeerCheckpointHead(ctx, "peer-d4e5f6") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 7, peerHead.Sequence) + assert.Equal(t, strings.Repeat("d", 64), peerHead.CheckpointSHA256) + assert.Equal(t, int64(456), peerHead.CheckpointSize) + + // Clean authority rows survive the swap and copied generations are advanced, + // so an in-flight claim against the source database cannot become valid in + // the replacement archive after the same session is dirtied again. + require.NoError(t, target.UpsertSession(Session{ + ID: "published", Project: "project", Machine: "local", Agent: "claude", + })) + pending, err = target.PendingArtifactExports(ctx, 10) + require.NoError(t, err) + var republished ArtifactExportQueueItem + for _, item := range pending { + if item.SessionID == "published" { + republished = item + } + } + require.Greater(t, republished.Generation, publishedClaim.Generation) + _, _, err = target.ApplyArtifactPublicationChanges(ctx, "desktop-a1b2c3", []ArtifactPublicationChange{{ + SessionID: "published", Generation: publishedClaim.Generation, + ManifestHash: "stale", SourceFingerprint: "stale", + }}) + require.ErrorIs(t, err, ErrArtifactExportClaimStale) +} + +func TestArtifactPublicationStateCopyAcceptsPreRevisionCheckpointHead(t *testing.T) { + dir := t.TempDir() + sourcePath := filepath.Join(dir, "legacy.db") + legacy, err := sql.Open("sqlite3", sourcePath) + require.NoError(t, err) + _, err = legacy.Exec(` + CREATE TABLE artifact_checkpoint_heads ( + origin TEXT PRIMARY KEY, + sequence INTEGER NOT NULL, + session_map_sha256 TEXT NOT NULL, + checkpoint_sha256 TEXT NOT NULL + ); + INSERT INTO artifact_checkpoint_heads VALUES ( + 'desktop-a1b2c3', 4, 'map', 'checkpoint' + );`) + require.NoError(t, err) + require.NoError(t, legacy.Close()) + + target := testDB(t) + require.NoError(t, target.CopySyncStateFrom(sourcePath)) + head, ok, err := target.GetArtifactCheckpointHead(t.Context(), "desktop-a1b2c3") + require.NoError(t, err) + require.True(t, ok) + assert.Equal(t, 4, head.Sequence) + assert.Zero(t, head.PublicationRevision) + assert.Zero(t, head.CheckpointSize) +} + +func TestArtifactPublicationExportCardinalityIsQueueBounded(t *testing.T) { + for _, unrelated := range []int{20, 2000} { + t.Run(strconv.Itoa(unrelated), func(t *testing.T) { + database := testDB(t) + for i := range unrelated { + _, err := database.getWriter().Exec( + `INSERT INTO sessions (id, project, machine, agent) + VALUES (?, 'project', 'peer-a1b2c3', 'claude')`, + "peer-"+strconv.Itoa(i), + ) + require.NoError(t, err) + } + require.NoError(t, database.UpsertSession(Session{ + ID: "dirty", Project: "project", Machine: "local", Agent: "claude", + })) + pending, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "dirty", pending[0].SessionID) + }) + } +} + +func TestArtifactPublicationQueueTracksOnlyLocallyOwnedContent(t *testing.T) { + database := testDB(t) + + require.NoError(t, database.UpsertSession(Session{ + ID: "local-session", Project: "project", Machine: "local", Agent: "claude", + })) + require.Equal(t, []string{"local-session"}, artifactExportQueueIDs(t, database)) + clearArtifactExportQueue(t, database) + + require.NoError(t, database.UpsertSession(Session{ + ID: "peer-session", Project: "project", Machine: "peer-a1b2c3", Agent: "claude", + })) + assert.Empty(t, artifactExportQueueIDs(t, database), "foreign inserts stay out of the local publication queue") + + _, err := database.getWriter().Exec( + `UPDATE sessions SET display_name = 'renamed' WHERE id = 'local-session'`, + ) + require.NoError(t, err) + require.Equal(t, []string{"local-session"}, artifactExportQueueIDs(t, database)) + clearArtifactExportQueue(t, database) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET machine = 'peer-a1b2c3' WHERE id = 'local-session'`, + ) + require.NoError(t, err) + require.Equal(t, []string{"local-session"}, artifactExportQueueIDs(t, database), + "local-to-foreign transition must publish removal") + clearArtifactExportQueue(t, database) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET machine = 'local' WHERE id = 'peer-session'`, + ) + require.NoError(t, err) + require.Equal(t, []string{"peer-session"}, artifactExportQueueIDs(t, database), + "foreign-to-local transition must publish content") + clearArtifactExportQueue(t, database) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'peer rename' WHERE id = 'local-session'`, + ) + require.NoError(t, err) + assert.Empty(t, artifactExportQueueIDs(t, database), "unchanged foreign updates stay out of the queue") +} + +func TestArtifactPublicationQueueKeepsFirstDirtyTimeForFIFO(t *testing.T) { + database := testDB(t) + require.NoError(t, database.UpsertSession(Session{ + ID: "dirty", Project: "project", Machine: "local", Agent: "claude", + })) + const firstDirty = "2026-01-02T03:04:05.000Z" + _, err := database.getWriter().Exec( + `UPDATE artifact_export_queue SET enqueued_at = ? WHERE session_id = 'dirty'`, + firstDirty, + ) + require.NoError(t, err) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'changed again' WHERE id = 'dirty'`, + ) + require.NoError(t, err) + pending, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, firstDirty, pending[0].EnqueuedAt, + "repeated changes retain FIFO position until acknowledgement") + assert.Greater(t, pending[0].Generation, int64(1), + "repeated changes advance the compare-and-ack generation") +} + +func TestArtifactPublicationQueueRefreshesDirtyTimeOnlyAfterAcknowledgement(t *testing.T) { + database := testDB(t) + require.NoError(t, database.UpsertSession(Session{ + ID: "dirty", Project: "project", Machine: "local", Agent: "claude", + })) + const oldDirty = "2020-01-02T03:04:05.000Z" + _, err := database.getWriter().Exec( + `UPDATE artifact_export_queue SET enqueued_at = ? WHERE session_id = 'dirty'`, + oldDirty, + ) + require.NoError(t, err) + claim, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.NoError(t, database.AcknowledgeArtifactExports(t.Context(), claim)) + + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'first clean mutation' WHERE id = 'dirty'`, + ) + require.NoError(t, err) + pending, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.NotEqual(t, oldDirty, pending[0].EnqueuedAt, + "clean-to-pending transition receives a fresh FIFO position") + firstDirty := pending[0].EnqueuedAt + + _, err = database.getWriter().Exec( + `UPDATE sessions SET display_name = 'second pending mutation' WHERE id = 'dirty'`, + ) + require.NoError(t, err) + pending, err = database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, firstDirty, pending[0].EnqueuedAt, + "repeated pending writes retain their FIFO position") +} + +func TestArtifactPublicationQueueTracksMessagesUsageAndCascadeDeletion(t *testing.T) { + database := testDB(t) + for _, session := range []Session{ + {ID: "local-session", Project: "project", Machine: "local", Agent: "claude"}, + {ID: "peer-session", Project: "project", Machine: "peer-a1b2c3", Agent: "claude"}, + } { + require.NoError(t, database.UpsertSession(session)) + } + clearArtifactExportQueue(t, database) + + require.NoError(t, database.InsertMessages([]Message{ + {SessionID: "local-session", Ordinal: 0, Role: "user", Content: "local"}, + {SessionID: "peer-session", Ordinal: 0, Role: "user", Content: "peer"}, + })) + require.Equal(t, []string{"local-session"}, artifactExportQueueIDs(t, database)) + pending, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, int64(2), pending[0].Generation, + "InsertMessages enqueues once for the owning session") + clearArtifactExportQueue(t, database) + + require.NoError(t, database.ReplaceSessionUsageEvents("local-session", []UsageEvent{{ + SessionID: "local-session", Source: "event", Model: "model", DedupKey: "local", + }})) + require.NoError(t, database.ReplaceSessionUsageEvents("peer-session", []UsageEvent{{ + SessionID: "peer-session", Source: "event", Model: "model", DedupKey: "peer", + }})) + require.Equal(t, []string{"local-session"}, artifactExportQueueIDs(t, database)) + pending, err = database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, int64(3), pending[0].Generation, + "usage replacement enqueues once independently of event rows") + clearArtifactExportQueue(t, database) + + _, err = database.getWriter().Exec(`DELETE FROM sessions WHERE id = 'local-session'`) + require.NoError(t, err) + require.Equal(t, []string{"local-session"}, artifactExportQueueIDs(t, database), + "the owner signal must survive child-row cascade ordering") +} + +func TestArtifactPublicationQueueMessageReplacementIsBatchBounded(t *testing.T) { + for _, count := range []int{2, 2000} { + t.Run(strconv.Itoa(count), func(t *testing.T) { + database := testDB(t) + require.NoError(t, database.UpsertSession(Session{ + ID: "session", Project: "project", Machine: "local", Agent: "claude", + })) + clearArtifactExportQueue(t, database) + messages := make([]Message, count) + for i := range messages { + messages[i] = Message{ + SessionID: "session", Ordinal: i, Role: "user", + Content: "message " + strconv.Itoa(i), + } + } + require.NoError(t, database.ReplaceSessionMessages("session", messages)) + pending, err := database.PendingArtifactExports(t.Context(), 1) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, int64(2), pending[0].Generation, + "one transaction advances the queue independently of message count") + }) + } +} + +func TestArtifactPublicationQueueIgnoresSessionBookkeepingUpdates(t *testing.T) { + database := testDB(t) + require.NoError(t, database.UpsertSession(Session{ + ID: "session", Project: "project", Machine: "local", Agent: "claude", + })) + clearArtifactExportQueue(t, database) + _, err := database.getWriter().Exec(` + UPDATE sessions SET + file_path = '/tmp/local', file_size = 42, file_mtime = 43, + next_ordinal = 9, last_entry_uuid = 'uuid', file_inode = 44, + file_device = 45, file_hash = 'hash', local_modified_at = 'now', + last_write_incremental = 1, secrets_rules_version = 'rules', + secret_leak_count = 2, sync_marker = 'marker' + WHERE id = 'session'`) + require.NoError(t, err) + assert.Empty(t, artifactExportQueueIDs(t, database)) +} + +func TestArtifactResetRepublishPendingUsesCompareAndSwapClear(t *testing.T) { + database := testDB(t) + pending := ArtifactResetRepublishPending{ + Version: 1, + RootFingerprint: strings.Repeat("a", 64), + Origin: "desktop-d4e5f6", + Token: strings.Repeat("b", 64), + BaselineHLC: "2026-07-22T120000.000000000Z-00000000000000000000", + } + require.NoError(t, database.SetArtifactResetRepublishPending(t.Context(), pending)) + + got, found, err := database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + require.True(t, found) + assert.Equal(t, pending, got) + + stale := pending + stale.Token = strings.Repeat("c", 64) + cleared, err := database.ClearArtifactResetRepublishPending(t.Context(), stale) + require.NoError(t, err) + assert.False(t, cleared) + _, found, err = database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.True(t, found) + + cleared, err = database.ClearArtifactResetRepublishPending(t.Context(), pending) + require.NoError(t, err) + assert.True(t, cleared) + _, found, err = database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.False(t, found) +} + +func TestArtifactImportQueueIsIdempotentAndRejectsIdentityConflicts(t *testing.T) { + database := testDB(t) + work := ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", + Name: artifactImportMetadataName("a"), + SHA256: strings.Repeat("a", 64), Size: 42, + Reason: "session not landed", RequiredFormatVersion: 1, + } + + require.NoError(t, database.EnqueueArtifactImport(t.Context(), work)) + require.NoError(t, database.EnqueueArtifactImport(t.Context(), work)) + count, oldest, err := database.ArtifactImportQueueStats(t.Context()) + require.NoError(t, err) + assert.Equal(t, 1, count) + assert.NotEmpty(t, oldest) + + conflict := work + conflict.SHA256 = strings.Repeat("b", 64) + require.ErrorContains(t, + database.EnqueueArtifactImport(t.Context(), conflict), "identity") + pending, err := database.PendingArtifactImports(t.Context(), 1, 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, work.SHA256, pending[0].SHA256) +} + +func TestArtifactImportQueueRejectsIncompleteWork(t *testing.T) { + database := testDB(t) + valid := ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", + Name: artifactImportMetadataName("a"), + SHA256: strings.Repeat("a", 64), Size: 1, + Reason: "retry", RequiredFormatVersion: 1, + } + tests := map[string]func(*ArtifactImportWork){ + "origin": func(work *ArtifactImportWork) { work.Origin = "" }, + "kind": func(work *ArtifactImportWork) { work.Kind = "segments" }, + "name": func(work *ArtifactImportWork) { work.Name = "elsewhere.json" }, + "identity": func(work *ArtifactImportWork) { work.SHA256 = "short" }, + "size": func(work *ArtifactImportWork) { work.Size = -1 }, + "reason": func(work *ArtifactImportWork) { work.Reason = "" }, + "format": func(work *ArtifactImportWork) { work.RequiredFormatVersion = 0 }, + "enqueueTime": func(work *ArtifactImportWork) { work.EnqueuedAt = "" }, + } + for name, mutate := range tests { + t.Run(name, func(t *testing.T) { + work := valid + mutate(&work) + if name == "enqueueTime" { + _, err := database.AcknowledgeArtifactImport(t.Context(), work) + assert.Error(t, err) + return + } + assert.Error(t, database.EnqueueArtifactImport(t.Context(), work)) + }) + } + checkpoint := valid + checkpoint.Kind = "checkpoints" + checkpoint.Name = "cp-4.json" + assert.Error(t, database.EnqueueArtifactImport(t.Context(), checkpoint)) + _, err := database.PendingArtifactImports(t.Context(), 0, 1) + assert.Error(t, err) +} + +func TestArtifactImportQueueBoundsFIFOAndFutureFormatEligibility(t *testing.T) { + database := testDB(t) + works := []ArtifactImportWork{ + { + Origin: "peer-a1b2c3", Kind: "meta", Name: artifactImportMetadataName("a"), + SHA256: strings.Repeat("a", 64), Size: 1, Reason: "first", + RequiredFormatVersion: 1, + }, + { + Origin: "peer-a1b2c3", Kind: "meta", Name: artifactImportMetadataName("b"), + SHA256: strings.Repeat("b", 64), Size: 2, Reason: "future", + RequiredFormatVersion: 2, + }, + { + Origin: "peer-a1b2c3", Kind: "meta", Name: artifactImportMetadataName("c"), + SHA256: strings.Repeat("c", 64), Size: 3, Reason: "second", + RequiredFormatVersion: 1, + }, + } + for _, work := range works { + require.NoError(t, database.EnqueueArtifactImport(t.Context(), work)) + } + _, err := database.getWriter().Exec(` + UPDATE artifact_import_queue SET enqueued_at = CASE name + WHEN ? THEN '2026-07-22T01:00:00.000Z' + WHEN ? THEN '2026-07-22T02:00:00.000Z' + ELSE '2026-07-22T03:00:00.000Z' END`, works[0].Name, works[1].Name) + require.NoError(t, err) + + _, err = database.PendingArtifactImports(t.Context(), 1, 0) + assert.Error(t, err) + _, err = database.PendingArtifactImports(t.Context(), 1, 1025) + assert.Error(t, err) + ready, err := database.PendingArtifactImports(t.Context(), 1, 1) + require.NoError(t, err) + require.Len(t, ready, 1) + assert.Equal(t, "first", ready[0].Reason) + ready, err = database.PendingArtifactImports(t.Context(), 1, 10) + require.NoError(t, err) + require.Len(t, ready, 2) + assert.Equal(t, []string{"first", "second"}, []string{ready[0].Reason, ready[1].Reason}) + all, err := database.PendingArtifactImports(t.Context(), 2, 10) + require.NoError(t, err) + require.Len(t, all, 3) + assert.Equal(t, "2026-07-22T01:00:00.000Z", all[0].EnqueuedAt) + count, oldest, err := database.ArtifactImportQueueStats(t.Context()) + require.NoError(t, err) + assert.Equal(t, 3, count) + assert.Equal(t, "2026-07-22T01:00:00.000Z", oldest) +} + +func TestArtifactImportQueueAcknowledgesExactClaimAndPersists(t *testing.T) { + path := filepath.Join(t.TempDir(), "archive.db") + database, err := Open(path) + require.NoError(t, err) + work := ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", + Name: artifactImportMetadataName("d"), + SHA256: strings.Repeat("d", 64), Size: 4, + Reason: "retry", RequiredFormatVersion: 1, + } + require.NoError(t, database.EnqueueArtifactImport(t.Context(), work)) + require.NoError(t, database.Close()) + + database, err = Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, database.Close()) }) + pending, err := database.PendingArtifactImports(t.Context(), 1, 10) + require.NoError(t, err) + require.Len(t, pending, 1) + stale := pending[0] + stale.EnqueuedAt = "2000-01-01T00:00:00.000Z" + acknowledged, err := database.AcknowledgeArtifactImport(t.Context(), stale) + require.NoError(t, err) + assert.False(t, acknowledged) + acknowledged, err = database.AcknowledgeArtifactImport(t.Context(), pending[0]) + require.NoError(t, err) + assert.True(t, acknowledged) + acknowledged, err = database.AcknowledgeArtifactImport(t.Context(), pending[0]) + require.NoError(t, err) + assert.False(t, acknowledged) +} + +func TestArtifactImportQueueAdditiveMigrationPreservesArchive(t *testing.T) { + path := filepath.Join(t.TempDir(), "archive.db") + database, err := Open(path) + require.NoError(t, err) + require.NoError(t, database.UpsertSession(Session{ + ID: "preserved", Project: "project", Machine: "local", Agent: "claude", + })) + require.NoError(t, database.Close()) + + legacy, err := sql.Open("sqlite3", path) + require.NoError(t, err) + _, err = legacy.Exec(`DROP TABLE artifact_import_queue`) + require.NoError(t, err) + require.NoError(t, legacy.Close()) + + database, err = Open(path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, database.Close()) }) + preserved, err := database.GetSession(t.Context(), "preserved") + require.NoError(t, err) + require.NotNil(t, preserved) + require.NoError(t, database.EnqueueArtifactImport(t.Context(), ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", + Name: artifactImportMetadataName("e"), + SHA256: strings.Repeat("e", 64), Size: 5, + Reason: "migration", RequiredFormatVersion: 1, + })) +} + +func TestArtifactImportQueueKeepsNewestCheckpointAndAllMetadata(t *testing.T) { + database := testDB(t) + checkpoint := func(sequence int, digest string) ArtifactImportWork { + return ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "checkpoints", + Name: fmt.Sprintf("cp-%010d.json", sequence), SHA256: strings.Repeat(digest, 64), + Size: int64(sequence), Reason: "dependency incomplete", RequiredFormatVersion: 1, + } + } + metadata := func(digest string) ArtifactImportWork { + return ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", Name: artifactImportMetadataName(digest), + SHA256: strings.Repeat(digest, 64), Size: 1, + Reason: "session not landed", RequiredFormatVersion: 1, + } + } + for _, work := range []ArtifactImportWork{ + checkpoint(4, "4"), metadata("a"), metadata("b"), checkpoint(5, "5"), + checkpoint(4, "4"), + } { + require.NoError(t, database.EnqueueArtifactImport(t.Context(), work)) + } + pending, err := database.PendingArtifactImports(t.Context(), 1, 10) + require.NoError(t, err) + require.Len(t, pending, 3) + var names []string + for _, work := range pending { + names = append(names, work.Name) + } + assert.ElementsMatch(t, []string{ + "cp-0000000005.json", artifactImportMetadataName("a"), + artifactImportMetadataName("b"), + }, names) +} + +func artifactImportMetadataName(digest string) string { + return "2026-07-22T120000.000000000Z-00000000000000000000-" + + strings.Repeat(digest, 64) + ".json" +} + +func artifactExportQueueIDs(t *testing.T, database *DB) []string { + t.Helper() + rows, err := database.getReader().Query( + `SELECT session_id FROM artifact_export_queue + WHERE pending = 1 ORDER BY session_id`, + ) + require.NoError(t, err) + defer rows.Close() + var ids []string + for rows.Next() { + var id string + require.NoError(t, rows.Scan(&id)) + ids = append(ids, id) + } + require.NoError(t, rows.Err()) + return ids +} + +func clearArtifactExportQueue(t *testing.T, database *DB) { + t.Helper() + _, err := database.getWriter().Exec(`UPDATE artifact_export_queue SET pending = 0`) + require.NoError(t, err) +} diff --git a/internal/db/db.go b/internal/db/db.go index 3ee07c3c3..d507fd724 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -1352,6 +1352,17 @@ var readOnlyRequiredTables = []string{ "recall_query_exposures", "recall_extract_generations", "recall_extract_progress", + "artifact_export_queue", + "artifact_import_queue", + "artifact_publications", + "artifact_publication_revisions", + "artifact_checkpoint_heads", + "artifact_checkpoint_floors", + "artifact_checkpoint_landings", + "artifact_checkpoint_landing_sessions", + "artifact_peer_checkpoint_heads", + "artifact_repair_queue", + "metadata_artifact_provenance", } var ( @@ -1657,6 +1668,14 @@ func schemaColumnMigrations() []schemaColumnMigration { "sessions", "deletion_cause", "ALTER TABLE sessions ADD COLUMN deletion_cause TEXT", }, + { + "artifact_checkpoint_heads", "publication_revision", + "ALTER TABLE artifact_checkpoint_heads ADD COLUMN publication_revision INTEGER NOT NULL DEFAULT 0", + }, + { + "artifact_checkpoint_heads", "checkpoint_size", + "ALTER TABLE artifact_checkpoint_heads ADD COLUMN checkpoint_size INTEGER NOT NULL DEFAULT 0", + }, { "messages", "is_system", "ALTER TABLE messages ADD COLUMN is_system INTEGER NOT NULL DEFAULT 0", @@ -2129,6 +2148,12 @@ func (db *DB) migrateColumns() error { if err := applySchemaColumnMigrations(w.QueryRow, w.Exec); err != nil { return err } + if _, err := w.Exec(` + INSERT OR IGNORE INTO artifact_export_queue(session_id) + SELECT id FROM sessions + WHERE machine = 'local' AND deleted_at IS NULL`); err != nil { + return fmt.Errorf("bootstrapping artifact export queue: %w", err) + } if err := installSyncMarkerSchemaLocked(w); err != nil { return err } @@ -2284,6 +2309,47 @@ func (db *DB) migrateColumns() error { return err } + if _, err := w.Exec(` + CREATE TABLE IF NOT EXISTS metadata_applied_events ( + origin TEXT NOT NULL, + order_key TEXT NOT NULL, + artifact_hash TEXT NOT NULL, + applied_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + PRIMARY KEY (origin, order_key) + ); + CREATE TABLE IF NOT EXISTS metadata_replay_state ( + session_gid TEXT NOT NULL, + field TEXT NOT NULL, + order_key TEXT NOT NULL, + hlc TEXT NOT NULL, + artifact_hash TEXT NOT NULL, + origin TEXT NOT NULL, + op TEXT NOT NULL, + value TEXT NOT NULL DEFAULT '', + updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + PRIMARY KEY (session_gid, field) + ); + CREATE TABLE IF NOT EXISTS metadata_conflicts ( + id INTEGER PRIMARY KEY, + session_gid TEXT NOT NULL, + field TEXT NOT NULL, + winning_order_key TEXT NOT NULL, + losing_order_key TEXT NOT NULL, + winning_origin TEXT NOT NULL, + losing_origin TEXT NOT NULL, + winning_op TEXT NOT NULL, + losing_op TEXT NOT NULL, + winning_value TEXT NOT NULL DEFAULT '', + losing_value TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + UNIQUE(session_gid, field, winning_order_key, losing_order_key) + ); + `); err != nil { + return fmt.Errorf( + "creating metadata replay tables: %w", err, + ) + } + if err := db.ensureUsageEventsSchemaLocked(w); err != nil { return err } @@ -2859,6 +2925,7 @@ func (db *DB) backfillMessageTokenCoverageLocked( } defer stmt.Close() + sessions := make(map[string]struct{}) for _, candidate := range candidates { if _, err := stmt.Exec( candidate.hasContext, candidate.hasOutput, candidate.id, @@ -2868,6 +2935,12 @@ func (db *DB) backfillMessageTokenCoverageLocked( candidate.id, err, ) } + sessions[candidate.sessionID] = struct{}{} + } + for sessionID := range sessions { + if err := enqueueArtifactExportTx(tx, sessionID); err != nil { + return 0, err + } } if err := tx.Commit(); err != nil { return 0, fmt.Errorf( @@ -2882,7 +2955,7 @@ func (db *DB) messageTokenCoverageBackfillCandidatesLocked( w *writerHandle, ) ([]messageTokenCoverageBackfillCandidate, error) { rows, err := w.Query( - `SELECT id, token_usage, context_tokens, output_tokens, + `SELECT id, session_id, token_usage, context_tokens, output_tokens, has_context_tokens, has_output_tokens FROM messages WHERE (has_context_tokens = 0 OR has_output_tokens = 0) @@ -2900,11 +2973,12 @@ func (db *DB) messageTokenCoverageBackfillCandidatesLocked( var candidates []messageTokenCoverageBackfillCandidate for rows.Next() { var id int64 + var sessionID string var tokenUsage string var contextTokens, outputTokens int var hasContextTokens, hasOutputTokens bool if err := rows.Scan( - &id, &tokenUsage, &contextTokens, + &id, &sessionID, &tokenUsage, &contextTokens, &outputTokens, &hasContextTokens, &hasOutputTokens, ); err != nil { @@ -2922,6 +2996,7 @@ func (db *DB) messageTokenCoverageBackfillCandidatesLocked( } candidates = append(candidates, messageTokenCoverageBackfillCandidate{ id: id, + sessionID: sessionID, hasContext: hasContext, hasOutput: hasOutput, }) @@ -2934,6 +3009,7 @@ func (db *DB) messageTokenCoverageBackfillCandidatesLocked( type messageTokenCoverageBackfillCandidate struct { id int64 + sessionID string hasContext bool hasOutput bool } @@ -3967,6 +4043,47 @@ func (db *DB) GetSyncState(key string) (string, error) { return value, err } +// SyncStateValues reads the non-empty values for exact sync-state keys. Queries +// are bounded so callers can bulk-resolve state without exceeding SQLite's +// historical variable limit. +func (db *DB) SyncStateValues(keys []string) (map[string]string, error) { + states := map[string]string{} + const batchSize = 900 + for start := 0; start < len(keys); start += batchSize { + end := min(start+batchSize, len(keys)) + batch := keys[start:end] + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(batch)), ",") + args := make([]any, len(batch)) + for i, key := range batch { + args[i] = key + } + rows, err := db.getReader().Query( + `SELECT key, value FROM pg_sync_state + WHERE key IN (`+placeholders+`) AND value <> ''`, + args..., + ) + if err != nil { + return nil, err + } + for rows.Next() { + var key, value string + if err := rows.Scan(&key, &value); err != nil { + rows.Close() + return nil, err + } + states[key] = value + } + if err := rows.Err(); err != nil { + rows.Close() + return nil, err + } + if err := rows.Close(); err != nil { + return nil, err + } + } + return states, nil +} + // SetSyncState writes a value to the pg_sync_state table. func (db *DB) SetSyncState(key, value string) error { db.mu.Lock() diff --git a/internal/db/db_test.go b/internal/db/db_test.go index f01c1a064..0f004cee8 100644 --- a/internal/db/db_test.go +++ b/internal/db/db_test.go @@ -1071,6 +1071,47 @@ func TestInsertMessages_PreservesToolResultEvents(t *testing.T) { assert.Equal(t, "subagent_notification", tc.ResultEvents[1].Source, "result event 1 source") } +func TestSessionSubagentSessionIDs(t *testing.T) { + d := testDB(t) + insertSession(t, d, "s-sub", "proj") + require.NoError(t, d.InsertMessages([]Message{ + { + SessionID: "s-sub", Ordinal: 0, Role: "assistant", + Content: "spawn", HasToolUse: true, + ToolCalls: []ToolCall{ + { + SessionID: "s-sub", ToolName: "Task", Category: "Task", + ToolUseID: "call-1", SubagentSessionID: "sub-a", + ResultEvents: []ToolResultEvent{ + { + ToolUseID: "call-1", SubagentSessionID: "sub-a", + Source: "subagent", Status: "ok", EventIndex: 0, + }, + { + ToolUseID: "call-1", SubagentSessionID: "sub-b", + Source: "subagent", Status: "ok", EventIndex: 1, + }, + }, + }, + { + SessionID: "s-sub", ToolName: "Read", Category: "file", + ToolUseID: "call-2", + }, + }, + }, + }), "InsertMessages") + + ids, err := d.SessionSubagentSessionIDs("s-sub") + require.NoError(t, err, "SessionSubagentSessionIDs") + // "sub-a" appears on both the tool call and a result event (deduped); + // "sub-b" only on a result event; the empty subagent id is excluded. + assert.ElementsMatch(t, []string{"sub-a", "sub-b"}, ids) + + none, err := d.SessionSubagentSessionIDs("missing") + require.NoError(t, err, "SessionSubagentSessionIDs missing") + assert.Empty(t, none) +} + func TestOpenPreservesDataAtCurrentVersion(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "test.db") @@ -5193,12 +5234,17 @@ func TestStarSession(t *testing.T) { assert.Equal(t, []string{"s1"}, ids, "listed = %v, want [s1]", ids) // Unstar. - err = d.UnstarSession("s1") + removed, err := d.UnstarSession("s1") require.NoError(t, err, "UnstarSession") + assert.True(t, removed, "UnstarSession should report removed star") ids, err = d.ListStarredSessionIDs(ctx) require.NoError(t, err, "ListStarredSessionIDs after unstar") assert.Empty(t, ids, "listed after unstar = %v, want []", ids) + removed, err = d.UnstarSession("s1") + require.NoError(t, err, "UnstarSession no-op") + assert.False(t, removed, "UnstarSession should report no-op") + // Star non-existent session returns false (no FK error). ok, err = d.StarSession("nonexistent") require.NoError(t, err, "StarSession nonexistent") @@ -5212,8 +5258,13 @@ func TestBulkStarSessions(t *testing.T) { insertSession(t, d, "s2", "proj") // Bulk star with mix of valid and invalid IDs. - err := d.BulkStarSessions([]string{"s1", "s2", "nonexistent"}) + starred, err := d.BulkStarSessions([]string{"s1", "s2", "nonexistent"}) require.NoError(t, err, "BulkStarSessions") + assert.ElementsMatch(t, []string{"s1", "s2"}, starred, "starred ids returned") + + starred, err = d.BulkStarSessions([]string{"s1", "s2"}) + require.NoError(t, err, "BulkStarSessions already starred") + assert.Empty(t, starred, "already-starred ids should not be returned") ids, err := d.ListStarredSessionIDs(ctx) require.NoError(t, err, "ListStarredSessionIDs") @@ -5560,6 +5611,25 @@ func TestCopySyncStateFrom_OnlyCopiesDurablePGKeys(t *testing.T) { srcDB := testDBAtPath(t, srcPath, "src") require.NoError(t, srcDB.SetSyncState("pg_push_marker_id", "marker-123"), "seed source marker") + require.NoError(t, srcDB.SetSyncState("artifact_origin_id", "laptop-a1b2c3"), + "seed source artifact origin") + require.NoError(t, srcDB.SetSyncState("artifact_metadata_hlc", "hlc-42"), + "seed source artifact hlc") + resetPending := ArtifactResetRepublishPending{ + Version: 1, + RootFingerprint: strings.Repeat("a", 64), + Origin: "laptop-a1b2c3", + Token: strings.Repeat("b", 64), + BaselineHLC: "2026-07-22T12:00:00.000000000Z-0000000000", + } + require.NoError(t, srcDB.SetArtifactResetRepublishPending(t.Context(), resetPending), + "seed source artifact reset marker") + require.NoError(t, + srcDB.SetSyncState("artifact_import:peer-b4c5d6:peer-b4c5d6~sess-1", "hash-imp"), + "seed source import watermark") + require.NoError(t, + srcDB.SetSyncState("artifact_export:laptop-a1b2c3:sess-2", "hash-exp"), + "seed source export watermark") require.NoError(t, srcDB.SetSyncState("last_sync_started_at", "old-start"), "seed source started") require.NoError(t, srcDB.SetSyncState("last_sync_finished_at", "old-finish"), @@ -5581,6 +5651,23 @@ func TestCopySyncStateFrom_OnlyCopiesDurablePGKeys(t *testing.T) { require.NoError(t, err, "GetSyncState pg_push_marker_id") assert.Equal(t, "marker-123", gotMarker) + durableArtifactKeys := map[string]string{ + "artifact_origin_id": "laptop-a1b2c3", + "artifact_metadata_hlc": "hlc-42", + "artifact_import:peer-b4c5d6:peer-b4c5d6~sess-1": "hash-imp", + "artifact_export:laptop-a1b2c3:sess-2": "hash-exp", + } + for key, want := range durableArtifactKeys { + got, err := dstDB.GetSyncState(key) + require.NoError(t, err, "GetSyncState %s", key) + assert.Equal(t, want, got, "artifact key %s must survive the copy", key) + } + gotResetPending, found, err := dstDB.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.True(t, found) + assert.Equal(t, resetPending, gotResetPending, + "artifact reset recovery authority must survive the copy") + gotStarted, err := dstDB.GetSyncState("last_sync_started_at") require.NoError(t, err, "GetSyncState last_sync_started_at") assert.Equal(t, "new-start", gotStarted) @@ -5613,6 +5700,109 @@ func TestCopySyncStateFrom_PropagatesErrors(t *testing.T) { assert.Equal(t, "safe", got) } +func TestCopyMetadataReplayFrom(t *testing.T) { + dir := t.TempDir() + ctx := context.Background() + + srcPath := filepath.Join(dir, "src.db") + srcDB := testDBAtPath(t, srcPath, "src") + + // Seed the LWW register and applied-event markers without touching + // session rows, exactly as local curation bookkeeping does. + winner := MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0000000002", + HLC: "hlc-2", + ArtifactHash: "hash-2", + SessionGID: "laptop-a1b2c3~sess-1", + LocalSessionID: "sess-1", + Field: "display_name", + Op: "rename", + Value: `{"display_name":"kept"}`, + } + res, err := srcDB.RecordLocalMetadataProjection(ctx, winner) + require.NoError(t, err, "record winning projection") + require.True(t, res.Applied, "winning projection applied") + + // A stale peer event loses and records a conflict row. + loser := winner + loser.EventOrigin = "peer-b4c5d6" + loser.OrderKey = "0000000001" + loser.HLC = "hlc-1" + loser.ArtifactHash = "hash-1" + loser.Value = `{"display_name":"stale"}` + res, err = srcDB.RecordLocalMetadataProjection(ctx, loser) + require.NoError(t, err, "record losing projection") + require.True(t, res.Conflict, "losing projection records a conflict") + require.NoError(t, srcDB.Close(), "Close src") + + dstPath := filepath.Join(dir, "dst.db") + dstDB := testDBAtPath(t, dstPath, "dst") + defer dstDB.Close() + + require.NoError(t, dstDB.CopyMetadataReplayFrom(srcPath), + "CopyMetadataReplayFrom") + + for _, ev := range []MetadataProjection{winner, loser} { + applied, err := dstDB.MetadataEventApplied(ctx, ev.EventOrigin, ev.OrderKey) + require.NoError(t, err, "MetadataEventApplied %s/%s", ev.EventOrigin, ev.OrderKey) + assert.True(t, applied, + "applied-event marker %s/%s must survive the copy", + ev.EventOrigin, ev.OrderKey) + } + for _, ev := range []MetadataProjection{winner, loser} { + provenance, err := dstDB.MetadataArtifactProvenanceForSession( + ctx, ev.EventOrigin, ev.SessionGID, ev.Op, + ) + require.NoError(t, err) + require.Len(t, provenance, 1, + "metadata artifact provenance must survive the copy") + assert.Equal(t, ev.OrderKey, provenance[0].OrderKey) + assert.Equal(t, ev.ArtifactHash, provenance[0].ArtifactHash) + } + + op, ok, err := dstDB.MetadataReplayStateOp(ctx, winner.SessionGID, winner.Field) + require.NoError(t, err, "MetadataReplayStateOp") + require.True(t, ok, "replay state must survive the copy") + assert.Equal(t, "rename", op) + + conflicts, err := dstDB.ListMetadataConflicts(ctx, []string{winner.SessionGID}) + require.NoError(t, err, "ListMetadataConflicts") + require.Len(t, conflicts, 1, "conflict row must survive the copy") + assert.Equal(t, "laptop-a1b2c3", conflicts[0].WinningOrigin) + assert.Equal(t, "peer-b4c5d6", conflicts[0].LosingOrigin) +} + +func TestCopyMetadataReplayFrom_NoSourceTables(t *testing.T) { + dir := t.TempDir() + ctx := context.Background() + + // Simulate an older source database that predates the metadata + // replay tables. + srcPath := filepath.Join(dir, "src.db") + srcDB := testDBAtPath(t, srcPath, "src") + for _, table := range []string{ + "metadata_applied_events", + "metadata_replay_state", + "metadata_conflicts", + } { + _, err := srcDB.getWriter().Exec("DROP TABLE " + table) + require.NoError(t, err, "drop %s", table) + } + require.NoError(t, srcDB.Close(), "Close src") + + dstPath := filepath.Join(dir, "dst.db") + dstDB := testDBAtPath(t, dstPath, "dst") + defer dstDB.Close() + + require.NoError(t, dstDB.CopyMetadataReplayFrom(srcPath), + "CopyMetadataReplayFrom with missing source tables") + + applied, err := dstDB.MetadataEventApplied(ctx, "laptop-a1b2c3", "0000000001") + require.NoError(t, err, "MetadataEventApplied") + assert.False(t, applied) +} + func TestCopySessionMetadataFrom(t *testing.T) { dir := t.TempDir() ctx := context.Background() @@ -7380,6 +7570,8 @@ func TestMigration_TerminationStatusColumn(t *testing.T) { // dropping a column referenced by an index. _, err = conn.Exec(`DROP INDEX IF EXISTS idx_sessions_termination_status`) requireNoError(t, err, "drop termination_status index") + _, err = conn.Exec(`DROP TRIGGER IF EXISTS artifact_sessions_update_queue`) + requireNoError(t, err, "drop artifact session update trigger") _, err = conn.Exec(`ALTER TABLE sessions DROP COLUMN termination_status`) requireNoError(t, err, "drop termination_status column") diff --git a/internal/db/insights_test.go b/internal/db/insights_test.go index 7726900b7..0b0478740 100644 --- a/internal/db/insights_test.go +++ b/internal/db/insights_test.go @@ -68,9 +68,7 @@ func TestInsights_CannedMetadataAndCacheLookup(t *testing.T) { if err != nil { t.Fatalf("GetCachedInsight: %v", err) } - if got == nil { - t.Fatal("expected cached insight") - } + require.NotNil(t, got, "expected cached insight") if got.ID != id || got.Kind != "prompt_maturity_review" || got.SchemaVersion != "llm_insight.v1" || got.ProvenanceJSON == "" || got.StructuredJSON == "" { diff --git a/internal/db/legacy_schema_test.go b/internal/db/legacy_schema_test.go index 9e84957fd..b2752b266 100644 --- a/internal/db/legacy_schema_test.go +++ b/internal/db/legacy_schema_test.go @@ -173,6 +173,11 @@ func TestOpenLegacySchemasPreservesArchiveAndRequestsResync(t *testing.T) { d, err := Open(path) require.NoError(t, err) assert.True(t, d.NeedsResync()) + pending, err := d.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "legacy-session", pending[0].SessionID, + "artifact bootstrap runs after the deleted_at migration") session := requireSessionExists(t, d, "legacy-session") assert.Equal(t, "project-a", session.Project) @@ -223,6 +228,10 @@ func TestOpenLegacySchemasPreservesArchiveAndRequestsResync(t *testing.T) { require.NoError(t, err) defer reopened.Close() require.True(t, reopened.NeedsResync()) + pending, err = reopened.PendingArtifactExports(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, "legacy-session", pending[0].SessionID) }) } } diff --git a/internal/db/messages.go b/internal/db/messages.go index 55b938502..b04a868df 100644 --- a/internal/db/messages.go +++ b/internal/db/messages.go @@ -399,6 +399,26 @@ func (db *DB) GetResumeModelCounts( return counts, nil } +// GetMessageForMetadataPin returns only the stable fields needed to publish a +// pin metadata event. It deliberately avoids loading message content, tool +// calls, and tool-result events for an otherwise single-row lookup. +func (db *DB) GetMessageForMetadataPin( + ctx context.Context, sessionID string, messageID int64, +) (*Message, error) { + row := db.getReader().QueryRowContext(ctx, ` + SELECT id, session_id, ordinal, COALESCE(source_uuid, '') + FROM messages + WHERE session_id = ? AND id = ?`, sessionID, messageID) + var msg Message + if err := row.Scan(&msg.ID, &msg.SessionID, &msg.Ordinal, &msg.SourceUUID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, fmt.Errorf("querying message metadata for pin: %w", err) + } + return &msg, nil +} + // EmbeddableUnit is one embedding document: a single embeddable user // message, or a run of contiguous embeddable assistant messages. type EmbeddableUnit struct { @@ -2360,6 +2380,40 @@ func (db *DB) ToolResultEventFingerprintWithTimestampNormalizer( return b.String(), rows.Err() } +// SessionSubagentSessionIDs returns the distinct non-empty subagent_session_id +// values referenced by a session's tool calls and tool result events. Used by +// PG push to decide whether a session whose content otherwise matches PG still +// needs its message rows replaced so subagent links can be re-resolved. +func (db *DB) SessionSubagentSessionIDs(sessionID string) ([]string, error) { + rows, err := db.getReader().Query(` + SELECT DISTINCT subagent_session_id FROM ( + SELECT subagent_session_id FROM tool_calls + WHERE session_id = ? + UNION + SELECT subagent_session_id FROM tool_result_events + WHERE session_id = ? + ) + WHERE subagent_session_id IS NOT NULL AND subagent_session_id != ''`, + sessionID, sessionID, + ) + if err != nil { + return nil, fmt.Errorf( + "querying subagent session ids for %s: %w", sessionID, err, + ) + } + defer rows.Close() + + var ids []string + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, err + } + ids = append(ids, id) + } + return ids, rows.Err() +} + // GetMessageByOrdinal returns a single message by session ID and ordinal. func (db *DB) GetMessageByOrdinal( sessionID string, ordinal int, diff --git a/internal/db/metadata_baseline.go b/internal/db/metadata_baseline.go new file mode 100644 index 000000000..ac2a92fb7 --- /dev/null +++ b/internal/db/metadata_baseline.go @@ -0,0 +1,382 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "fmt" +) + +const metadataBaselinePageSize = 128 + +// MetadataBaselineSnapshot captures existing local user curation that predates +// artifact metadata event recording. +type MetadataBaselineSnapshot struct { + Renames []MetadataBaselineRename + StarredSessionIDs []string + SoftDeletedIDs []string + Pins []MetadataBaselinePin +} + +// MetadataBaselineRename is one session display-name override. +type MetadataBaselineRename struct { + SessionID string + DisplayName *string +} + +// MetadataBaselinePin is one pinned message represented in metadata-event +// coordinates. +type MetadataBaselinePin struct { + SessionID string + SourceUUID string + Ordinal int + Note *string +} + +// VisitMetadataBaselinePages visits current local curation in fixed-size +// keyset pages. Every query cursor is fully read and closed before visit runs, +// so the callback may create artifacts and record replay state on the same DB. +func (db *DB) VisitMetadataBaselinePages( + ctx context.Context, + visit func(MetadataBaselineSnapshot) error, +) error { + if visit == nil { + return errors.New("metadata baseline page visitor is required") + } + for _, visitKind := range []func(context.Context, func(MetadataBaselineSnapshot) error) error{ + db.visitMetadataBaselineRenamePages, + db.visitMetadataBaselineStarPages, + db.visitMetadataBaselineSoftDeletePages, + db.visitMetadataBaselinePinPages, + } { + if err := visitKind(ctx, visit); err != nil { + return err + } + } + return nil +} + +func (db *DB) visitMetadataBaselineRenamePages( + ctx context.Context, visit func(MetadataBaselineSnapshot) error, +) error { + afterSessionID := "" + for { + if err := ctx.Err(); err != nil { + return err + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT id, display_name + FROM sessions + WHERE display_name IS NOT NULL AND id > ? + ORDER BY id + LIMIT ?`, afterSessionID, metadataBaselinePageSize) + if err != nil { + return fmt.Errorf("listing baseline rename page: %w", err) + } + page := MetadataBaselineSnapshot{ + Renames: make([]MetadataBaselineRename, 0, metadataBaselinePageSize), + } + for rows.Next() { + var sessionID, displayName string + if err := rows.Scan(&sessionID, &displayName); err != nil { + rows.Close() + return fmt.Errorf("scanning baseline rename page: %w", err) + } + displayNameCopy := displayName + page.Renames = append(page.Renames, MetadataBaselineRename{ + SessionID: sessionID, DisplayName: &displayNameCopy, + }) + } + if err := closeMetadataBaselineRows(rows, "rename"); err != nil { + return err + } + if len(page.Renames) == 0 { + return nil + } + if err := visit(page); err != nil { + return err + } + afterSessionID = page.Renames[len(page.Renames)-1].SessionID + if len(page.Renames) < metadataBaselinePageSize { + return nil + } + } +} + +func (db *DB) visitMetadataBaselineStarPages( + ctx context.Context, visit func(MetadataBaselineSnapshot) error, +) error { + afterSessionID := "" + for { + if err := ctx.Err(); err != nil { + return err + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT ss.session_id + FROM starred_sessions ss + JOIN sessions s ON s.id = ss.session_id + WHERE ss.session_id > ? + ORDER BY ss.session_id + LIMIT ?`, afterSessionID, metadataBaselinePageSize) + if err != nil { + return fmt.Errorf("listing baseline star page: %w", err) + } + page := MetadataBaselineSnapshot{ + StarredSessionIDs: make([]string, 0, metadataBaselinePageSize), + } + for rows.Next() { + var sessionID string + if err := rows.Scan(&sessionID); err != nil { + rows.Close() + return fmt.Errorf("scanning baseline star page: %w", err) + } + page.StarredSessionIDs = append(page.StarredSessionIDs, sessionID) + } + if err := closeMetadataBaselineRows(rows, "star"); err != nil { + return err + } + if len(page.StarredSessionIDs) == 0 { + return nil + } + if err := visit(page); err != nil { + return err + } + afterSessionID = page.StarredSessionIDs[len(page.StarredSessionIDs)-1] + if len(page.StarredSessionIDs) < metadataBaselinePageSize { + return nil + } + } +} + +func (db *DB) visitMetadataBaselineSoftDeletePages( + ctx context.Context, visit func(MetadataBaselineSnapshot) error, +) error { + afterSessionID := "" + for { + if err := ctx.Err(); err != nil { + return err + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT id + FROM sessions + WHERE deleted_at IS NOT NULL AND id > ? + ORDER BY id + LIMIT ?`, afterSessionID, metadataBaselinePageSize) + if err != nil { + return fmt.Errorf("listing baseline soft-delete page: %w", err) + } + page := MetadataBaselineSnapshot{ + SoftDeletedIDs: make([]string, 0, metadataBaselinePageSize), + } + for rows.Next() { + var sessionID string + if err := rows.Scan(&sessionID); err != nil { + rows.Close() + return fmt.Errorf("scanning baseline soft-delete page: %w", err) + } + page.SoftDeletedIDs = append(page.SoftDeletedIDs, sessionID) + } + if err := closeMetadataBaselineRows(rows, "soft-delete"); err != nil { + return err + } + if len(page.SoftDeletedIDs) == 0 { + return nil + } + if err := visit(page); err != nil { + return err + } + afterSessionID = page.SoftDeletedIDs[len(page.SoftDeletedIDs)-1] + if len(page.SoftDeletedIDs) < metadataBaselinePageSize { + return nil + } + } +} + +type metadataBaselinePinPageRow struct { + pin MetadataBaselinePin + pinID int64 +} + +func (db *DB) visitMetadataBaselinePinPages( + ctx context.Context, visit func(MetadataBaselineSnapshot) error, +) error { + afterSessionID := "" + afterOrdinal := 0 + var afterPinID int64 + for { + if err := ctx.Err(); err != nil { + return err + } + rows, err := db.getReader().QueryContext(ctx, ` + SELECT p.session_id, COALESCE(m.source_uuid, ''), + p.ordinal, p.note, p.id + FROM pinned_messages p + JOIN sessions s ON s.id = p.session_id + JOIN messages m ON m.id = p.message_id AND m.session_id = p.session_id + WHERE (p.session_id, p.ordinal, p.id) > (?, ?, ?) + ORDER BY p.session_id, p.ordinal, p.id + LIMIT ?`, + afterSessionID, afterOrdinal, afterPinID, + metadataBaselinePageSize, + ) + if err != nil { + return fmt.Errorf("listing baseline pin page: %w", err) + } + pageRows := make([]metadataBaselinePinPageRow, 0, metadataBaselinePageSize) + for rows.Next() { + var row metadataBaselinePinPageRow + var note sql.NullString + if err := rows.Scan( + &row.pin.SessionID, &row.pin.SourceUUID, &row.pin.Ordinal, ¬e, &row.pinID, + ); err != nil { + rows.Close() + return fmt.Errorf("scanning baseline pin page: %w", err) + } + if note.Valid { + noteCopy := note.String + row.pin.Note = ¬eCopy + } + pageRows = append(pageRows, row) + } + if err := closeMetadataBaselineRows(rows, "pin"); err != nil { + return err + } + if len(pageRows) == 0 { + return nil + } + page := MetadataBaselineSnapshot{ + Pins: make([]MetadataBaselinePin, len(pageRows)), + } + for index := range pageRows { + page.Pins[index] = pageRows[index].pin + } + if err := visit(page); err != nil { + return err + } + last := pageRows[len(pageRows)-1] + afterSessionID = last.pin.SessionID + afterOrdinal = last.pin.Ordinal + afterPinID = last.pinID + if len(pageRows) < metadataBaselinePageSize { + return nil + } + } +} + +func closeMetadataBaselineRows(rows *sql.Rows, kind string) error { + rowsErr := rows.Err() + closeErr := rows.Close() + if err := errors.Join(rowsErr, closeErr); err != nil { + return fmt.Errorf("iterating baseline %s page: %w", kind, err) + } + return nil +} + +// MetadataBaselineSnapshot returns the current curation rows that need baseline +// metadata events during artifact sync initialization. +func (db *DB) MetadataBaselineSnapshot(ctx context.Context) (MetadataBaselineSnapshot, error) { + var snap MetadataBaselineSnapshot + + // Curation queries do not filter on deleted_at: a session sitting in + // trash at opt-in still baselines its name, star, and pins, so a later + // restore reaches peers with them instead of only the soft delete. + renameRows, err := db.getReader().QueryContext(ctx, ` + SELECT id, display_name + FROM sessions + WHERE display_name IS NOT NULL + ORDER BY id`) + if err != nil { + return snap, fmt.Errorf("listing baseline renames: %w", err) + } + defer renameRows.Close() + for renameRows.Next() { + var sessionID string + var displayName string + if err := renameRows.Scan(&sessionID, &displayName); err != nil { + return snap, fmt.Errorf("scanning baseline rename: %w", err) + } + displayNameCopy := displayName + snap.Renames = append(snap.Renames, MetadataBaselineRename{ + SessionID: sessionID, + DisplayName: &displayNameCopy, + }) + } + if err := renameRows.Err(); err != nil { + return snap, fmt.Errorf("iterating baseline renames: %w", err) + } + + starRows, err := db.getReader().QueryContext(ctx, ` + SELECT ss.session_id + FROM starred_sessions ss + JOIN sessions s + ON s.id = ss.session_id + ORDER BY ss.session_id`) + if err != nil { + return snap, fmt.Errorf("listing baseline stars: %w", err) + } + defer starRows.Close() + for starRows.Next() { + var sessionID string + if err := starRows.Scan(&sessionID); err != nil { + return snap, fmt.Errorf("scanning baseline star: %w", err) + } + snap.StarredSessionIDs = append(snap.StarredSessionIDs, sessionID) + } + if err := starRows.Err(); err != nil { + return snap, fmt.Errorf("iterating baseline stars: %w", err) + } + + deletedRows, err := db.getReader().QueryContext(ctx, ` + SELECT id + FROM sessions + WHERE deleted_at IS NOT NULL + ORDER BY id`) + if err != nil { + return snap, fmt.Errorf("listing baseline soft deletes: %w", err) + } + defer deletedRows.Close() + for deletedRows.Next() { + var sessionID string + if err := deletedRows.Scan(&sessionID); err != nil { + return snap, fmt.Errorf("scanning baseline soft delete: %w", err) + } + snap.SoftDeletedIDs = append(snap.SoftDeletedIDs, sessionID) + } + if err := deletedRows.Err(); err != nil { + return snap, fmt.Errorf("iterating baseline soft deletes: %w", err) + } + + pinRows, err := db.getReader().QueryContext(ctx, ` + SELECT p.session_id, COALESCE(m.source_uuid, ''), + m.ordinal, p.note + FROM pinned_messages p + JOIN sessions s + ON s.id = p.session_id + JOIN messages m + ON m.id = p.message_id + AND m.session_id = p.session_id + ORDER BY p.session_id, m.ordinal, p.id`) + if err != nil { + return snap, fmt.Errorf("listing baseline pins: %w", err) + } + defer pinRows.Close() + for pinRows.Next() { + var pin MetadataBaselinePin + var note sql.NullString + if err := pinRows.Scan( + &pin.SessionID, &pin.SourceUUID, &pin.Ordinal, ¬e, + ); err != nil { + return snap, fmt.Errorf("scanning baseline pin: %w", err) + } + if note.Valid { + noteCopy := note.String + pin.Note = ¬eCopy + } + snap.Pins = append(snap.Pins, pin) + } + if err := pinRows.Err(); err != nil { + return snap, fmt.Errorf("iterating baseline pins: %w", err) + } + + return snap, nil +} diff --git a/internal/db/metadata_baseline_test.go b/internal/db/metadata_baseline_test.go new file mode 100644 index 000000000..1f2e525be --- /dev/null +++ b/internal/db/metadata_baseline_test.go @@ -0,0 +1,188 @@ +package db + +import ( + "context" + "fmt" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestVisitMetadataBaselinePagesBoundsSmallAndLargeCuration(t *testing.T) { + for _, size := range []int{4, 257} { + t.Run(fmt.Sprintf("rows_%d", size), func(t *testing.T) { + database := testDB(t) + seedMetadataBaselineCuration(t, database, size) + counts := MetadataBaselineSnapshot{} + maxPage := 0 + pages := 0 + + err := database.VisitMetadataBaselinePages(t.Context(), func(page MetadataBaselineSnapshot) error { + pages++ + pageRows := len(page.Renames) + len(page.StarredSessionIDs) + + len(page.SoftDeletedIDs) + len(page.Pins) + maxPage = max(maxPage, pageRows) + counts.Renames = append(counts.Renames, page.Renames...) + counts.StarredSessionIDs = append( + counts.StarredSessionIDs, page.StarredSessionIDs..., + ) + counts.SoftDeletedIDs = append(counts.SoftDeletedIDs, page.SoftDeletedIDs...) + counts.Pins = append(counts.Pins, page.Pins...) + return database.SetSyncState(fmt.Sprintf("baseline_page_%d", pages), "visited") + }) + require.NoError(t, err) + assert.Len(t, counts.Renames, size) + assert.Len(t, counts.StarredSessionIDs, size) + assert.Len(t, counts.SoftDeletedIDs, size) + assert.Len(t, counts.Pins, size) + assert.Equal(t, min(size, 128), maxPage, + "retained page size must not grow with total curation cardinality") + assert.Positive(t, pages) + }) + } +} + +func TestVisitMetadataBaselinePagesStopsBetweenPagesOnCancellation(t *testing.T) { + database := testDB(t) + seedMetadataBaselineCuration(t, database, 129) + ctx, cancel := context.WithCancel(t.Context()) + visited := 0 + pages := 0 + + err := database.VisitMetadataBaselinePages(ctx, func(page MetadataBaselineSnapshot) error { + pages++ + visited += len(page.Renames) + len(page.StarredSessionIDs) + + len(page.SoftDeletedIDs) + len(page.Pins) + cancel() + return nil + }) + + require.ErrorIs(t, err, context.Canceled) + assert.Equal(t, 1, pages) + assert.Equal(t, 128, visited) +} + +func TestVisitMetadataBaselinePagesPinsUseCompositeKeysetWithoutOffsetDrift(t *testing.T) { + database := testDB(t) + sessionID := "many-pins" + require.NoError(t, database.UpsertSession(Session{ + ID: sessionID, Project: "project-a", Machine: "local", Agent: "claude", + MessageCount: 257, CreatedAt: "2026-01-01T00:00:00Z", + })) + messages := make([]Message, 257) + for index := range messages { + messages[index] = Message{ + SessionID: sessionID, Ordinal: index, Role: "user", Content: "hello", + ContentLength: 5, SourceUUID: fmt.Sprintf("uuid-%04d", index), + } + } + require.NoError(t, database.InsertMessages(messages)) + messages, err := database.GetAllMessages(t.Context(), sessionID) + require.NoError(t, err) + require.Len(t, messages, 257) + for index := range messages { + note := fmt.Sprintf("note-%04d", index) + _, err := database.PinMessage(sessionID, messages[index].ID, ¬e) + require.NoError(t, err) + } + + var ordinals []int + maxPage := 0 + pinPages := 0 + removedVisitedPin := false + err = database.VisitMetadataBaselinePages(t.Context(), func(page MetadataBaselineSnapshot) error { + if len(page.Pins) == 0 { + return nil + } + pinPages++ + maxPage = max(maxPage, len(page.Pins)) + for _, pin := range page.Pins { + ordinals = append(ordinals, pin.Ordinal) + } + if !removedVisitedPin { + removedVisitedPin = true + return database.UnpinMessage(sessionID, messages[0].ID) + } + return nil + }) + require.NoError(t, err) + require.Len(t, ordinals, 257, + "deleting an already-visited row must not shift a later keyset page") + for index, ordinal := range ordinals { + assert.Equal(t, index, ordinal) + } + assert.Equal(t, 3, pinPages) + assert.Equal(t, 128, maxPage) +} + +// A session that was renamed, starred, and pinned before the machine first +// opted into artifact sync must baseline that curation even while it sits in +// trash: only the soft delete would otherwise publish, and a later restore +// would reach peers without the name, star, or pin. +func TestMetadataBaselineSnapshotIncludesTrashedSessions(t *testing.T) { + d := testDB(t) + ctx := context.Background() + + require.NoError(t, d.UpsertSession(Session{ + ID: "s1", Project: "proj", Machine: "local", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, d.InsertMessages([]Message{{ + SessionID: "s1", Ordinal: 0, Role: "user", Content: "hi", + ContentLength: 2, SourceUUID: "uuid-1", + }})) + name := "Kept name" + require.NoError(t, d.RenameSession("s1", &name)) + starred, err := d.StarSession("s1") + require.NoError(t, err) + require.True(t, starred) + msgs, err := d.GetAllMessages(ctx, "s1") + require.NoError(t, err) + require.Len(t, msgs, 1) + note := "kept pin" + _, err = d.PinMessage("s1", msgs[0].ID, ¬e) + require.NoError(t, err) + require.NoError(t, d.SoftDeleteSession("s1")) + + snap, err := d.MetadataBaselineSnapshot(ctx) + require.NoError(t, err) + + require.Len(t, snap.Renames, 1) + assert.Equal(t, "s1", snap.Renames[0].SessionID) + require.NotNil(t, snap.Renames[0].DisplayName) + assert.Equal(t, name, *snap.Renames[0].DisplayName) + assert.Equal(t, []string{"s1"}, snap.StarredSessionIDs) + assert.Equal(t, []string{"s1"}, snap.SoftDeletedIDs) + require.Len(t, snap.Pins, 1) + assert.Equal(t, "s1", snap.Pins[0].SessionID) + assert.Equal(t, "uuid-1", snap.Pins[0].SourceUUID) +} + +func seedMetadataBaselineCuration(t *testing.T, database *DB, count int) { + t.Helper() + ctx := t.Context() + for index := range count { + sessionID := fmt.Sprintf("session-%04d", index) + require.NoError(t, database.UpsertSession(Session{ + ID: sessionID, Project: "project-a", Machine: "local", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, database.InsertMessages([]Message{{ + SessionID: sessionID, Ordinal: 0, Role: "user", Content: "hello", + ContentLength: 5, SourceUUID: fmt.Sprintf("uuid-%04d", index), + }})) + name := fmt.Sprintf("Curated %04d", index) + require.NoError(t, database.RenameSession(sessionID, &name)) + starred, err := database.StarSession(sessionID) + require.NoError(t, err) + require.True(t, starred) + messages, err := database.GetAllMessages(ctx, sessionID) + require.NoError(t, err) + require.Len(t, messages, 1) + note := fmt.Sprintf("note-%04d", index) + _, err = database.PinMessage(sessionID, messages[0].ID, ¬e) + require.NoError(t, err) + require.NoError(t, database.SoftDeleteSession(sessionID)) + } +} diff --git a/internal/db/metadata_replay.go b/internal/db/metadata_replay.go new file mode 100644 index 000000000..8369f9c09 --- /dev/null +++ b/internal/db/metadata_replay.go @@ -0,0 +1,1030 @@ +package db + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "strings" + "time" +) + +// ErrMetadataTargetUnavailable means a metadata event depends on session or +// message content that is not durable locally yet. +var ErrMetadataTargetUnavailable = errors.New("metadata target unavailable") + +// MetadataPinProjection identifies a pinned message during metadata replay. +type MetadataPinProjection struct { + SourceUUID string `json:"source_uuid,omitempty"` + Ordinal int `json:"ordinal"` + Note *string `json:"note,omitempty"` +} + +// MetadataProjection is one decoded artifact metadata event ready for replay. +type MetadataProjection struct { + EventOrigin string + OrderKey string + HLC string + ArtifactHash string + SessionGID string + LocalSessionID string + Field string + Op string + Value string + DisplayName *string + Pin *MetadataPinProjection +} + +// MetadataApplyResult summarizes how replay handled an event. +type MetadataApplyResult struct { + Applied bool + Skipped bool + Conflict bool + Duplicate bool +} + +// MetadataConflict is a losing metadata value recorded during deterministic +// replay. +type MetadataConflict struct { + ID int64 `json:"id"` + SessionGID string `json:"session_gid"` + Field string `json:"field"` + WinningOrderKey string `json:"winning_order_key"` + LosingOrderKey string `json:"losing_order_key"` + WinningOrigin string `json:"winning_origin"` + LosingOrigin string `json:"losing_origin"` + WinningOp string `json:"winning_op"` + LosingOp string `json:"losing_op"` + WinningValue string `json:"winning_value"` + LosingValue string `json:"losing_value"` + CreatedAt string `json:"created_at"` +} + +type metadataReplayState struct { + OrderKey string + HLC string + ArtifactHash string + Origin string + Op string + Value string +} + +// MetadataEventIdentity identifies one immutable artifact metadata event. +type MetadataEventIdentity struct { + Origin string + OrderKey string +} + +const metadataReplayWinnerPageSize = 128 + +// VisitMetadataReplayWinnersAuthoredBy visits the current per-field replay +// winners whose winning event was authored by origin. Keyset pages are fully +// read and closed before visit runs, so callers may write new canonical events +// without retaining the read cursor or materializing all replay state. +func (db *DB) VisitMetadataReplayWinnersAuthoredBy( + ctx context.Context, + origin string, + visit func(MetadataProjection) error, +) error { + if strings.TrimSpace(origin) == "" { + return errors.New("metadata replay winner origin is required") + } + if visit == nil { + return errors.New("metadata replay winner visitor is required") + } + afterSessionGID, afterField := "", "" + for { + rows, err := db.getReader().QueryContext(ctx, + `SELECT session_gid, field, order_key, hlc, artifact_hash, origin, op, value + FROM metadata_replay_state + WHERE origin = ? + AND (session_gid > ? OR (session_gid = ? AND field > ?)) + ORDER BY session_gid, field + LIMIT ?`, + origin, afterSessionGID, afterSessionGID, afterField, + metadataReplayWinnerPageSize, + ) + if err != nil { + return fmt.Errorf("listing metadata replay winners for %s: %w", origin, err) + } + winners := make([]MetadataProjection, 0, metadataReplayWinnerPageSize) + for rows.Next() { + var sessionGID, field string + var state metadataReplayState + if err := rows.Scan( + &sessionGID, &field, &state.OrderKey, &state.HLC, + &state.ArtifactHash, &state.Origin, &state.Op, &state.Value, + ); err != nil { + rows.Close() + return fmt.Errorf("scanning metadata replay winner: %w", err) + } + winner, err := metadataReplayStateProjection(sessionGID, sessionGID, field, state) + if err != nil { + rows.Close() + return err + } + winners = append(winners, winner) + } + rowsErr := rows.Err() + closeErr := rows.Close() + if err := errors.Join(rowsErr, closeErr); err != nil { + return fmt.Errorf("iterating metadata replay winners: %w", err) + } + if len(winners) == 0 { + return nil + } + for _, winner := range winners { + if err := ctx.Err(); err != nil { + return err + } + if err := visit(winner); err != nil { + return err + } + } + last := winners[len(winners)-1] + afterSessionGID, afterField = last.SessionGID, last.Field + if len(winners) < metadataReplayWinnerPageSize { + return nil + } + } +} + +// MetadataArtifactProvenance is the point-read index for one validated +// metadata artifact. It is recorded independently of visible replay state so +// local bookkeeping repair never has to scan the append-only ledger. +type MetadataArtifactProvenance struct { + Origin string + OrderKey string + ArtifactHash string + SessionGID string + Op string +} + +// RecordMetadataArtifactProvenance idempotently records one immutable event +// identity and rejects a different hash, session, or operation at the same +// origin/order key. +func (db *DB) RecordMetadataArtifactProvenance( + ctx context.Context, provenance MetadataArtifactProvenance, +) error { + if provenance.Origin == "" || provenance.OrderKey == "" || + provenance.ArtifactHash == "" || provenance.SessionGID == "" || provenance.Op == "" { + return errors.New("complete metadata artifact provenance is required") + } + db.mu.Lock() + defer db.mu.Unlock() + result, err := db.getWriter().ExecContext(ctx, ` + INSERT INTO metadata_artifact_provenance( + origin, order_key, artifact_hash, session_gid, op + ) VALUES (?, ?, ?, ?, ?) + ON CONFLICT(origin, order_key) DO UPDATE SET + artifact_hash = excluded.artifact_hash, + session_gid = excluded.session_gid, + op = excluded.op + WHERE metadata_artifact_provenance.artifact_hash = excluded.artifact_hash + AND metadata_artifact_provenance.session_gid = excluded.session_gid + AND metadata_artifact_provenance.op = excluded.op`, + provenance.Origin, provenance.OrderKey, provenance.ArtifactHash, + provenance.SessionGID, provenance.Op, + ) + if err != nil { + return fmt.Errorf("recording metadata artifact provenance: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("reading metadata artifact provenance result: %w", err) + } + if rows != 1 { + return fmt.Errorf("metadata artifact provenance %s/%s conflicts with its immutable identity", + provenance.Origin, provenance.OrderKey) + } + return nil +} + +// MetadataArtifactProvenanceForSession returns only one origin/session's +// indexed events, optionally filtered by operation, in ledger order. +func (db *DB) MetadataArtifactProvenanceForSession( + ctx context.Context, origin, sessionGID string, ops ...string, +) ([]MetadataArtifactProvenance, error) { + if origin == "" || sessionGID == "" { + return nil, errors.New("metadata artifact provenance origin and session are required") + } + query := `SELECT origin, order_key, artifact_hash, session_gid, op + FROM metadata_artifact_provenance WHERE origin = ? AND session_gid = ?` + args := []any{origin, sessionGID} + if len(ops) > 0 { + query += " AND op IN (" + strings.TrimRight(strings.Repeat("?,", len(ops)), ",") + ")" + for _, op := range ops { + args = append(args, op) + } + } + query += " ORDER BY order_key" + rows, err := db.getReader().QueryContext(ctx, query, args...) + if err != nil { + return nil, fmt.Errorf("listing metadata artifact provenance: %w", err) + } + defer rows.Close() + var result []MetadataArtifactProvenance + for rows.Next() { + var provenance MetadataArtifactProvenance + if err := rows.Scan(&provenance.Origin, &provenance.OrderKey, + &provenance.ArtifactHash, &provenance.SessionGID, &provenance.Op); err != nil { + return nil, fmt.Errorf("scanning metadata artifact provenance: %w", err) + } + result = append(result, provenance) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating metadata artifact provenance: %w", err) + } + return result, nil +} + +// MetadataAppliedEventIdentities bulk-loads the events already durably handled +// by metadata replay. +func (db *DB) MetadataAppliedEventIdentities( + ctx context.Context, +) (map[MetadataEventIdentity]struct{}, error) { + rows, err := db.getReader().QueryContext(ctx, + `SELECT origin, order_key FROM metadata_applied_events`, + ) + if err != nil { + return nil, fmt.Errorf("listing applied metadata events: %w", err) + } + defer rows.Close() + + identities := make(map[MetadataEventIdentity]struct{}) + for rows.Next() { + var identity MetadataEventIdentity + if err := rows.Scan(&identity.Origin, &identity.OrderKey); err != nil { + return nil, fmt.Errorf("scanning applied metadata event: %w", err) + } + identities[identity] = struct{}{} + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating applied metadata events: %w", err) + } + return identities, nil +} + +// MetadataEventApplied reports whether an artifact metadata event was already +// durably handled. +func (db *DB) MetadataEventApplied(ctx context.Context, origin, orderKey string) (bool, error) { + var exists int + err := db.getReader().QueryRowContext(ctx, + `SELECT 1 FROM metadata_applied_events + WHERE origin = ? AND order_key = ?`, + origin, orderKey, + ).Scan(&exists) + if err == sql.ErrNoRows { + return false, nil + } + if err != nil { + return false, fmt.Errorf("checking metadata event %s/%s: %w", origin, orderKey, err) + } + return true, nil +} + +// ListMetadataConflicts returns conflict rows for one or more global session +// identifiers. +func (db *DB) ListMetadataConflicts( + ctx context.Context, + sessionGIDs []string, +) ([]MetadataConflict, error) { + ids := uniqueNonEmptyStrings(sessionGIDs) + if len(ids) == 0 { + return []MetadataConflict{}, nil + } + placeholders := strings.TrimRight(strings.Repeat("?,", len(ids)), ",") + args := make([]any, len(ids)) + for i, id := range ids { + args[i] = id + } + rows, err := db.getReader().QueryContext(ctx, + `SELECT id, session_gid, field, winning_order_key, losing_order_key, + winning_origin, losing_origin, winning_op, losing_op, + winning_value, losing_value, created_at + FROM metadata_conflicts + WHERE session_gid IN (`+placeholders+`) + AND winning_origin <> losing_origin + ORDER BY created_at DESC, id DESC`, + args..., + ) + if err != nil { + return nil, fmt.Errorf("listing metadata conflicts: %w", err) + } + defer rows.Close() + + conflicts := []MetadataConflict{} + for rows.Next() { + var c MetadataConflict + if err := rows.Scan( + &c.ID, &c.SessionGID, &c.Field, + &c.WinningOrderKey, &c.LosingOrderKey, + &c.WinningOrigin, &c.LosingOrigin, + &c.WinningOp, &c.LosingOp, + &c.WinningValue, &c.LosingValue, + &c.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("scanning metadata conflict: %w", err) + } + conflicts = append(conflicts, c) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating metadata conflicts: %w", err) + } + return conflicts, nil +} + +// CountMetadataConflicts returns the total number of recorded metadata +// conflicts across all sessions. +func (db *DB) CountMetadataConflicts(ctx context.Context) (int, error) { + var count int + err := db.getReader().QueryRowContext(ctx, + `SELECT COUNT(*) FROM metadata_conflicts + WHERE winning_origin <> losing_origin`, + ).Scan(&count) + if err != nil { + return 0, fmt.Errorf("counting metadata conflicts: %w", err) + } + return count, nil +} + +// MarkMetadataEventApplied records a metadata event that was intentionally +// skipped, such as an unknown future op. +func (db *DB) MarkMetadataEventApplied(ctx context.Context, origin, orderKey, hash string) error { + db.mu.Lock() + defer db.mu.Unlock() + _, err := db.getWriter().ExecContext(ctx, + `INSERT OR IGNORE INTO metadata_applied_events + (origin, order_key, artifact_hash) + VALUES (?, ?, ?)`, + origin, orderKey, hash, + ) + if err != nil { + return fmt.Errorf("marking metadata event %s/%s applied: %w", origin, orderKey, err) + } + return nil +} + +// ApplyMetadataProjection applies one known metadata event if it wins the +// per-field LWW register, recording conflicts and the applied-event marker in +// the same transaction. +func (db *DB) ApplyMetadataProjection( + ctx context.Context, + ev MetadataProjection, +) (MetadataApplyResult, error) { + return db.applyMetadataProjection(ctx, ev, true) +} + +// RecordLocalMetadataProjection records the LWW register, conflict rows, and +// applied-event marker for a locally-originated metadata event whose session +// mutation the caller has already applied. It runs the same per-field LWW +// bookkeeping as replay but does not re-apply the mutation, so a later peer +// event with a lower order key cannot silently overwrite a newer local edit. +func (db *DB) RecordLocalMetadataProjection( + ctx context.Context, + ev MetadataProjection, +) (MetadataApplyResult, error) { + return db.applyMetadataProjection(ctx, ev, false) +} + +// MetadataReplayStateOp returns the current LWW operation recorded for a +// metadata field. +func (db *DB) MetadataReplayStateOp( + ctx context.Context, + sessionGID string, + field string, +) (string, bool, error) { + var op string + err := db.getReader().QueryRowContext(ctx, + `SELECT op FROM metadata_replay_state + WHERE session_gid = ? AND field = ?`, + sessionGID, field, + ).Scan(&op) + if err == sql.ErrNoRows { + return "", false, nil + } + if err != nil { + return "", false, fmt.Errorf("reading metadata replay state: %w", err) + } + return op, true, nil +} + +// ReapplyMetadataReplayState reapplies the current visible metadata projection +// for a session from the durable replay register. It does not alter LWW state or +// applied-event markers; it only repairs fields that content import may have +// overwritten or invalidated while replacing session/message rows. +func (db *DB) ReapplyMetadataReplayState( + ctx context.Context, + sessionGID string, + localSessionID string, +) (int, error) { + if err := db.requireWritable(); err != nil { + return 0, err + } + if strings.TrimSpace(sessionGID) == "" || strings.TrimSpace(localSessionID) == "" { + return 0, errors.New("metadata replay session id is required") + } + + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return 0, fmt.Errorf("begin metadata reapply tx: %w", err) + } + defer func() { _ = tx.Rollback() }() + + rows, err := tx.QueryContext(ctx, + `SELECT field, order_key, hlc, artifact_hash, origin, op, value + FROM metadata_replay_state + WHERE session_gid = ? + ORDER BY field`, + sessionGID, + ) + if err != nil { + return 0, fmt.Errorf("reading metadata replay state: %w", err) + } + projections := []MetadataProjection{} + for rows.Next() { + var field string + var state metadataReplayState + if err := rows.Scan( + &field, &state.OrderKey, &state.HLC, &state.ArtifactHash, + &state.Origin, &state.Op, &state.Value, + ); err != nil { + rows.Close() + return 0, fmt.Errorf("scanning metadata replay state: %w", err) + } + ev, err := metadataReplayStateProjection(sessionGID, localSessionID, field, state) + if err != nil { + rows.Close() + return 0, err + } + projections = append(projections, ev) + } + if err := rows.Close(); err != nil { + return 0, fmt.Errorf("closing metadata replay state rows: %w", err) + } + if err := rows.Err(); err != nil { + return 0, fmt.Errorf("iterating metadata replay state: %w", err) + } + + applied := 0 + for _, ev := range projections { + if err := ctx.Err(); err != nil { + return applied, err + } + if err := applyMetadataProjectionTx(ctx, tx, ev); err != nil { + if errors.Is(err, ErrMetadataTargetUnavailable) { + continue + } + return applied, fmt.Errorf("reapplying metadata replay state: %w", err) + } + applied++ + } + if err := tx.Commit(); err != nil { + return applied, fmt.Errorf("commit metadata reapply tx: %w", err) + } + return applied, nil +} + +func recordMetadataArtifactProvenanceTx( + ctx context.Context, tx *sql.Tx, provenance MetadataArtifactProvenance, +) error { + result, err := tx.ExecContext(ctx, ` + INSERT INTO metadata_artifact_provenance( + origin, order_key, artifact_hash, session_gid, op + ) VALUES (?, ?, ?, ?, ?) + ON CONFLICT(origin, order_key) DO UPDATE SET + artifact_hash = excluded.artifact_hash, + session_gid = excluded.session_gid, + op = excluded.op + WHERE metadata_artifact_provenance.artifact_hash = excluded.artifact_hash + AND metadata_artifact_provenance.session_gid = excluded.session_gid + AND metadata_artifact_provenance.op = excluded.op`, + provenance.Origin, provenance.OrderKey, provenance.ArtifactHash, + provenance.SessionGID, provenance.Op, + ) + if err != nil { + return fmt.Errorf("recording metadata artifact provenance: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("reading metadata artifact provenance result: %w", err) + } + if rows != 1 { + return fmt.Errorf("metadata artifact provenance %s/%s conflicts with its immutable identity", + provenance.Origin, provenance.OrderKey) + } + return nil +} + +func (db *DB) applyMetadataProjection( + ctx context.Context, + ev MetadataProjection, + applyMutation bool, +) (MetadataApplyResult, error) { + if ev.EventOrigin == "" || ev.OrderKey == "" || ev.ArtifactHash == "" { + return MetadataApplyResult{}, errors.New("metadata projection event identity is required") + } + if ev.SessionGID == "" || ev.LocalSessionID == "" || ev.Field == "" || ev.Op == "" { + return MetadataApplyResult{}, errors.New("metadata projection target is required") + } + + db.mu.Lock() + defer db.mu.Unlock() + tx, err := db.getWriter().BeginTx(ctx, nil) + if err != nil { + return MetadataApplyResult{}, fmt.Errorf("begin metadata replay tx: %w", err) + } + defer func() { _ = tx.Rollback() }() + if err := recordMetadataArtifactProvenanceTx(ctx, tx, MetadataArtifactProvenance{ + Origin: ev.EventOrigin, OrderKey: ev.OrderKey, ArtifactHash: ev.ArtifactHash, + SessionGID: ev.SessionGID, Op: ev.Op, + }); err != nil { + return MetadataApplyResult{}, err + } + + already, err := metadataEventAppliedTx(ctx, tx, ev.EventOrigin, ev.OrderKey) + if err != nil { + return MetadataApplyResult{}, err + } + if already { + if err := tx.Commit(); err != nil { + return MetadataApplyResult{}, fmt.Errorf("commit metadata replay duplicate: %w", err) + } + return MetadataApplyResult{Skipped: true, Duplicate: true}, nil + } + + current, hasCurrent, err := metadataReplayStateTx(ctx, tx, ev.SessionGID, ev.Field) + if err != nil { + return MetadataApplyResult{}, err + } + result := MetadataApplyResult{} + if hasCurrent && ev.OrderKey <= current.OrderKey { + if metadataStateDiffers(current.Op, current.Value, ev.Op, ev.Value) && + metadataConflictOriginsDiffer(current.Origin, ev.EventOrigin) { + if err := insertMetadataConflictTx(ctx, tx, metadataConflict{ + sessionGID: ev.SessionGID, + field: ev.Field, + winningOrderKey: current.OrderKey, + losingOrderKey: ev.OrderKey, + winningOrigin: current.Origin, + losingOrigin: ev.EventOrigin, + winningOp: current.Op, + losingOp: ev.Op, + winningValue: current.Value, + losingValue: ev.Value, + }); err != nil { + return MetadataApplyResult{}, err + } + result.Conflict = true + } + if err := markMetadataEventAppliedTx(ctx, tx, ev.EventOrigin, ev.OrderKey, ev.ArtifactHash); err != nil { + return MetadataApplyResult{}, err + } + if err := tx.Commit(); err != nil { + return MetadataApplyResult{}, fmt.Errorf("commit metadata replay loser: %w", err) + } + result.Skipped = true + return result, nil + } + + if hasCurrent && metadataStateDiffers(current.Op, current.Value, ev.Op, ev.Value) && + metadataConflictOriginsDiffer(ev.EventOrigin, current.Origin) { + if err := insertMetadataConflictTx(ctx, tx, metadataConflict{ + sessionGID: ev.SessionGID, + field: ev.Field, + winningOrderKey: ev.OrderKey, + losingOrderKey: current.OrderKey, + winningOrigin: ev.EventOrigin, + losingOrigin: current.Origin, + winningOp: ev.Op, + losingOp: current.Op, + winningValue: ev.Value, + losingValue: current.Value, + }); err != nil { + return MetadataApplyResult{}, err + } + result.Conflict = true + } + if applyMutation { + if err := applyMetadataProjectionTx(ctx, tx, ev); err != nil { + return MetadataApplyResult{}, err + } + } + if err := upsertMetadataReplayStateTx(ctx, tx, ev); err != nil { + return MetadataApplyResult{}, err + } + if err := markMetadataEventAppliedTx(ctx, tx, ev.EventOrigin, ev.OrderKey, ev.ArtifactHash); err != nil { + return MetadataApplyResult{}, err + } + if err := tx.Commit(); err != nil { + return MetadataApplyResult{}, fmt.Errorf("commit metadata replay: %w", err) + } + result.Applied = true + return result, nil +} + +func metadataEventAppliedTx(ctx context.Context, tx *sql.Tx, origin, orderKey string) (bool, error) { + var exists int + err := tx.QueryRowContext(ctx, + `SELECT 1 FROM metadata_applied_events + WHERE origin = ? AND order_key = ?`, + origin, orderKey, + ).Scan(&exists) + if err == sql.ErrNoRows { + return false, nil + } + if err != nil { + return false, fmt.Errorf("checking metadata event %s/%s: %w", origin, orderKey, err) + } + return true, nil +} + +func metadataReplayStateTx( + ctx context.Context, + tx *sql.Tx, + sessionGID, field string, +) (metadataReplayState, bool, error) { + var state metadataReplayState + err := tx.QueryRowContext(ctx, + `SELECT order_key, hlc, artifact_hash, origin, op, value + FROM metadata_replay_state + WHERE session_gid = ? AND field = ?`, + sessionGID, field, + ).Scan( + &state.OrderKey, &state.HLC, &state.ArtifactHash, + &state.Origin, &state.Op, &state.Value, + ) + if err == sql.ErrNoRows { + return metadataReplayState{}, false, nil + } + if err != nil { + return metadataReplayState{}, false, fmt.Errorf("reading metadata replay state: %w", err) + } + return state, true, nil +} + +func metadataReplayStateProjection( + sessionGID string, + localSessionID string, + field string, + state metadataReplayState, +) (MetadataProjection, error) { + ev := MetadataProjection{ + EventOrigin: state.Origin, + OrderKey: state.OrderKey, + HLC: state.HLC, + ArtifactHash: state.ArtifactHash, + SessionGID: sessionGID, + LocalSessionID: localSessionID, + Field: field, + Op: state.Op, + Value: state.Value, + } + switch state.Op { + case "rename": + var payload struct { + DisplayName *string `json:"display_name"` + } + if err := json.Unmarshal([]byte(state.Value), &payload); err != nil { + return MetadataProjection{}, fmt.Errorf("decoding rename metadata replay state: %w", err) + } + ev.DisplayName = payload.DisplayName + case "pin", "unpin": + var pin MetadataPinProjection + if err := json.Unmarshal([]byte(state.Value), &pin); err != nil { + return MetadataProjection{}, fmt.Errorf("decoding pin metadata replay state: %w", err) + } + ev.Pin = &pin + } + return ev, nil +} + +// metadataEventWallTime renders a replayed event's HLC wall-clock portion in +// the sessions.deleted_at column format, so trash retention is anchored to +// when the deletion happened on the authoring machine rather than when each +// peer imported the event. The HLC shape ("-<20-digit logical>", wall +// layout without ":" separators) is pinned by the artifact format contract in +// internal/artifact. Returns "" when the HLC has no parseable wall portion. +func metadataEventWallTime(hlc string) string { + idx := strings.LastIndex(hlc, "-") + if idx <= 0 { + return "" + } + wall, err := time.Parse("2006-01-02T150405.000000000Z", hlc[:idx]) + if err != nil { + return "" + } + return wall.UTC().Format("2006-01-02T15:04:05.000Z") +} + +func applyMetadataProjectionTx(ctx context.Context, tx *sql.Tx, ev MetadataProjection) error { + switch ev.Op { + case "rename": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + _, err := tx.ExecContext(ctx, + `UPDATE sessions + SET display_name = ?, + local_modified_at = strftime('%Y-%m-%dT%H:%M:%fZ','now') + WHERE id = ?`, + ev.DisplayName, ev.LocalSessionID, + ) + return err + case "soft_delete": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + _, err := tx.ExecContext(ctx, + `UPDATE sessions + SET deleted_at = COALESCE(deleted_at, NULLIF(?, ''), strftime('%Y-%m-%dT%H:%M:%fZ','now')), + local_modified_at = strftime('%Y-%m-%dT%H:%M:%fZ','now') + WHERE id = ?`, + metadataEventWallTime(ev.HLC), ev.LocalSessionID, + ) + return err + case "restore": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + _, err := tx.ExecContext(ctx, + `UPDATE sessions + SET deleted_at = NULL, + local_modified_at = strftime('%Y-%m-%dT%H:%M:%fZ','now') + WHERE id = ?`, + ev.LocalSessionID, + ) + return err + case "star": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + _, err := tx.ExecContext(ctx, + `INSERT OR IGNORE INTO starred_sessions (session_id) + VALUES (?)`, + ev.LocalSessionID, + ) + return err + case "unstar": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + _, err := tx.ExecContext(ctx, + `DELETE FROM starred_sessions WHERE session_id = ?`, + ev.LocalSessionID, + ) + return err + case "pin": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + if ev.Pin == nil { + return errors.New("pin metadata event missing pin payload") + } + return applyMetadataPinTx(ctx, tx, ev.LocalSessionID, *ev.Pin) + case "unpin": + if err := requireMetadataSessionTx(ctx, tx, ev.LocalSessionID); err != nil { + return err + } + if ev.Pin == nil { + return errors.New("unpin metadata event missing pin payload") + } + return unpinMetadataTx(ctx, tx, ev.LocalSessionID, *ev.Pin) + case "purge": + return applyMetadataPurgeTx(ctx, tx, ev.LocalSessionID) + default: + return fmt.Errorf("unsupported metadata op %q", ev.Op) + } +} + +func applyMetadataPinTx(ctx context.Context, tx *sql.Tx, sessionID string, pin MetadataPinProjection) error { + msg, ok, err := metadataPinTargetTx(ctx, tx, sessionID, pin) + if err != nil { + return err + } + if !ok { + return fmt.Errorf("%w: pin target %s ordinal %d", + ErrMetadataTargetUnavailable, sessionID, pin.Ordinal) + } + _, err = tx.ExecContext(ctx, + `INSERT INTO pinned_messages (session_id, message_id, ordinal, note) + VALUES (?, ?, ?, ?) + ON CONFLICT(session_id, message_id) DO UPDATE SET note = excluded.note`, + sessionID, msg.id, msg.ordinal, pin.Note, + ) + return err +} + +func applyMetadataPurgeTx(ctx context.Context, tx *sql.Tx, sessionID string) error { + aliasIDs, err := sessionAliasIDsTx(tx, "id = ?", sessionID) + if err != nil { + return err + } + if _, err := tx.ExecContext(ctx, + `INSERT OR IGNORE INTO excluded_sessions (id) VALUES (?)`, + sessionID, + ); err != nil { + return err + } + for _, aliasID := range aliasIDs { + if err := excludeSessionIDTx(tx, aliasID); err != nil { + return fmt.Errorf("excluding metadata purge alias %s: %w", aliasID, err) + } + } + _, err = tx.ExecContext(ctx, + `DELETE FROM sessions WHERE id = ?`, + sessionID, + ) + return err +} + +func requireMetadataSessionTx(ctx context.Context, tx *sql.Tx, id string) error { + var exists int + err := tx.QueryRowContext(ctx, + `SELECT 1 FROM sessions WHERE id = ?`, + id, + ).Scan(&exists) + if err == sql.ErrNoRows { + return fmt.Errorf("%w: session %s", ErrMetadataTargetUnavailable, id) + } + if err != nil { + return fmt.Errorf("checking metadata session %s: %w", id, err) + } + return nil +} + +type metadataPinTarget struct { + id int64 + ordinal int +} + +func metadataPinTargetTx( + ctx context.Context, + tx *sql.Tx, + sessionID string, + pin MetadataPinProjection, +) (metadataPinTarget, bool, error) { + if pin.SourceUUID != "" { + target, ok, err := metadataPinTargetByQueryTx(ctx, tx, + `SELECT id, ordinal FROM messages + WHERE session_id = ? AND source_uuid = ? + ORDER BY ordinal LIMIT 1`, + sessionID, pin.SourceUUID, + ) + if err != nil || ok { + return target, ok, err + } + } + return metadataPinTargetByQueryTx(ctx, tx, + `SELECT id, ordinal FROM messages + WHERE session_id = ? AND ordinal = ? + ORDER BY id LIMIT 1`, + sessionID, pin.Ordinal, + ) +} + +func metadataPinTargetByQueryTx( + ctx context.Context, + tx *sql.Tx, + query string, + args ...any, +) (metadataPinTarget, bool, error) { + var target metadataPinTarget + err := tx.QueryRowContext(ctx, query, args...).Scan(&target.id, &target.ordinal) + if err == sql.ErrNoRows { + return metadataPinTarget{}, false, nil + } + if err != nil { + return metadataPinTarget{}, false, fmt.Errorf("finding metadata pin target: %w", err) + } + return target, true, nil +} + +func unpinMetadataTx( + ctx context.Context, + tx *sql.Tx, + sessionID string, + pin MetadataPinProjection, +) error { + if pin.SourceUUID != "" { + res, err := tx.ExecContext(ctx, + `DELETE FROM pinned_messages + WHERE session_id = ? + AND message_id IN ( + SELECT id FROM messages + WHERE session_id = ? AND source_uuid = ? + )`, + sessionID, sessionID, pin.SourceUUID, + ) + if err != nil { + return err + } + if n, _ := res.RowsAffected(); n > 0 { + return nil + } + } + _, err := tx.ExecContext(ctx, + `DELETE FROM pinned_messages + WHERE session_id = ? + AND message_id IN ( + SELECT id FROM messages + WHERE session_id = ? AND ordinal = ? + )`, + sessionID, sessionID, pin.Ordinal, + ) + return err +} + +func upsertMetadataReplayStateTx(ctx context.Context, tx *sql.Tx, ev MetadataProjection) error { + _, err := tx.ExecContext(ctx, + `INSERT INTO metadata_replay_state + (session_gid, field, order_key, hlc, artifact_hash, origin, op, value, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, strftime('%Y-%m-%dT%H:%M:%fZ','now')) + ON CONFLICT(session_gid, field) DO UPDATE SET + order_key = excluded.order_key, + hlc = excluded.hlc, + artifact_hash = excluded.artifact_hash, + origin = excluded.origin, + op = excluded.op, + value = excluded.value, + updated_at = excluded.updated_at`, + ev.SessionGID, ev.Field, ev.OrderKey, ev.HLC, ev.ArtifactHash, + ev.EventOrigin, ev.Op, ev.Value, + ) + if err != nil { + return fmt.Errorf("upserting metadata replay state: %w", err) + } + return nil +} + +func markMetadataEventAppliedTx( + ctx context.Context, + tx *sql.Tx, + origin, orderKey, hash string, +) error { + _, err := tx.ExecContext(ctx, + `INSERT OR IGNORE INTO metadata_applied_events + (origin, order_key, artifact_hash) + VALUES (?, ?, ?)`, + origin, orderKey, hash, + ) + if err != nil { + return fmt.Errorf("marking metadata event %s/%s applied: %w", origin, orderKey, err) + } + return nil +} + +type metadataConflict struct { + sessionGID string + field string + winningOrderKey string + losingOrderKey string + winningOrigin string + losingOrigin string + winningOp string + losingOp string + winningValue string + losingValue string +} + +func insertMetadataConflictTx(ctx context.Context, tx *sql.Tx, c metadataConflict) error { + _, err := tx.ExecContext(ctx, + `INSERT OR IGNORE INTO metadata_conflicts + (session_gid, field, winning_order_key, losing_order_key, + winning_origin, losing_origin, winning_op, losing_op, + winning_value, losing_value) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + c.sessionGID, c.field, c.winningOrderKey, c.losingOrderKey, + c.winningOrigin, c.losingOrigin, c.winningOp, c.losingOp, + c.winningValue, c.losingValue, + ) + if err != nil { + return fmt.Errorf("inserting metadata conflict: %w", err) + } + return nil +} + +func metadataStateDiffers(aOp, aValue, bOp, bValue string) bool { + return aOp != bOp || aValue != bValue +} + +func metadataConflictOriginsDiffer(winningOrigin, losingOrigin string) bool { + return winningOrigin != losingOrigin +} + +func uniqueNonEmptyStrings(values []string) []string { + seen := make(map[string]bool, len(values)) + unique := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" || seen[value] { + continue + } + seen[value] = true + unique = append(unique, value) + } + return unique +} diff --git a/internal/db/metadata_replay_test.go b/internal/db/metadata_replay_test.go new file mode 100644 index 000000000..d0f49b72c --- /dev/null +++ b/internal/db/metadata_replay_test.go @@ -0,0 +1,343 @@ +package db + +import ( + "context" + "database/sql" + "fmt" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestApplyMetadataProjectionSessionOps(t *testing.T) { + ctx := context.Background() + d := testDB(t) + insertSession(t, d, "s1", "alpha") + + applyMetadataProjectionForTest(t, d, MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0001-a", + HLC: "0001", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "starred", + Op: "star", + Value: "star", + }) + assert.Equal(t, 1, metadataTableCount(t, d, "starred_sessions", "session_id = 's1'")) + + applyMetadataProjectionForTest(t, d, MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0002-a", + HLC: "0002", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "starred", + Op: "unstar", + Value: "unstar", + }) + assert.Equal(t, 0, metadataTableCount(t, d, "starred_sessions", "session_id = 's1'")) + + applyMetadataProjectionForTest(t, d, MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0003-a", + HLC: "0003", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "deleted_at", + Op: "soft_delete", + Value: "soft_delete", + }) + got, err := d.GetSession(ctx, "s1") + require.NoError(t, err) + assert.Nil(t, got) + + applyMetadataProjectionForTest(t, d, MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0004-a", + HLC: "0004", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "deleted_at", + Op: "restore", + Value: "restore", + }) + got, err = d.GetSession(ctx, "s1") + require.NoError(t, err) + require.NotNil(t, got) + assert.Equal(t, "s1", got.ID) + + applyMetadataProjectionForTest(t, d, MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0005-a", + HLC: "0005", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "purge", + Op: "purge", + Value: "purge", + }) + got, err = d.GetSessionFull(ctx, "s1") + require.NoError(t, err) + assert.Nil(t, got) + assert.Equal(t, 1, metadataTableCount(t, d, "excluded_sessions", "id = 's1'")) +} + +func TestApplyMetadataProjectionRequiresPinTargetBeforeMarkingApplied(t *testing.T) { + d := testDB(t) + insertSession(t, d, "s1", "alpha") + + result, err := d.ApplyMetadataProjection(context.Background(), MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0001-a", + HLC: "0001", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "pin:source_uuid:missing", + Op: "pin", + Value: `{"ordinal":1,"source_uuid":"missing"}`, + Pin: &MetadataPinProjection{ + SourceUUID: "missing", + Ordinal: 1, + }, + }) + + require.ErrorIs(t, err, ErrMetadataTargetUnavailable) + assert.False(t, result.Applied) + applied, checkErr := d.MetadataEventApplied(context.Background(), "laptop-a1b2c3", "0001-a") + require.NoError(t, checkErr) + assert.False(t, applied) + assert.Equal(t, 0, metadataTableCount(t, d, "metadata_replay_state", "session_gid = 'desktop-d4e5f6~s1'")) +} + +func TestApplyMetadataProjectionDoesNotConflictSameOriginSequentialEdits(t *testing.T) { + d := testDB(t) + insertSession(t, d, "s1", "alpha") + firstName := "one" + secondName := "two" + + events := []MetadataProjection{ + { + EventOrigin: "laptop-a1b2c3", + OrderKey: "0001-a", + HLC: "0001", + ArtifactHash: "a", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "starred", + Op: "star", + Value: "star", + }, + { + EventOrigin: "laptop-a1b2c3", + OrderKey: "0002-a", + HLC: "0002", + ArtifactHash: "b", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "starred", + Op: "unstar", + Value: "unstar", + }, + { + EventOrigin: "laptop-a1b2c3", + OrderKey: "0003-a", + HLC: "0003", + ArtifactHash: "c", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "display_name", + Op: "rename", + Value: `{"display_name":"one"}`, + DisplayName: &firstName, + }, + { + EventOrigin: "laptop-a1b2c3", + OrderKey: "0004-a", + HLC: "0004", + ArtifactHash: "d", + SessionGID: "desktop-d4e5f6~s1", + LocalSessionID: "s1", + Field: "display_name", + Op: "rename", + Value: `{"display_name":"two"}`, + DisplayName: &secondName, + }, + } + + for _, ev := range events { + result, err := d.ApplyMetadataProjection(context.Background(), ev) + require.NoError(t, err) + assert.True(t, result.Applied) + assert.False(t, result.Conflict) + } + + assert.Equal(t, 0, metadataTableCount(t, d, "metadata_conflicts", "1 = 1")) +} + +func TestMetadataConflictQueriesIgnoreSameOriginRows(t *testing.T) { + ctx := context.Background() + d := testDB(t) + _, err := d.getWriter().ExecContext(ctx, + `INSERT INTO metadata_conflicts + (session_gid, field, winning_order_key, losing_order_key, + winning_origin, losing_origin, winning_op, losing_op, + winning_value, losing_value) + VALUES + ('desktop-d4e5f6~s1', 'display_name', '0002-a', '0001-a', + 'desktop-d4e5f6', 'desktop-d4e5f6', 'rename', 'rename', + '{"display_name":"two"}', '{"display_name":"one"}'), + ('desktop-d4e5f6~s1', 'display_name', '0003-b', '0002-a', + 'laptop-a1b2c3', 'desktop-d4e5f6', 'rename', 'rename', + '{"display_name":"peer"}', '{"display_name":"two"}')`, + ) + require.NoError(t, err) + + conflicts, err := d.ListMetadataConflicts(ctx, []string{"desktop-d4e5f6~s1"}) + require.NoError(t, err) + require.Len(t, conflicts, 1) + assert.Equal(t, "laptop-a1b2c3", conflicts[0].WinningOrigin) + assert.Equal(t, "desktop-d4e5f6", conflicts[0].LosingOrigin) + + count, err := d.CountMetadataConflicts(ctx) + require.NoError(t, err) + assert.Equal(t, 1, count) +} + +func TestApplyMetadataProjectionPurgeExcludesFallbackAlias(t *testing.T) { + d := testDB(t) + filePath := "/tmp/vibe/session_20260616_083518_abc123/messages.jsonl" + insertSession(t, d, "vibe:canonical-1", "alpha", func(s *Session) { + s.Agent = "vibe" + s.FilePath = &filePath + }) + + applyMetadataProjectionForTest(t, d, MetadataProjection{ + EventOrigin: "laptop-a1b2c3", + OrderKey: "0001-a", + HLC: "0001", + ArtifactHash: "a", + SessionGID: "laptop-a1b2c3~vibe:canonical-1", + LocalSessionID: "vibe:canonical-1", + Field: "purge", + Op: "purge", + Value: "purge", + }) + + assert.Equal(t, 1, metadataTableCount(t, d, "excluded_sessions", "id = 'vibe:canonical-1'")) + assert.Equal(t, 1, metadataTableCount(t, d, "excluded_sessions", "id = 'vibe:session_20260616_083518_abc123'")) +} + +func TestMetadataArtifactProvenanceIsIndexedOrderedAndImmutable(t *testing.T) { + database := testDB(t) + ctx := t.Context() + rows := []MetadataArtifactProvenance{ + {Origin: "desk-a1b2c3", OrderKey: "0002", ArtifactHash: "hash-2", SessionGID: "desk-a1b2c3~s1", Op: "soft_delete"}, + {Origin: "desk-a1b2c3", OrderKey: "0001", ArtifactHash: "hash-1", SessionGID: "desk-a1b2c3~s1", Op: "star"}, + {Origin: "desk-a1b2c3", OrderKey: "0003", ArtifactHash: "hash-3", SessionGID: "desk-a1b2c3~other", Op: "star"}, + } + for _, row := range rows { + require.NoError(t, database.RecordMetadataArtifactProvenance(ctx, row)) + } + require.NoError(t, database.RecordMetadataArtifactProvenance(ctx, rows[0]), + "exact provenance replay is idempotent") + + got, err := database.MetadataArtifactProvenanceForSession( + ctx, "desk-a1b2c3", "desk-a1b2c3~s1", + ) + require.NoError(t, err) + require.Len(t, got, 2) + assert.Equal(t, "0001", got[0].OrderKey) + assert.Equal(t, "star", got[0].Op) + assert.Equal(t, "0002", got[1].OrderKey) + assert.Equal(t, "soft_delete", got[1].Op) + + filtered, err := database.MetadataArtifactProvenanceForSession( + ctx, "desk-a1b2c3", "desk-a1b2c3~s1", "soft_delete", + ) + require.NoError(t, err) + require.Len(t, filtered, 1) + assert.Equal(t, "0002", filtered[0].OrderKey) + + conflict := rows[0] + conflict.ArtifactHash = "different" + require.Error(t, database.RecordMetadataArtifactProvenance(ctx, conflict)) + unchanged, err := database.MetadataArtifactProvenanceForSession( + ctx, "desk-a1b2c3", "desk-a1b2c3~s1", "soft_delete", + ) + require.NoError(t, err) + assert.Equal(t, rows[0], unchanged[0]) +} + +func TestVisitMetadataReplayWinnersAuthoredByCrossesBoundedPagesAndFiltersOrigin(t *testing.T) { + database := testDB(t) + ctx := t.Context() + localOrigin := "desk-a1b2c3" + for index := range 257 { + identity := fmt.Sprintf("%064x", index+1) + result, err := database.RecordLocalMetadataProjection(ctx, MetadataProjection{ + EventOrigin: localOrigin, + OrderKey: fmt.Sprintf("%020d-%s", index+1, identity), + HLC: fmt.Sprintf("%020d", index+1), + ArtifactHash: identity, + SessionGID: fmt.Sprintf("%s~session-%03d", localOrigin, index), + LocalSessionID: fmt.Sprintf("session-%03d", index), + Field: "starred", + Op: "unstar", + Value: "unstar", + }) + require.NoError(t, err) + assert.True(t, result.Applied) + } + foreignIdentity := fmt.Sprintf("%064x", 999) + _, err := database.RecordLocalMetadataProjection(ctx, MetadataProjection{ + EventOrigin: "peer-d4e5f6", + OrderKey: "999-peer", + HLC: "999", + ArtifactHash: foreignIdentity, + SessionGID: localOrigin + "~foreign-winner", + LocalSessionID: "foreign-winner", + Field: "starred", + Op: "star", + Value: "star", + }) + require.NoError(t, err) + + var visited []string + err = database.VisitMetadataReplayWinnersAuthoredBy(ctx, localOrigin, + func(winner MetadataProjection) error { + assert.Equal(t, localOrigin, winner.EventOrigin) + visited = append(visited, winner.SessionGID) + return nil + }) + require.NoError(t, err) + require.Len(t, visited, 257) + assert.Equal(t, localOrigin+"~session-000", visited[0]) + assert.Equal(t, localOrigin+"~session-256", visited[len(visited)-1]) +} + +func applyMetadataProjectionForTest(t *testing.T, d *DB, ev MetadataProjection) { + t.Helper() + result, err := d.ApplyMetadataProjection(context.Background(), ev) + require.NoError(t, err) + assert.True(t, result.Applied) + assert.False(t, result.Duplicate) +} + +func metadataTableCount(t *testing.T, d *DB, table, where string) int { + t.Helper() + var count int + err := d.Reader().QueryRow("SELECT COUNT(*) FROM " + table + " WHERE " + where).Scan(&count) + if err == sql.ErrNoRows { + return 0 + } + require.NoError(t, err) + return count +} diff --git a/internal/db/orphaned.go b/internal/db/orphaned.go index 1d9c5eb33..8aea580a0 100644 --- a/internal/db/orphaned.go +++ b/internal/db/orphaned.go @@ -323,9 +323,12 @@ func (d *DB) CopyTrashedDataFrom(sourcePath string) (int, error) { return count, nil } -// CopySyncStateFrom copies pg_sync_state rows from the source database into the -// current database. ResyncAll uses this to preserve durable local sync metadata -// such as the PG push owner marker across the temp-DB swap. +// CopySyncStateFrom copies durable synchronization authority from the source +// database into the current database. Alongside selected pg_sync_state rows it +// preserves artifact publication work, publications, checkpoint heads and +// floors, and repair work across the temp-database resync swap. Transient +// bookkeeping such as last_sync_* timestamps is deliberately left behind so +// the rebuilt DB reports its own sync times. func (d *DB) CopySyncStateFrom(sourcePath string) error { d.mu.Lock() defer d.mu.Unlock() @@ -346,24 +349,234 @@ func (d *DB) CopySyncStateFrom(sourcePath string) error { _, _ = execWithoutCancel(ctx, conn, "DETACH DATABASE old_db") }() - // Older databases may have no pg_sync_state table. - var tableExists int - err = conn.QueryRowContext( - ctx, "SELECT 1 FROM old_db.sqlite_master WHERE type='table' AND name='pg_sync_state'", - ).Scan(&tableExists) + tx, err := conn.BeginTx(ctx, nil) if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return nil + return fmt.Errorf("beginning sync state copy: %w", err) + } + defer func() { _ = tx.Rollback() }() + + if oldDBHasTable(ctx, tx, "pg_sync_state") { + if _, err := tx.ExecContext(ctx, ` + INSERT OR REPLACE INTO main.pg_sync_state (key, value) + SELECT key, value FROM old_db.pg_sync_state + WHERE key = 'pg_push_marker_id' + OR key LIKE 'artifact\_%' ESCAPE '\'`); err != nil { + return fmt.Errorf("copying sync state: %w", err) } - return fmt.Errorf("probing pg_sync_state table: %w", err) } - _, err = conn.ExecContext(ctx, ` - INSERT OR REPLACE INTO main.pg_sync_state (key, value) - SELECT key, value FROM old_db.pg_sync_state - WHERE key = 'pg_push_marker_id'`) + headRevisionExpr := "0" + if oldDBHasColumn(ctx, tx, "artifact_checkpoint_heads", "publication_revision") { + headRevisionExpr = "publication_revision" + } + headSizeExpr := "0" + if oldDBHasColumn(ctx, tx, "artifact_checkpoint_heads", "checkpoint_size") { + headSizeExpr = "checkpoint_size" + } + if oldDBHasTable(ctx, tx, "artifact_import_queue") { + var conflicting bool + if err := tx.QueryRowContext(ctx, ` + SELECT EXISTS( + SELECT 1 + FROM main.artifact_import_queue AS current + JOIN old_db.artifact_import_queue AS previous + USING (origin, kind, name) + WHERE current.sha256 <> previous.sha256 + OR current.size <> previous.size + )`).Scan(&conflicting); err != nil { + return fmt.Errorf("checking artifact import identities: %w", err) + } + if conflicting { + return errors.New("copying artifact import queue found conflicting identity") + } + } + artifactCopies := []struct { + table string + sql string + }{ + { + "artifact_export_queue", + `INSERT INTO main.artifact_export_queue(session_id, enqueued_at, generation, pending) + SELECT session_id, enqueued_at, generation + 1, pending + FROM old_db.artifact_export_queue WHERE true + ON CONFLICT(session_id) DO UPDATE SET + enqueued_at = CASE + WHEN artifact_export_queue.pending = 1 AND excluded.pending = 1 + THEN min(artifact_export_queue.enqueued_at, excluded.enqueued_at) + WHEN artifact_export_queue.pending = 1 + THEN artifact_export_queue.enqueued_at + ELSE excluded.enqueued_at + END, + generation = max(artifact_export_queue.generation, excluded.generation) + 1, + pending = max(artifact_export_queue.pending, excluded.pending)`, + }, + { + "artifact_import_queue", + `INSERT INTO main.artifact_import_queue( + origin, kind, name, sha256, size, reason, + required_format_version, enqueued_at) + SELECT origin, kind, name, sha256, size, reason, + required_format_version, enqueued_at + FROM old_db.artifact_import_queue WHERE true + ON CONFLICT(origin, kind, name) DO UPDATE SET + reason = excluded.reason, + required_format_version = max( + artifact_import_queue.required_format_version, + excluded.required_format_version), + enqueued_at = min( + artifact_import_queue.enqueued_at, excluded.enqueued_at)`, + }, + { + "artifact_publications", + `INSERT OR REPLACE INTO main.artifact_publications( + origin, session_id, manifest_hash, source_fingerprint) + SELECT origin, session_id, manifest_hash, source_fingerprint + FROM old_db.artifact_publications`, + }, + { + "artifact_publication_revisions", + `INSERT INTO main.artifact_publication_revisions(origin, revision) + SELECT origin, revision FROM old_db.artifact_publication_revisions WHERE true + ON CONFLICT(origin) DO UPDATE SET + revision = max(artifact_publication_revisions.revision, excluded.revision)`, + }, + { + "artifact_checkpoint_heads", + `INSERT OR REPLACE INTO main.artifact_checkpoint_heads( + origin, sequence, publication_revision, session_map_sha256, + checkpoint_sha256, checkpoint_size) + SELECT origin, sequence, ` + headRevisionExpr + `, session_map_sha256, + checkpoint_sha256, ` + headSizeExpr + ` + FROM old_db.artifact_checkpoint_heads`, + }, + { + "artifact_checkpoint_floors", + `INSERT INTO main.artifact_checkpoint_floors(origin, sequence) + SELECT origin, sequence FROM old_db.artifact_checkpoint_floors WHERE true + ON CONFLICT(origin) DO UPDATE SET + sequence = max(artifact_checkpoint_floors.sequence, excluded.sequence)`, + }, + { + "artifact_checkpoint_landings", + `INSERT OR REPLACE INTO main.artifact_checkpoint_landings(origin, sequence) + SELECT origin, sequence FROM old_db.artifact_checkpoint_landings`, + }, + { + "artifact_peer_checkpoint_heads", + `INSERT OR REPLACE INTO main.artifact_peer_checkpoint_heads( + origin, sequence, checkpoint_sha256, checkpoint_size) + SELECT origin, sequence, checkpoint_sha256, checkpoint_size + FROM old_db.artifact_peer_checkpoint_heads`, + }, + { + "artifact_checkpoint_landing_sessions", + `INSERT OR REPLACE INTO main.artifact_checkpoint_landing_sessions( + origin, gid, manifest_hash) + SELECT origin, gid, manifest_hash + FROM old_db.artifact_checkpoint_landing_sessions`, + }, + { + "artifact_repair_queue", + `INSERT OR REPLACE INTO main.artifact_repair_queue( + origin, kind, name, sha256, size, detected_at) + SELECT origin, kind, name, sha256, size, detected_at + FROM old_db.artifact_repair_queue`, + }, + } + for _, copy := range artifactCopies { + if !oldDBHasTable(ctx, tx, copy.table) { + continue + } + if _, err := tx.ExecContext(ctx, copy.sql); err != nil { + return fmt.Errorf("copying %s: %w", copy.table, err) + } + } + if _, err := tx.ExecContext(ctx, ` + DELETE FROM main.artifact_import_queue AS candidate + WHERE candidate.kind = 'checkpoints' + AND EXISTS ( + SELECT 1 FROM main.artifact_import_queue AS newer + WHERE newer.origin = candidate.origin + AND newer.kind = 'checkpoints' + AND newer.name > candidate.name + )`); err != nil { + return fmt.Errorf("retiring copied artifact checkpoint work: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing sync state copy: %w", err) + } + return nil +} + +// CopyMetadataReplayFrom copies the durable artifact metadata replay tables +// (metadata provenance, applied events, replay state, and conflicts) from +// the source database. ResyncAll uses this so previously applied peer +// metadata events are not replayed against an empty LWW register after a +// full rebuild, which would let old events overwrite newer local state. +func (d *DB) CopyMetadataReplayFrom(sourcePath string) error { + d.mu.Lock() + defer d.mu.Unlock() + + ctx := context.Background() + conn, err := d.getWriter().Conn(ctx) if err != nil { - return fmt.Errorf("copying sync state: %w", err) + return fmt.Errorf("acquiring connection: %w", err) + } + defer conn.Close() + + if _, err := conn.ExecContext( + ctx, "ATTACH DATABASE ? AS old_db", sourcePath, + ); err != nil { + return fmt.Errorf("attaching source db: %w", err) + } + defer func() { + _, _ = execWithoutCancel(ctx, conn, "DETACH DATABASE old_db") + }() + + copies := []struct { + table string + columns string + }{ + { + "metadata_artifact_provenance", + "origin, order_key, artifact_hash, session_gid, op", + }, + { + "metadata_applied_events", + "origin, order_key, artifact_hash, applied_at", + }, + { + "metadata_replay_state", + "session_gid, field, order_key, hlc, artifact_hash, " + + "origin, op, value, updated_at", + }, + { + "metadata_conflicts", + "session_gid, field, winning_order_key, losing_order_key, " + + "winning_origin, losing_origin, winning_op, losing_op, " + + "winning_value, losing_value, created_at", + }, + } + for _, c := range copies { + // Older databases may predate the metadata replay tables. + var tableExists int + err := conn.QueryRowContext(ctx, + "SELECT 1 FROM old_db.sqlite_master WHERE type='table' AND name=?", + c.table, + ).Scan(&tableExists) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + continue + } + return fmt.Errorf("probing %s table: %w", c.table, err) + } + if _, err := conn.ExecContext(ctx, fmt.Sprintf( + `INSERT OR IGNORE INTO main.%s (%s) + SELECT %s FROM old_db.%s`, + c.table, c.columns, c.columns, c.table, + )); err != nil { + return fmt.Errorf("copying %s: %w", c.table, err) + } } return nil } diff --git a/internal/db/read_only_test.go b/internal/db/read_only_test.go index b012cd954..35df42119 100644 --- a/internal/db/read_only_test.go +++ b/internal/db/read_only_test.go @@ -234,7 +234,8 @@ func TestOpenReadOnlyWriteMethodsReturnErrReadOnly(t *testing.T) { return readonly.InsertMessages(nil) }) requireReadOnlyOp(t, "BulkStarSessions", func() error { - return readonly.BulkStarSessions(nil) + _, err := readonly.BulkStarSessions(nil) + return err }) requireReadOnlyOp(t, "DeleteParserExcludedSessions", func() error { _, err := readonly.DeleteParserExcludedSessions(nil) @@ -268,6 +269,7 @@ func TestOpenReadOnlyWriteMethodsReturnErrReadOnly(t *testing.T) { func TestOpenReadOnlyRejectsMissingMigratedColumn(t *testing.T) { path := createClosedTestDB(t, tempDBPath(t, "sessions.db"), nil) + execRawSQLite(t, path, "DROP TRIGGER IF EXISTS artifact_sessions_update_queue") execRawSQLite(t, path, "ALTER TABLE sessions DROP COLUMN display_name") requireOpenReadOnlyFails(t, path, "schema missing sessions.display_name") } @@ -299,10 +301,25 @@ func TestReadOnlySchemaCompatibilityRejectsMissingReadColumn(t *testing.T) { {"extract generation", "recall_extract_generations", "state"}, {"extract progress stamp", "recall_extract_progress", "content_stamped_at"}, + {"artifact export authority", "artifact_export_queue", "pending"}, + {"artifact checkpoint revision", "artifact_checkpoint_heads", "publication_revision"}, + {"artifact checkpoint size", "artifact_checkpoint_heads", "checkpoint_size"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { conn := openReadOnlySchemaProbe(t) + if tt.table == "sessions" { + _, err := conn.Exec("DROP TRIGGER IF EXISTS artifact_sessions_update_queue") + require.NoError(t, err) + } + if tt.table == "artifact_export_queue" { + _, err := conn.Exec(` + DROP INDEX IF EXISTS idx_artifact_export_queue_pending; + DROP TRIGGER IF EXISTS artifact_sessions_insert_queue; + DROP TRIGGER IF EXISTS artifact_sessions_update_queue; + DROP TRIGGER IF EXISTS artifact_sessions_delete_queue`) + require.NoError(t, err) + } _, err := conn.Exec( "ALTER TABLE " + tt.table + " DROP COLUMN " + tt.column) require.NoError(t, err) @@ -328,6 +345,9 @@ func TestOpenReadOnlyRejectsMissingReadTable(t *testing.T) { {"recall_query_exposures", "query_id"}, {"recall_extract_generations", "fingerprint"}, {"recall_extract_progress", "session_id"}, + {"artifact_export_queue", "session_id"}, + {"artifact_import_queue", "origin"}, + {"artifact_publication_revisions", "origin"}, } for _, tt := range tests { t.Run(tt.table, func(t *testing.T) { diff --git a/internal/db/schema.sql b/internal/db/schema.sql index 43dba0942..c99a80f19 100644 --- a/internal/db/schema.sql +++ b/internal/db/schema.sql @@ -118,6 +118,207 @@ CREATE TABLE IF NOT EXISTS messages ( UNIQUE(session_id, ordinal) ); +-- Durable, bounded artifact publication state. The export queue intentionally +-- has no foreign key: a deleted locally-owned session remains pending until a +-- checkpoint publishes its removal. Acknowledged rows remain as generation +-- authority, so this table is bounded by historical archive session IDs rather +-- than only the currently dirty set. +CREATE TABLE IF NOT EXISTS artifact_export_queue ( + session_id TEXT PRIMARY KEY, + enqueued_at TEXT NOT NULL + DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + -- Compare-and-ack token. Repeated writes retain their FIFO timestamp but + -- advance generation, including multiple writes in one SQLite millisecond. + generation INTEGER NOT NULL DEFAULT 1, + -- Acknowledgement clears pending but retains the row as durable generation + -- authority, preventing an old claim from becoming valid after requeue. + pending INTEGER NOT NULL DEFAULT 1 CHECK (pending IN (0, 1)) +); +CREATE INDEX IF NOT EXISTS idx_artifact_export_queue_pending + ON artifact_export_queue(pending, enqueued_at, session_id); + +-- Durable exact-reference work that could not finish during a bounded artifact +-- import pass. Rows are retained until compare-and-delete acknowledgement; +-- future-format work becomes eligible when the running reader advances. +CREATE TABLE IF NOT EXISTS artifact_import_queue ( + origin TEXT NOT NULL, + kind TEXT NOT NULL, + name TEXT NOT NULL, + sha256 TEXT NOT NULL, + size INTEGER NOT NULL CHECK (size >= 0), + reason TEXT NOT NULL, + required_format_version INTEGER NOT NULL CHECK (required_format_version >= 1), + enqueued_at TEXT NOT NULL + DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + PRIMARY KEY (origin, kind, name) +); +CREATE INDEX IF NOT EXISTS idx_artifact_import_queue_pending + ON artifact_import_queue( + required_format_version, enqueued_at, origin, kind, name + ); + +CREATE TABLE IF NOT EXISTS artifact_publications ( + origin TEXT NOT NULL, + session_id TEXT NOT NULL, + manifest_hash TEXT NOT NULL, + source_fingerprint TEXT NOT NULL, + PRIMARY KEY(origin, session_id) +); + +CREATE TABLE IF NOT EXISTS artifact_publication_revisions ( + origin TEXT PRIMARY KEY, + revision INTEGER NOT NULL +); + +CREATE TABLE IF NOT EXISTS artifact_checkpoint_heads ( + origin TEXT PRIMARY KEY, + sequence INTEGER NOT NULL, + publication_revision INTEGER NOT NULL, + session_map_sha256 TEXT NOT NULL, + checkpoint_sha256 TEXT NOT NULL, + checkpoint_size INTEGER NOT NULL +); + +CREATE TABLE IF NOT EXISTS artifact_checkpoint_floors ( + origin TEXT PRIMARY KEY, + sequence INTEGER NOT NULL +); + +CREATE TABLE IF NOT EXISTS artifact_checkpoint_landings ( + origin TEXT PRIMARY KEY, + sequence INTEGER NOT NULL +); + +CREATE TABLE IF NOT EXISTS artifact_peer_checkpoint_heads ( + origin TEXT PRIMARY KEY, + sequence INTEGER NOT NULL, + checkpoint_sha256 TEXT NOT NULL, + checkpoint_size INTEGER NOT NULL +); + +CREATE TABLE IF NOT EXISTS artifact_checkpoint_landing_sessions ( + origin TEXT NOT NULL, + gid TEXT NOT NULL, + manifest_hash TEXT NOT NULL, + PRIMARY KEY(origin, gid), + FOREIGN KEY(origin) REFERENCES artifact_checkpoint_landings(origin) + ON DELETE CASCADE +); + +CREATE TABLE IF NOT EXISTS artifact_repair_queue ( + origin TEXT NOT NULL, + kind TEXT NOT NULL, + name TEXT NOT NULL, + sha256 TEXT NOT NULL, + size INTEGER NOT NULL, + detected_at TEXT NOT NULL + DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + PRIMARY KEY(origin, kind, name) +); +CREATE INDEX IF NOT EXISTS idx_artifact_repair_queue_detected + ON artifact_repair_queue(detected_at, origin, kind, name); + +-- Session ownership is represented by machine='local'. Queue both sides of an +-- ownership transition: OLD local publishes a removal and NEW local publishes +-- content. A BEFORE DELETE trigger preserves the owner signal before child +-- cascades run. +DROP TRIGGER IF EXISTS artifact_sessions_insert_queue; +DROP TRIGGER IF EXISTS artifact_sessions_update_queue; +DROP TRIGGER IF EXISTS artifact_sessions_delete_queue; +DROP TRIGGER IF EXISTS artifact_messages_insert_queue; +DROP TRIGGER IF EXISTS artifact_messages_update_queue; +DROP TRIGGER IF EXISTS artifact_messages_delete_queue; +DROP TRIGGER IF EXISTS artifact_usage_events_insert_queue; +DROP TRIGGER IF EXISTS artifact_usage_events_update_queue; +DROP TRIGGER IF EXISTS artifact_usage_events_delete_queue; + +CREATE TRIGGER IF NOT EXISTS artifact_sessions_insert_queue +AFTER INSERT ON sessions WHEN NEW.machine = 'local' BEGIN + INSERT INTO artifact_export_queue(session_id) VALUES (NEW.id) + ON CONFLICT(session_id) DO UPDATE SET + enqueued_at = CASE WHEN pending = 0 + THEN strftime('%Y-%m-%dT%H:%M:%fZ','now') ELSE enqueued_at END, + generation = generation + 1, + pending = 1; +END; + +CREATE TRIGGER IF NOT EXISTS artifact_sessions_update_queue +AFTER UPDATE ON sessions +WHEN (OLD.machine = 'local' OR NEW.machine = 'local') AND ( + OLD.project IS NOT NEW.project OR + OLD.machine IS NOT NEW.machine OR + OLD.agent IS NOT NEW.agent OR + OLD.agent_label IS NOT NEW.agent_label OR + OLD.entrypoint IS NOT NEW.entrypoint OR + OLD.first_message IS NOT NEW.first_message OR + OLD.display_name IS NOT NEW.display_name OR + OLD.session_name IS NOT NEW.session_name OR + OLD.started_at IS NOT NEW.started_at OR + OLD.ended_at IS NOT NEW.ended_at OR + OLD.message_count IS NOT NEW.message_count OR + OLD.user_message_count IS NOT NEW.user_message_count OR + OLD.transcript_revision IS NOT NEW.transcript_revision OR + OLD.parent_session_id IS NOT NEW.parent_session_id OR + OLD.relationship_type IS NOT NEW.relationship_type OR + OLD.total_output_tokens IS NOT NEW.total_output_tokens OR + OLD.peak_context_tokens IS NOT NEW.peak_context_tokens OR + OLD.has_total_output_tokens IS NOT NEW.has_total_output_tokens OR + OLD.has_peak_context_tokens IS NOT NEW.has_peak_context_tokens OR + OLD.is_automated IS NOT NEW.is_automated OR + OLD.tool_failure_signal_count IS NOT NEW.tool_failure_signal_count OR + OLD.tool_retry_count IS NOT NEW.tool_retry_count OR + OLD.edit_churn_count IS NOT NEW.edit_churn_count OR + OLD.consecutive_failure_max IS NOT NEW.consecutive_failure_max OR + OLD.outcome IS NOT NEW.outcome OR + OLD.outcome_confidence IS NOT NEW.outcome_confidence OR + OLD.ended_with_role IS NOT NEW.ended_with_role OR + OLD.final_failure_streak IS NOT NEW.final_failure_streak OR + OLD.signals_pending_since IS NOT NEW.signals_pending_since OR + OLD.compaction_count IS NOT NEW.compaction_count OR + OLD.mid_task_compaction_count IS NOT NEW.mid_task_compaction_count OR + OLD.context_pressure_max IS NOT NEW.context_pressure_max OR + OLD.health_score IS NOT NEW.health_score OR + OLD.health_grade IS NOT NEW.health_grade OR + OLD.has_tool_calls IS NOT NEW.has_tool_calls OR + OLD.has_context_data IS NOT NEW.has_context_data OR + OLD.quality_signal_version IS NOT NEW.quality_signal_version OR + OLD.short_prompt_count IS NOT NEW.short_prompt_count OR + OLD.unstructured_start IS NOT NEW.unstructured_start OR + OLD.missing_success_criteria_count IS NOT NEW.missing_success_criteria_count OR + OLD.missing_verification_count IS NOT NEW.missing_verification_count OR + OLD.duplicate_prompt_count IS NOT NEW.duplicate_prompt_count OR + OLD.no_code_context_count IS NOT NEW.no_code_context_count OR + OLD.runaway_tool_loop_count IS NOT NEW.runaway_tool_loop_count OR + OLD.data_version IS NOT NEW.data_version OR + OLD.cwd IS NOT NEW.cwd OR + OLD.git_branch IS NOT NEW.git_branch OR + OLD.source_session_id IS NOT NEW.source_session_id OR + OLD.source_version IS NOT NEW.source_version OR + OLD.transcript_fidelity IS NOT NEW.transcript_fidelity OR + OLD.parser_malformed_lines IS NOT NEW.parser_malformed_lines OR + OLD.is_truncated IS NOT NEW.is_truncated OR + OLD.deleted_at IS NOT NEW.deleted_at OR + OLD.created_at IS NOT NEW.created_at OR + OLD.termination_status IS NOT NEW.termination_status +) BEGIN + INSERT INTO artifact_export_queue(session_id) VALUES (NEW.id) + ON CONFLICT(session_id) DO UPDATE SET + enqueued_at = CASE WHEN pending = 0 + THEN strftime('%Y-%m-%dT%H:%M:%fZ','now') ELSE enqueued_at END, + generation = generation + 1, + pending = 1; +END; + +CREATE TRIGGER IF NOT EXISTS artifact_sessions_delete_queue +BEFORE DELETE ON sessions WHEN OLD.machine = 'local' BEGIN + INSERT INTO artifact_export_queue(session_id) VALUES (OLD.id) + ON CONFLICT(session_id) DO UPDATE SET + enqueued_at = CASE WHEN pending = 0 + THEN strftime('%Y-%m-%dT%H:%M:%fZ','now') ELSE enqueued_at END, + generation = generation + 1, + pending = 1; +END; + -- Stats table maintained by triggers CREATE TABLE IF NOT EXISTS stats ( key TEXT PRIMARY KEY, @@ -492,6 +693,8 @@ CREATE TABLE IF NOT EXISTS pinned_messages ( CREATE INDEX IF NOT EXISTS idx_pinned_session ON pinned_messages(session_id); +CREATE INDEX IF NOT EXISTS idx_pinned_session_ordinal_id + ON pinned_messages(session_id, ordinal, id); -- idx_pinned_message backs the ON DELETE CASCADE from messages(id). -- The UNIQUE(session_id, message_id) constraint creates an index -- ordered (session_id, message_id), which the FK lookup on @@ -843,6 +1046,59 @@ CREATE TABLE IF NOT EXISTS pg_sync_state ( value TEXT NOT NULL ); +-- Metadata replay: durable record of handled artifact events. This lets +-- import scan the small append-only feed repeatedly without replaying +-- duplicates or relying on an unsafe max-HLC watermark. +CREATE TABLE IF NOT EXISTS metadata_applied_events ( + origin TEXT NOT NULL, + order_key TEXT NOT NULL, + artifact_hash TEXT NOT NULL, + applied_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + PRIMARY KEY (origin, order_key) +); + +CREATE TABLE IF NOT EXISTS metadata_artifact_provenance ( + origin TEXT NOT NULL, + order_key TEXT NOT NULL, + artifact_hash TEXT NOT NULL, + session_gid TEXT NOT NULL, + op TEXT NOT NULL, + PRIMARY KEY (origin, order_key) +); +CREATE INDEX IF NOT EXISTS idx_metadata_artifact_provenance_session + ON metadata_artifact_provenance(origin, session_gid, op, order_key); + +-- Per-field LWW winners for metadata replay. +CREATE TABLE IF NOT EXISTS metadata_replay_state ( + session_gid TEXT NOT NULL, + field TEXT NOT NULL, + order_key TEXT NOT NULL, + hlc TEXT NOT NULL, + artifact_hash TEXT NOT NULL, + origin TEXT NOT NULL, + op TEXT NOT NULL, + value TEXT NOT NULL DEFAULT '', + updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + PRIMARY KEY (session_gid, field) +); + +-- Losing metadata values that were overridden by deterministic LWW replay. +CREATE TABLE IF NOT EXISTS metadata_conflicts ( + id INTEGER PRIMARY KEY, + session_gid TEXT NOT NULL, + field TEXT NOT NULL, + winning_order_key TEXT NOT NULL, + losing_order_key TEXT NOT NULL, + winning_origin TEXT NOT NULL, + losing_origin TEXT NOT NULL, + winning_op TEXT NOT NULL, + losing_op TEXT NOT NULL, + winning_value TEXT NOT NULL DEFAULT '', + losing_value TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ','now')), + UNIQUE(session_gid, field, winning_order_key, losing_order_key) +); + -- Model pricing for cost calculation CREATE TABLE IF NOT EXISTS model_pricing ( model_pattern TEXT PRIMARY KEY, diff --git a/internal/db/session_batch.go b/internal/db/session_batch.go index fc0139fa6..84c78af2d 100644 --- a/internal/db/session_batch.go +++ b/internal/db/session_batch.go @@ -325,6 +325,12 @@ func writeOneSessionBatchTx( } replaceMessages := write.ReplaceMessages || (deletionCause.Valid && deletionCause.String == deletionCauseSourceMissing) + queueGenerationBefore, queueExistedBefore, err := artifactExportGenerationTx( + tx, write.Session.ID, + ) + if err != nil { + return 0, err + } replacementTranscriptChanged := false if replaceMessages && sessionExists { stored, err := sessionMessagesTx( @@ -355,7 +361,7 @@ func writeOneSessionBatchTx( } } if err := replaceSessionUsageEventsTx( - tx, write.Session.ID, write.UsageEvents, + tx, write.Session.ID, write.UsageEvents, false, ); err != nil { return 0, err } @@ -451,6 +457,18 @@ func writeOneSessionBatchTx( write.Signals.SecretLeakCount, write.Signals.SecretsRulesVersion); err != nil { return 0, err } + queueGenerationAfter, queueExistsAfter, err := artifactExportGenerationTx( + tx, write.Session.ID, + ) + if err != nil { + return 0, err + } + if queueExistedBefore == queueExistsAfter && + queueGenerationBefore == queueGenerationAfter { + if err := enqueueArtifactExportTx(tx, write.Session.ID); err != nil { + return 0, err + } + } return len(msgs), nil } diff --git a/internal/db/sessions.go b/internal/db/sessions.go index dcd513d19..f524aee9f 100644 --- a/internal/db/sessions.go +++ b/internal/db/sessions.go @@ -2844,6 +2844,35 @@ func sqliteLikeEscape(value string) string { return value } +// ListOwnedSessionIDsForExport returns the IDs of locally-owned, non-deleted +// sessions for artifact export, ordered by id. Unlike ListSessions it does not +// apply the sidebar visibility filter (message_count > 0), so zero-message +// usage-only sessions are still published. +func (db *DB) ListOwnedSessionIDsForExport(ctx context.Context) ([]string, error) { + rows, err := db.getReader().QueryContext(ctx, + `SELECT id FROM sessions + WHERE machine = 'local' AND deleted_at IS NULL + ORDER BY id`, + ) + if err != nil { + return nil, fmt.Errorf("listing sessions for artifact export: %w", err) + } + defer rows.Close() + + var ids []string + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, fmt.Errorf("scanning export session ID: %w", err) + } + ids = append(ids, id) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating export session IDs: %w", err) + } + return ids, nil +} + // GetDataVersionByPath returns the minimum data_version for non-source-missing // sessions matching a file_path. Returns 0 when no eligible session exists. func (db *DB) GetDataVersionByPath(path string) int { @@ -3271,6 +3300,32 @@ func (db *DB) GetBranches( return branches, rows.Err() } +// MachineSessionCounts returns the number of non-deleted sessions per machine, +// keyed by machine name. Child sessions are included so the count reflects the +// full corpus owned by or imported from each origin. +func (db *DB) MachineSessionCounts(ctx context.Context) (map[string]int, error) { + rows, err := db.getReader().QueryContext(ctx, + `SELECT machine, COUNT(*) FROM sessions + WHERE deleted_at IS NULL + GROUP BY machine`, + ) + if err != nil { + return nil, fmt.Errorf("counting sessions per machine: %w", err) + } + defer rows.Close() + + counts := map[string]int{} + for rows.Next() { + var machine string + var count int + if err := rows.Scan(&machine, &count); err != nil { + return nil, fmt.Errorf("scanning machine session count: %w", err) + } + counts[machine] = count + } + return counts, rows.Err() +} + // scanSessionRows iterates rows and scans each using // scanSessionRow. func scanSessionRows(rows *sql.Rows) ([]Session, error) { @@ -3423,8 +3478,16 @@ func (db *DB) SoftDeleteSession(id string) error { // tombstones are skipped; recoverable source-missing tombstones are converted. // Returns the count of newly deleted or converted rows. func (db *DB) SoftDeleteSessions(ids []string) (int, error) { + deleted, err := db.SoftDeleteSessionsReturningIDs(ids) + return len(deleted), err +} + +// SoftDeleteSessionsReturningIDs marks multiple sessions as deleted by setting +// deleted_at and returns the IDs that were newly deleted. Sessions that are +// already soft-deleted are skipped. +func (db *DB) SoftDeleteSessionsReturningIDs(ids []string) ([]string, error) { if len(ids) == 0 { - return 0, nil + return []string{}, nil } db.mu.Lock() @@ -3432,11 +3495,11 @@ func (db *DB) SoftDeleteSessions(ids []string) (int, error) { tx, err := db.getWriter().Begin() if err != nil { - return 0, fmt.Errorf("beginning soft-delete tx: %w", err) + return nil, fmt.Errorf("beginning soft-delete tx: %w", err) } defer func() { _ = tx.Rollback() }() - total := 0 + deleted := make([]string, 0, len(ids)) const batchSize = 500 for i := 0; i < len(ids); i += batchSize { end := min(i+batchSize, len(ids)) @@ -3448,7 +3511,38 @@ func (db *DB) SoftDeleteSessions(ids []string) (int, error) { } placeholders := strings.Repeat(",?", len(batch))[1:] - res, err := tx.Exec( + rows, err := tx.Query( + `SELECT id FROM sessions + WHERE id IN (`+placeholders+`) + AND (deleted_at IS NULL OR deletion_cause = '`+deletionCauseSourceMissing+`')`, + args..., + ) + if err != nil { + return nil, fmt.Errorf("loading soft-delete batch ids: %w", err) + } + batchIDs := map[string]struct{}{} + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + _ = rows.Close() + return nil, fmt.Errorf("scanning soft-delete batch id: %w", err) + } + batchIDs[id] = struct{}{} + } + if err := rows.Close(); err != nil { + return nil, fmt.Errorf("closing soft-delete batch ids: %w", err) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating soft-delete batch ids: %w", err) + } + for _, id := range batch { + if _, ok := batchIDs[id]; ok { + deleted = append(deleted, id) + delete(batchIDs, id) + } + } + + if _, err := tx.Exec( `UPDATE sessions SET deleted_at = strftime('%Y-%m-%dT%H:%M:%fZ','now'), deletion_cause = NULL, @@ -3456,18 +3550,64 @@ func (db *DB) SoftDeleteSessions(ids []string) (int, error) { WHERE id IN (`+placeholders+`) AND (deleted_at IS NULL OR deletion_cause = '`+deletionCauseSourceMissing+`')`, args..., - ) - if err != nil { - return 0, fmt.Errorf("soft-deleting batch: %w", err) + ); err != nil { + return nil, fmt.Errorf("soft-deleting batch: %w", err) } - n, _ := res.RowsAffected() - total += int(n) } if err := tx.Commit(); err != nil { - return 0, fmt.Errorf("committing soft-delete tx: %w", err) + return nil, fmt.Errorf("committing soft-delete tx: %w", err) } - return total, nil + return deleted, nil +} + +// TrashedSessionIDs returns requested sessions that currently exist in the +// trash. The result preserves first-seen input order and omits duplicates. +func (db *DB) TrashedSessionIDs(ids []string) ([]string, error) { + if len(ids) == 0 { + return []string{}, nil + } + trashed := make(map[string]struct{}, len(ids)) + const batchSize = 500 + for i := 0; i < len(ids); i += batchSize { + end := min(i+batchSize, len(ids)) + batch := ids[i:end] + args := make([]any, len(batch)) + for j, id := range batch { + args[j] = id + } + placeholders := strings.Repeat(",?", len(batch))[1:] + rows, err := db.getReader().Query( + `SELECT id FROM sessions + WHERE id IN (`+placeholders+`) AND deleted_at IS NOT NULL`, + args..., + ) + if err != nil { + return nil, fmt.Errorf("loading trashed session ids: %w", err) + } + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + _ = rows.Close() + return nil, fmt.Errorf("scanning trashed session id: %w", err) + } + trashed[id] = struct{}{} + } + if err := rows.Close(); err != nil { + return nil, fmt.Errorf("closing trashed session ids: %w", err) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterating trashed session ids: %w", err) + } + } + out := make([]string, 0, len(trashed)) + for _, id := range ids { + if _, ok := trashed[id]; ok { + out = append(out, id) + delete(trashed, id) + } + } + return out, nil } // RestoreSession clears deleted_at, making the session visible again. diff --git a/internal/db/starred.go b/internal/db/starred.go index ab9f38e2c..91f10e013 100644 --- a/internal/db/starred.go +++ b/internal/db/starred.go @@ -3,6 +3,7 @@ package db import ( "context" "database/sql" + "errors" "fmt" "strings" ) @@ -40,18 +41,22 @@ func (db *DB) StarSession(sessionID string) (bool, error) { return true, nil // already starred } -// UnstarSession removes a session's star. -func (db *DB) UnstarSession(sessionID string) error { +// UnstarSession removes a session's star and reports whether a row was removed. +func (db *DB) UnstarSession(sessionID string) (bool, error) { db.mu.Lock() defer db.mu.Unlock() - _, err := db.getWriter().Exec( + res, err := db.getWriter().Exec( "DELETE FROM starred_sessions WHERE session_id = ?", sessionID, ) if err != nil { - return fmt.Errorf("unstarring session %s: %w", sessionID, err) + return false, fmt.Errorf("unstarring session %s: %w", sessionID, err) } - return nil + n, err := res.RowsAffected() + if err != nil { + return false, fmt.Errorf("checking unstar result for %s: %w", sessionID, err) + } + return n > 0, nil } // ListStarredSessionIDs returns all starred session IDs. @@ -142,12 +147,12 @@ func curationScopeWhere(alias string, projects, excludeProjects []string) (strin // BulkStarSessions stars multiple sessions in a single transaction. // Used for migrating localStorage stars to the database. -func (db *DB) BulkStarSessions(sessionIDs []string) error { +func (db *DB) BulkStarSessions(sessionIDs []string) ([]string, error) { if err := db.requireWritable(); err != nil { - return err + return nil, err } if len(sessionIDs) == 0 { - return nil + return nil, nil } db.mu.Lock() @@ -155,27 +160,50 @@ func (db *DB) BulkStarSessions(sessionIDs []string) error { tx, err := db.getWriter().Begin() if err != nil { - return fmt.Errorf("beginning transaction: %w", err) + return nil, fmt.Errorf("beginning transaction: %w", err) } defer func() { _ = tx.Rollback() }() - // Use INSERT ... SELECT ... WHERE EXISTS so that stale IDs - // (sessions pruned or deleted from disk) are silently skipped - // instead of causing a foreign key violation that aborts the - // entire migration transaction. - stmt, err := tx.Prepare(` - INSERT OR IGNORE INTO starred_sessions (session_id) - SELECT ? WHERE EXISTS (SELECT 1 FROM sessions WHERE id = ?)`) + // Check existence separately from the insert so stale IDs (sessions pruned + // or deleted from disk) are silently skipped instead of aborting the + // migration transaction, and so the caller learns which sessions were + // actually starred and need a converging metadata event. + exists, err := tx.Prepare(`SELECT 1 FROM sessions WHERE id = ?`) + if err != nil { + return nil, fmt.Errorf("preparing existence statement: %w", err) + } + defer exists.Close() + insert, err := tx.Prepare( + `INSERT OR IGNORE INTO starred_sessions (session_id) VALUES (?)`) if err != nil { - return fmt.Errorf("preparing statement: %w", err) + return nil, fmt.Errorf("preparing statement: %w", err) } - defer stmt.Close() + defer insert.Close() + starred := make([]string, 0, len(sessionIDs)) for _, id := range sessionIDs { - if _, err := stmt.Exec(id, id); err != nil { - return fmt.Errorf("starring session %s: %w", id, err) + var one int + switch err := exists.QueryRow(id).Scan(&one); { + case errors.Is(err, sql.ErrNoRows): + continue + case err != nil: + return nil, fmt.Errorf("checking session %s: %w", id, err) + } + res, err := insert.Exec(id) + if err != nil { + return nil, fmt.Errorf("starring session %s: %w", id, err) + } + rowsAffected, err := res.RowsAffected() + if err != nil { + return nil, fmt.Errorf("checking star insert result for %s: %w", id, err) + } + if rowsAffected > 0 { + starred = append(starred, id) } } - return tx.Commit() + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("committing star transaction: %w", err) + } + return starred, nil } diff --git a/internal/db/store.go b/internal/db/store.go index 88f992c44..b65faeae4 100644 --- a/internal/db/store.go +++ b/internal/db/store.go @@ -41,6 +41,7 @@ type Store interface { GetMessagesWindow(ctx context.Context, sessionID string, w MessageWindow) ([]Message, error) GetAllMessages(ctx context.Context, sessionID string) ([]Message, error) GetResumeModelCounts(ctx context.Context, sessionID string) ([]ModelCount, error) + GetMessageForMetadataPin(ctx context.Context, sessionID string, messageID int64) (*Message, error) GetSessionActivity(ctx context.Context, sessionID string) (*SessionActivityResponse, error) // Timing. @@ -59,6 +60,9 @@ type Store interface { GetSessionVersion(id string) (count int, version int64, ok bool) // Metadata. + ListMetadataConflicts(ctx context.Context, sessionGIDs []string) ([]MetadataConflict, error) + CountMetadataConflicts(ctx context.Context) (int, error) + MachineSessionCounts(ctx context.Context) (map[string]int, error) GetStats(ctx context.Context, excludeOneShot, excludeAutomated bool) (Stats, error) GetProjects(ctx context.Context, excludeOneShot, excludeAutomated bool) ([]ProjectInfo, error) GetActiveProjectLabels(ctx context.Context) ([]string, error) @@ -94,9 +98,9 @@ type Store interface { // Stars. StarSession(sessionID string) (bool, error) - UnstarSession(sessionID string) error + UnstarSession(sessionID string) (bool, error) ListStarredSessionIDs(ctx context.Context) ([]string, error) - BulkStarSessions(sessionIDs []string) error + BulkStarSessions(sessionIDs []string) ([]string, error) // Pins. PinMessage(sessionID string, messageID int64, note *string) (int64, error) @@ -130,6 +134,7 @@ type Store interface { RenameSession(id string, displayName *string) error SoftDeleteSession(id string) error SoftDeleteSessions(ids []string) (int, error) + SoftDeleteSessionsReturningIDs(ids []string) ([]string, error) RestoreSession(id string) (int64, error) DeleteSessionIfTrashed(id string) (int64, error) ListTrashedSessions(ctx context.Context) ([]Session, error) diff --git a/internal/db/store_contract_test.go b/internal/db/store_contract_test.go index 5dae0d6d8..0860c731d 100644 --- a/internal/db/store_contract_test.go +++ b/internal/db/store_contract_test.go @@ -129,6 +129,7 @@ func TestStoreContract(t *testing.T) { {"stars_and_pins", contractStarsAndPins}, {"analytics_trends_and_usage", contractAnalyticsTrendsAndUsage}, {"local_only_methods", contractLocalOnlyMethods}, + {"machine_counts_and_conflicts", contractMachineCountsAndConflicts}, } for _, backend := range storeContractBackends() { @@ -243,6 +244,32 @@ func contractSessionsCursorFiltersAndDates( require.Equal(t, []string{"linux", "mac"}, machines) } +func contractMachineCountsAndConflicts( + t *testing.T, + store Store, + _ storeContractFixture, + _ storeContractBackend, +) { + t.Helper() + ctx := context.Background() + + counts, err := store.MachineSessionCounts(ctx) + require.NoError(t, err) + // The seed places non-deleted sessions on both machines. + require.Positive(t, counts["linux"]) + require.Positive(t, counts["mac"]) + for machine := range counts { + require.Contains(t, []string{"linux", "mac"}, machine, + "unexpected machine in session counts: %q", machine) + } + + // No conflicts are seeded; every backend reports zero (read-only + // mirrors do not carry the local metadata ledger at all). + conflicts, err := store.CountMetadataConflicts(ctx) + require.NoError(t, err) + require.Equal(t, 0, conflicts) +} + func contractMessagesOrderingAndToolResults( t *testing.T, store Store, @@ -392,13 +419,23 @@ func contractStarsAndPins( ok, err := store.StarSession(fixture.alphaID) require.NoError(t, err) require.True(t, ok) - require.NoError(t, store.BulkStarSessions([]string{fixture.gammaID, "missing-session"})) + bulkStarred, err := store.BulkStarSessions([]string{fixture.gammaID, "missing-session"}) + require.NoError(t, err) + require.Equal(t, []string{fixture.gammaID}, bulkStarred) + bulkStarred, err = store.BulkStarSessions([]string{fixture.alphaID, fixture.gammaID}) + require.NoError(t, err) + require.Empty(t, bulkStarred) stars, err := store.ListStarredSessionIDs(ctx) require.NoError(t, err) require.ElementsMatch(t, []string{fixture.alphaID, fixture.gammaID}, stars) - require.NoError(t, store.UnstarSession(fixture.gammaID)) + removed, err := store.UnstarSession(fixture.gammaID) + require.NoError(t, err) + require.True(t, removed) + removed, err = store.UnstarSession(fixture.gammaID) + require.NoError(t, err) + require.False(t, removed) stars, err = store.ListStarredSessionIDs(ctx) require.NoError(t, err) require.Equal(t, []string{fixture.alphaID}, stars) diff --git a/internal/db/sync_state_test.go b/internal/db/sync_state_test.go new file mode 100644 index 000000000..6ae05b6aa --- /dev/null +++ b/internal/db/sync_state_test.go @@ -0,0 +1,94 @@ +package db + +import ( + "fmt" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSyncStateValuesReadsOnlyExactKeysAcrossBatches(t *testing.T) { + database := testDB(t) + for i := range 905 { + key := fmt.Sprintf("artifact_import:desk:desk~%04d", i) + require.NoError(t, database.SetSyncState(key, fmt.Sprintf("hash-%04d", i))) + } + require.NoError(t, database.SetSyncState("unrelated", "keep-out")) + + got, err := database.SyncStateValues([]string{ + "artifact_import:desk:desk~0000", + "artifact_import:desk:desk~0899", + "artifact_import:desk:desk~0904", + "missing", + }) + require.NoError(t, err) + assert.Equal(t, map[string]string{ + "artifact_import:desk:desk~0000": "hash-0000", + "artifact_import:desk:desk~0899": "hash-0899", + "artifact_import:desk:desk~0904": "hash-0904", + }, got) +} + +func TestCopySyncStatePreservesArtifactImportQueue(t *testing.T) { + sourcePath := filepath.Join(t.TempDir(), "source.db") + source, err := Open(sourcePath) + require.NoError(t, err) + meta := ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", + Name: artifactImportMetadataName("a"), + SHA256: strings.Repeat("a", 64), Size: 11, + Reason: "future metadata", RequiredFormatVersion: 2, + } + checkpoint5 := ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "checkpoints", Name: "cp-0000000005.json", + SHA256: strings.Repeat("5", 64), Size: 55, + Reason: "missing segment", RequiredFormatVersion: 1, + } + require.NoError(t, source.EnqueueArtifactImport(t.Context(), meta)) + require.NoError(t, source.EnqueueArtifactImport(t.Context(), checkpoint5)) + require.NoError(t, source.Close()) + + destination := testDB(t) + require.NoError(t, destination.EnqueueArtifactImport(t.Context(), ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "checkpoints", Name: "cp-0000000004.json", + SHA256: strings.Repeat("4", 64), Size: 44, + Reason: "older checkpoint", RequiredFormatVersion: 1, + })) + require.NoError(t, destination.CopySyncStateFrom(sourcePath)) + + pending, err := destination.PendingArtifactImports(t.Context(), 2, 10) + require.NoError(t, err) + require.Len(t, pending, 2) + var names []string + for _, work := range pending { + names = append(names, work.Name) + } + assert.ElementsMatch(t, []string{meta.Name, checkpoint5.Name}, names) +} + +func TestCopySyncStateRejectsArtifactImportIdentityConflict(t *testing.T) { + sourcePath := filepath.Join(t.TempDir(), "source.db") + source, err := Open(sourcePath) + require.NoError(t, err) + work := ArtifactImportWork{ + Origin: "peer-a1b2c3", Kind: "meta", + Name: artifactImportMetadataName("a"), + SHA256: strings.Repeat("a", 64), Size: 11, + Reason: "source", RequiredFormatVersion: 1, + } + require.NoError(t, source.EnqueueArtifactImport(t.Context(), work)) + require.NoError(t, source.Close()) + + destination := testDB(t) + _, err = destination.getWriter().Exec(` + INSERT INTO artifact_import_queue( + origin, kind, name, sha256, size, reason, required_format_version + ) VALUES (?, ?, ?, ?, ?, ?, ?)`, work.Origin, work.Kind, work.Name, + strings.Repeat("b", 64), work.Size, "destination", 1) + require.NoError(t, err) + + require.ErrorContains(t, destination.CopySyncStateFrom(sourcePath), "conflicting identity") +} diff --git a/internal/db/usage_events.go b/internal/db/usage_events.go index ef0f34878..86112e025 100644 --- a/internal/db/usage_events.go +++ b/internal/db/usage_events.go @@ -74,7 +74,7 @@ func (db *DB) ReplaceSessionUsageEvents( } defer func() { _ = tx.Rollback() }() - if err := replaceSessionUsageEventsTx(tx, sessionID, events); err != nil { + if err := replaceSessionUsageEventsTx(tx, sessionID, events, true); err != nil { return err } @@ -82,7 +82,7 @@ func (db *DB) ReplaceSessionUsageEvents( } func replaceSessionUsageEventsTx( - tx *sql.Tx, sessionID string, events []UsageEvent, + tx *sql.Tx, sessionID string, events []UsageEvent, enqueueArtifact bool, ) error { if _, err := tx.Exec( `DELETE FROM usage_events WHERE session_id = ?`, @@ -147,6 +147,11 @@ func replaceSessionUsageEventsTx( sessionID, err, ) } + if enqueueArtifact { + if err := enqueueArtifactExportTx(tx, sessionID); err != nil { + return err + } + } return nil } diff --git a/internal/duckdb/curation.go b/internal/duckdb/curation.go index d0bfd5ab3..c51fa9330 100644 --- a/internal/duckdb/curation.go +++ b/internal/duckdb/curation.go @@ -12,8 +12,8 @@ func (s *Store) StarSession(sessionID string) (bool, error) { return false, db.ErrReadOnly } -func (s *Store) UnstarSession(sessionID string) error { - return db.ErrReadOnly +func (s *Store) UnstarSession(sessionID string) (bool, error) { + return false, db.ErrReadOnly } func (s *Store) ListStarredSessionIDs(ctx context.Context) ([]string, error) { @@ -35,8 +35,8 @@ func (s *Store) ListStarredSessionIDs(ctx context.Context) ([]string, error) { return ids, rows.Err() } -func (s *Store) BulkStarSessions(sessionIDs []string) error { - return db.ErrReadOnly +func (s *Store) BulkStarSessions(sessionIDs []string) ([]string, error) { + return nil, db.ErrReadOnly } func (s *Store) PinMessage(sessionID string, messageID int64, note *string) (int64, error) { diff --git a/internal/duckdb/messages.go b/internal/duckdb/messages.go index 5f9e78e7f..a09781e70 100644 --- a/internal/duckdb/messages.go +++ b/internal/duckdb/messages.go @@ -3,6 +3,7 @@ package duckdb import ( "context" "database/sql" + "errors" "fmt" "slices" "strings" @@ -267,6 +268,25 @@ func (s *Store) GetResumeModelCounts( return counts, nil } +// GetMessageForMetadataPin returns only the stable message identity fields +// needed for metadata pin events. +func (s *Store) GetMessageForMetadataPin( + ctx context.Context, sessionID string, messageID int64, +) (*db.Message, error) { + row := s.queryRowContext(ctx, ` + SELECT id, session_id, ordinal, COALESCE(source_uuid, '') + FROM messages + WHERE session_id = ? AND id = ?`, sessionID, messageID) + var msg db.Message + if err := row.Scan(&msg.ID, &msg.SessionID, &msg.Ordinal, &msg.SourceUUID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, fmt.Errorf("querying message metadata for pin: %w", err) + } + return &msg, nil +} + func scanMessages(rows *sql.Rows) ([]db.Message, error) { var msgs []db.Message for rows.Next() { diff --git a/internal/duckdb/metadata.go b/internal/duckdb/metadata.go new file mode 100644 index 000000000..dd2736052 --- /dev/null +++ b/internal/duckdb/metadata.go @@ -0,0 +1,22 @@ +package duckdb + +import ( + "context" + + "go.kenn.io/agentsview/internal/db" +) + +// ListMetadataConflicts returns no rows for DuckDB read mode because the local +// artifact metadata ledger is not part of the analytical mirror. +func (s *Store) ListMetadataConflicts( + context.Context, + []string, +) ([]db.MetadataConflict, error) { + return []db.MetadataConflict{}, nil +} + +// CountMetadataConflicts returns zero for DuckDB read mode because the local +// artifact metadata ledger is not part of the analytical mirror. +func (s *Store) CountMetadataConflicts(context.Context) (int, error) { + return 0, nil +} diff --git a/internal/duckdb/push_bounded_test.go b/internal/duckdb/push_bounded_test.go index 6ce4e4230..5f11b5843 100644 --- a/internal/duckdb/push_bounded_test.go +++ b/internal/duckdb/push_bounded_test.go @@ -350,7 +350,9 @@ func TestReplaceCurationBoundedByLocalCurationSizeNotMirrorSize(t *testing.T) { assertMirrorTableCountWhere(t, path, "starred_sessions", "session_id = ?", "sess-1", 1) assertMirrorTableCountWhere(t, path, "pinned_messages", "session_id = ?", "sess-2", 1) - require.NoError(t, local.UnstarSession("sess-1")) + removed, err := local.UnstarSession("sess-1") + require.NoError(t, err) + require.True(t, removed) require.NoError(t, local.UnpinMessage("sess-2", msgs[0].ID)) // A mutating incremental push (a required precondition for // replaceCuration to be worth asserting on): appending a message diff --git a/internal/duckdb/rebuild_test.go b/internal/duckdb/rebuild_test.go index 89a65c944..2cdf8fccb 100644 --- a/internal/duckdb/rebuild_test.go +++ b/internal/duckdb/rebuild_test.go @@ -494,7 +494,9 @@ func TestCurationToggleRevertRaceLeavesMirrorConsistent(t *testing.T) { ok, err := local.StarSession(ids[1]) require.NoError(t, err) require.True(t, ok) - require.NoError(t, local.UnstarSession(ids[1])) + removed, err := local.UnstarSession(ids[1]) + require.NoError(t, err) + require.True(t, removed) written, err := s.replaceCuration(ctx, snap) require.NoError(t, err) diff --git a/internal/duckdb/store.go b/internal/duckdb/store.go index dcc16c9ad..b3ee173e3 100644 --- a/internal/duckdb/store.go +++ b/internal/duckdb/store.go @@ -789,6 +789,31 @@ func (s *Store) GetAgents(ctx context.Context, excludeOneShot, excludeAutomated return out, rows.Err() } +// MachineSessionCounts returns the number of non-deleted sessions per machine, +// keyed by machine name. +func (s *Store) MachineSessionCounts(ctx context.Context) (map[string]int, error) { + rows, err := s.duck.QueryContext(ctx, + `SELECT machine, COUNT(*) FROM sessions + WHERE deleted_at IS NULL + GROUP BY machine`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + counts := map[string]int{} + for rows.Next() { + var machine string + var count int + if err := rows.Scan(&machine, &count); err != nil { + return nil, err + } + counts[machine] = count + } + return counts, rows.Err() +} + func (s *Store) GetMachines(ctx context.Context, excludeOneShot, excludeAutomated bool) ([]string, error) { rows, err := s.queryContext(ctx, `SELECT DISTINCT machine FROM sessions WHERE `+ diff --git a/internal/duckdb/store_contract_test.go b/internal/duckdb/store_contract_test.go index 2ace1b1ed..bf441195d 100644 --- a/internal/duckdb/store_contract_test.go +++ b/internal/duckdb/store_contract_test.go @@ -308,6 +308,15 @@ func duckContractSessionsCursorsAndMetadata( machines, err := store.GetMachines(ctx, false, false) require.NoError(t, err) require.Equal(t, []string{"test-machine"}, machines) + + counts, err := store.MachineSessionCounts(ctx) + require.NoError(t, err) + require.Len(t, counts, 1) + require.Positive(t, counts["test-machine"]) + + conflicts, err := store.CountMetadataConflicts(ctx) + require.NoError(t, err) + require.Equal(t, 0, conflicts) } func assertDuckJSONTranscriptRevision(t *testing.T, value any, want string) { @@ -403,8 +412,10 @@ func duckContractReadOnlyCuration( ok, err := store.StarSession(fixture.betaID) require.ErrorIs(t, err, db.ErrReadOnly) require.False(t, ok) - require.ErrorIs(t, store.UnstarSession(fixture.alphaID), db.ErrReadOnly) - require.ErrorIs(t, store.BulkStarSessions([]string{fixture.betaID}), db.ErrReadOnly) + _, err = store.UnstarSession(fixture.alphaID) + require.ErrorIs(t, err, db.ErrReadOnly) + _, bulkErr := store.BulkStarSessions([]string{fixture.betaID}) + require.ErrorIs(t, bulkErr, db.ErrReadOnly) pinID, err := store.PinMessage(fixture.alphaID, 1, nil) require.ErrorIs(t, err, db.ErrReadOnly) diff --git a/internal/duckdb/store_test.go b/internal/duckdb/store_test.go index 95f398c4c..971f4d979 100644 --- a/internal/duckdb/store_test.go +++ b/internal/duckdb/store_test.go @@ -859,8 +859,10 @@ func TestStoreCurationMethods(t *testing.T) { ok, err := store.StarSession(fixture.betaID) require.ErrorIs(t, err, db.ErrReadOnly) assert.False(t, ok) - require.ErrorIs(t, store.BulkStarSessions([]string{fixture.betaID}), db.ErrReadOnly) - require.ErrorIs(t, store.UnstarSession(fixture.alphaID), db.ErrReadOnly) + _, bulkErr := store.BulkStarSessions([]string{fixture.betaID}) + require.ErrorIs(t, bulkErr, db.ErrReadOnly) + _, err = store.UnstarSession(fixture.alphaID) + require.ErrorIs(t, err, db.ErrReadOnly) starred, err = store.ListStarredSessionIDs(ctx) require.NoError(t, err) assert.Equal(t, []string{fixture.alphaID}, starred) diff --git a/internal/duckdb/stubs.go b/internal/duckdb/stubs.go index 21cd2e0cf..1907a438f 100644 --- a/internal/duckdb/stubs.go +++ b/internal/duckdb/stubs.go @@ -15,9 +15,12 @@ func (s *Store) GetInsight(_ context.Context, _ int64) (*db.Insight, error) { re func (s *Store) GetCachedInsight(_ context.Context, _ string) (*db.Insight, error) { return nil, nil } -func (s *Store) RenameSession(_ string, _ *string) error { return db.ErrReadOnly } -func (s *Store) SoftDeleteSession(_ string) error { return db.ErrReadOnly } -func (s *Store) SoftDeleteSessions(_ []string) (int, error) { return 0, db.ErrReadOnly } +func (s *Store) RenameSession(_ string, _ *string) error { return db.ErrReadOnly } +func (s *Store) SoftDeleteSession(_ string) error { return db.ErrReadOnly } +func (s *Store) SoftDeleteSessions(_ []string) (int, error) { return 0, db.ErrReadOnly } +func (s *Store) SoftDeleteSessionsReturningIDs(_ []string) ([]string, error) { + return nil, db.ErrReadOnly +} func (s *Store) RestoreSession(_ string) (int64, error) { return 0, db.ErrReadOnly } func (s *Store) DeleteSessionIfTrashed(_ string) (int64, error) { return 0, db.ErrReadOnly } func (s *Store) EmptyTrash() (int, error) { return 0, db.ErrReadOnly } diff --git a/internal/e2e/artifact_sync_test.go b/internal/e2e/artifact_sync_test.go new file mode 100644 index 000000000..63544651e --- /dev/null +++ b/internal/e2e/artifact_sync_test.go @@ -0,0 +1,379 @@ +//go:build e2e + +package e2e + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "net/url" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/dbtest" + "go.kenn.io/agentsview/internal/parser" + "go.kenn.io/agentsview/internal/server" + syncpkg "go.kenn.io/agentsview/internal/sync" +) + +const e2eToken = "artifact-e2e-token" + +func TestArtifactSyncTwoInstanceFolderAndHTTP(t *testing.T) { + ctx := context.Background() + root := preservedWorkspace(t) + share := filepath.Join(root, "share") + require.NoError(t, os.MkdirAll(share, 0o755)) + + a := openE2ENode(t, filepath.Join(root, "node-a"), "laptop-a1b2c3") + defer a.Close() + b := openE2ENode(t, filepath.Join(root, "node-b"), "desktop-d4e5f6") + defer b.Close() + + seedE2ESession(t, a.DB, "sess-1", "alpha", "world") + + folderSync(t, a, share) + folderSync(t, b, share) + + importedID := a.Origin + "~sess-1" + assertSessionProject(t, b.DB, importedID, "alpha") + assertSearchFinds(t, b.DB, "world", importedID) + + renameWithMetadata(t, a, "sess-1", "Alpha from laptop") + folderSync(t, a, share) + folderSync(t, b, share) + assertSessionDisplayName(t, b.DB, importedID, "Alpha from laptop") + + renameWithMetadata(t, b, importedID, "Bravo from desktop") + folderSync(t, b, share) + folderSync(t, a, share) + assertSessionDisplayName(t, a.DB, "sess-1", "Bravo from desktop") + + b.Close() + b = openE2ENode(t, filepath.Join(root, "node-b"), "desktop-d4e5f6") + defer b.Close() + + require.NoError(t, a.DB.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + {SessionID: "sess-1", Ordinal: 1, Role: "assistant", Content: "planet", ContentLength: 6}, + })) + _, err := artifact.ExportToStore(ctx, a.DB, a.Repository.Content(), artifact.ExportOptions{ + Origin: a.Origin, + Full: true, + }) + require.NoError(t, err) + postOriginArtifacts(t, a, b, a.Origin) + + assertMessagesContain(t, b.DB, importedID, "planet") + assertSearchFinds(t, b.DB, "planet", importedID) + + first := writeForeignRename(t, b, "writer-a1b2c3", importedID, "Fork one") + second := writeForeignRename(t, b, "writer-b1b2c3", importedID, "Fork two") + coordinator := artifact.NewStoreImportCoordinator( + b.DB, b.Repository.Content(), b.Origin, + ) + require.NoError(t, coordinator.RecordChanged(ctx, first)) + require.NoError(t, coordinator.RecordChanged(ctx, second)) + imported, err := coordinator.Finalize(ctx) + require.NoError(t, err) + assert.NotZero(t, imported.Metadata) + + conflicts, err := b.DB.ListMetadataConflicts(ctx, []string{importedID}) + require.NoError(t, err) + require.NotEmpty(t, conflicts) + assert.Equal(t, "display_name", conflicts[0].Field) + + apiConflicts := getMetadataConflicts(t, b, importedID) + require.NotEmpty(t, apiConflicts.Conflicts) + assert.Equal(t, importedID, apiConflicts.Conflicts[0].SessionGID) +} + +type e2eNode struct { + DataDir string + DBPath string + Origin string + DB *db.DB + Server *httptest.Server + App *server.Server + Repository *artifact.Repository +} + +func openE2ENode(t *testing.T, dataDir, origin string) *e2eNode { + t.Helper() + require.NoError(t, os.MkdirAll(dataDir, 0o755)) + dbPath := filepath.Join(dataDir, "sessions.db") + database, err := db.Open(dbPath) + require.NoError(t, err) + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + + emptyAgentDir := filepath.Join(dataDir, "empty-agent-dir") + require.NoError(t, os.MkdirAll(emptyAgentDir, 0o755)) + broadcaster := server.NewBroadcaster(0) + engine := syncpkg.NewEngine(database, syncpkg.EngineConfig{ + AgentDirs: map[parser.AgentType][]string{ + parser.AgentClaude: {emptyAgentDir}, + }, + Machine: origin, + Emitter: broadcaster, + }) + cfg := config.Config{ + Host: "127.0.0.1", + Port: 0, + DataDir: dataDir, + DBPath: dbPath, + WriteTimeout: 30 * time.Second, + RequireAuth: true, + AuthToken: e2eToken, + ArtifactOriginID: origin, + } + srv := server.New(cfg, database, engine, + server.WithBroadcaster(broadcaster), + server.WithArtifactStore(repository.Content()), + ) + return &e2eNode{ + DataDir: dataDir, + DBPath: dbPath, + Origin: origin, + DB: database, + Server: httptest.NewServer(srv.Handler()), + App: srv, + Repository: repository, + } +} + +func (n *e2eNode) Close() { + if n == nil { + return + } + if n.Server != nil { + n.Server.Close() + n.Server = nil + } + if n.App != nil { + _ = n.App.Shutdown(context.Background()) + n.App = nil + } + if n.Repository != nil { + _ = n.Repository.Close() + n.Repository = nil + } + if n.DB != nil { + n.DB.Close() + n.DB = nil + } +} + +func preservedWorkspace(t *testing.T) string { + t.Helper() + root, err := os.MkdirTemp("", "agentsview-artifact-e2e-*") + require.NoError(t, err) + t.Cleanup(func() { + if t.Failed() { + t.Logf("preserved artifact sync e2e workspace: %s", root) + return + } + require.NoError(t, os.RemoveAll(root)) + }) + return root +} + +func seedE2ESession(t *testing.T, database *db.DB, id, project, assistantText string) { + t.Helper() + started := "2026-06-14T01:02:03Z" + ended := "2026-06-14T01:03:03Z" + first := "hello" + dbtest.SeedSession(t, database, id, project, func(s *db.Session) { + s.MessageCount = 2 + s.UserMessageCount = 1 + s.FirstMessage = &first + s.StartedAt = &started + s.EndedAt = &ended + }) + require.NoError(t, database.ReplaceSessionMessages(id, []db.Message{ + {SessionID: id, Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + {SessionID: id, Ordinal: 1, Role: "assistant", Content: assistantText, ContentLength: len(assistantText)}, + })) +} + +func folderSync(t *testing.T, n *e2eNode, share string) artifact.SyncResult { + t.Helper() + res, err := artifact.SyncWithRepository(context.Background(), n.DB, n.Repository, artifact.SyncOptions{ + DataDir: n.DataDir, + Target: share, + Origin: n.Origin, + }) + require.NoError(t, err) + return res +} + +func renameWithMetadata(t *testing.T, n *e2eNode, sessionID, displayName string) { + t.Helper() + require.NoError(t, n.DB.RenameSession(sessionID, &displayName)) + // A local edit records its own replay register in this node's db, exactly as + // the rename handler does. + appendRenameArtifact(t, n.DB, n.Repository.Content(), n.Origin, sessionID, displayName) +} + +func writeForeignRename( + t *testing.T, n *e2eNode, origin, sessionID, displayName string, +) artifact.Entry { + t.Helper() + // A foreign origin's event arrives as an artifact file written by another + // machine. Record it through a throwaway db so it is not pre-marked applied + // in this node's db, leaving the real import to replay it. + scratch, err := db.Open(filepath.Join(t.TempDir(), "scratch.db")) + require.NoError(t, err) + t.Cleanup(func() { scratch.Close() }) + return appendRenameArtifact(t, scratch, n.Repository.Content(), origin, sessionID, displayName) +} + +func appendRenameArtifact( + t *testing.T, database *db.DB, store artifact.ArtifactStore, origin, sessionID, displayName string, +) artifact.Entry { + t.Helper() + value, err := json.Marshal(struct { + DisplayName *string `json:"display_name"` + }{DisplayName: &displayName}) + require.NoError(t, err) + recorder := artifact.NewMetadataRecorder(database, artifact.MetadataRecorderOptions{ + Store: store, + Origin: origin, + }) + record, err := recorder.Append(context.Background(), artifact.MetadataEventInput{ + SessionID: sessionID, + Op: artifact.MetadataOpRename, + Value: json.RawMessage(value), + }) + require.NoError(t, err) + entry, err := store.Stat(t.Context(), record.Ref) + require.NoError(t, err) + return entry +} + +func postOriginArtifacts(t *testing.T, from, to *e2eNode, origin string) { + t.Helper() + for _, kind := range []artifact.Kind{ + artifact.KindSegments, + artifact.KindRaw, + artifact.KindManifests, + artifact.KindMeta, + artifact.KindCheckpoints, + } { + iterator, err := from.Repository.Content().Entries(t.Context(), origin, kind) + require.NoError(t, err) + for { + entries, nextErr := iterator.Next(t.Context(), 100) + require.True(t, nextErr == nil || errors.Is(nextErr, io.EOF)) + for _, entry := range entries { + _, reader, err := from.Repository.Content().Open(t.Context(), entry.Ref) + require.NoError(t, err) + wire, err := artifact.ToWireRef(entry.Ref) + require.NoError(t, err) + var body bytes.Buffer + require.NoError(t, artifact.EncodeWire(t.Context(), entry.Ref, reader, &body)) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + postArtifact(t, to, wire, body.Bytes()) + } + if errors.Is(nextErr, io.EOF) { + break + } + } + require.NoError(t, iterator.Close()) + } +} + +func postArtifact(t *testing.T, to *e2eNode, wire artifact.WireRef, body []byte) { + t.Helper() + req, err := http.NewRequest( + http.MethodPost, + to.Server.URL+"/api/v1/artifacts/"+wire.Origin+"/"+string(wire.Kind)+"/"+url.PathEscape(wire.Name), + bytes.NewReader(body), + ) + require.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+e2eToken) + req.Header.Set("Content-Type", "application/octet-stream") + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) +} + +type metadataConflictsResponse struct { + Conflicts []db.MetadataConflict `json:"conflicts"` +} + +func getMetadataConflicts(t *testing.T, n *e2eNode, sessionID string) metadataConflictsResponse { + t.Helper() + req, err := http.NewRequest( + http.MethodGet, + n.Server.URL+"/api/v1/sessions/"+url.PathEscape(sessionID)+"/metadata-conflicts", + nil, + ) + require.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+e2eToken) + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) + var out metadataConflictsResponse + require.NoError(t, json.NewDecoder(resp.Body).Decode(&out)) + return out +} + +func assertSessionProject(t *testing.T, database *db.DB, sessionID, project string) { + t.Helper() + sess, err := database.GetSessionFull(context.Background(), sessionID) + require.NoError(t, err) + require.NotNil(t, sess) + assert.Equal(t, project, sess.Project) +} + +func assertSessionDisplayName(t *testing.T, database *db.DB, sessionID, displayName string) { + t.Helper() + sess, err := database.GetSessionFull(context.Background(), sessionID) + require.NoError(t, err) + require.NotNil(t, sess) + require.NotNil(t, sess.DisplayName) + assert.Equal(t, displayName, *sess.DisplayName) +} + +func assertMessagesContain(t *testing.T, database *db.DB, sessionID, text string) { + t.Helper() + msgs, err := database.GetAllMessages(context.Background(), sessionID) + require.NoError(t, err) + for _, msg := range msgs { + if msg.Content == text { + return + } + } + require.Fail(t, fmt.Sprintf("session %s messages did not contain %q", sessionID, text)) +} + +func assertSearchFinds(t *testing.T, database *db.DB, query, sessionID string) { + t.Helper() + page, err := database.Search(context.Background(), db.SearchFilter{ + Query: query, + Limit: 10, + }) + require.NoError(t, err) + for _, result := range page.Results { + if result.SessionID == sessionID { + return + } + } + require.Fail(t, fmt.Sprintf("search %q did not find %s", query, sessionID)) +} diff --git a/internal/postgres/collision_pgtest_test.go b/internal/postgres/collision_pgtest_test.go index e1ea596d8..415b7815b 100644 --- a/internal/postgres/collision_pgtest_test.go +++ b/internal/postgres/collision_pgtest_test.go @@ -4,15 +4,914 @@ package postgres import ( "context" + "database/sql" + "fmt" "path/filepath" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/db" ) +func openArtifactTestRepository(t *testing.T, ctx context.Context) *artifact.Repository { + t.Helper() + repository, err := artifact.OpenRepository(ctx, t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + return repository +} + +func importExportedArtifactCheckpoint( + t *testing.T, + ctx context.Context, + database *db.DB, + repository *artifact.Repository, + sourceOrigin string, + localOrigin string, + export artifact.ExportResult, +) artifact.ImportResult { + t.Helper() + ref, err := artifact.NewRef(sourceOrigin, artifact.KindCheckpoints, + fmt.Sprintf("cp-%010d.json", export.CheckpointSequence)) + require.NoError(t, err) + entry, err := repository.Content().Stat(ctx, ref) + require.NoError(t, err) + coordinator := artifact.NewStoreImportCoordinator(database, repository.Content(), localOrigin) + require.NoError(t, coordinator.RecordChanged(ctx, entry)) + result, err := coordinator.Finalize(ctx) + require.NoError(t, err) + return result +} + +func TestPushArtifactNativeAndImportedCopiesShareOriginIdentity(t *testing.T) { + pgURL := testPGURL(t) + ctx := context.Background() + const originA = "origin-a1b2c3" + const originB = "origin-b4c5d6" + const nativeID = "native-id" + const canonicalID = originA + "~" + nativeID + const childID = "child-id" + const canonicalChildID = originA + "~" + childID + + for _, tc := range []struct { + name string + schema string + importerFirst bool + }{ + { + name: "origin pushes first", + schema: "agentsview_artifact_identity_origin_first_test", + }, + { + name: "importer pushes first", + schema: "agentsview_artifact_identity_importer_first_test", + importerFirst: true, + }, + } { + t.Run(tc.name, func(t *testing.T) { + pg, err := Open(pgURL, tc.schema, true) + require.NoError(t, err, "Open") + t.Cleanup(func() { require.NoError(t, pg.Close()) }) + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + tc.schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, tc.schema), "EnsureSchema") + + originDB, err := db.Open(filepath.Join(t.TempDir(), "origin.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, originDB.Close()) }) + require.NoError(t, artifact.AdoptOrigin(originDB, originA)) + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: nativeID, Project: "alpha", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, originDB.ReplaceSessionMessages(nativeID, []db.Message{{ + SessionID: nativeID, Ordinal: 0, Role: "user", + Content: "hello", ContentLength: 5, + }})) + parentID := nativeID + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: childID, Project: "alpha", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:01Z", ParentSessionID: &parentID, + })) + require.NoError(t, originDB.ReplaceSessionMessages(childID, []db.Message{{ + SessionID: childID, Ordinal: 0, Role: "user", + Content: "child", ContentLength: 5, + }})) + + artifactRepository := openArtifactTestRepository(t, ctx) + exportResult, err := artifact.ExportToStore(ctx, originDB, artifactRepository.Content(), artifact.ExportOptions{ + Origin: originA, + Full: true, + }) + require.NoError(t, err) + importerDB, err := db.Open(filepath.Join(t.TempDir(), "importer.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, importerDB.Close()) }) + require.NoError(t, artifact.AdoptOrigin(importerDB, originB)) + importResult := importExportedArtifactCheckpoint( + t, ctx, importerDB, artifactRepository, originA, originB, exportResult, + ) + require.Equal(t, 2, importResult.Sessions) + require.Equal(t, 2, importResult.Messages) + + // Advance the origin after the importer captured its artifact snapshot. + // An origin-first push must not be rolled back when that stale replica + // subsequently pushes the same canonical owner marker. + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: nativeID, Project: "fresh-origin", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, originDB.ReplaceSessionMessages(nativeID, []db.Message{{ + SessionID: nativeID, Ordinal: 0, Role: "user", + Content: "fresh", ContentLength: 5, + }})) + + originSync := &Sync{ + pg: pg, local: originDB, machine: "host-a", + schema: tc.schema, schemaDone: true, + } + importerSync := &Sync{ + pg: pg, local: importerDB, machine: "host-b", + schema: tc.schema, schemaDone: true, + } + pushers := []*Sync{originSync, importerSync} + if tc.importerFirst { + pushers[0], pushers[1] = pushers[1], pushers[0] + } + for _, pusher := range pushers { + result, pushErr := pusher.Push(ctx, false, nil) + require.NoError(t, pushErr) + assert.Zero(t, result.Errors) + assert.Zero(t, result.SkippedConflicts) + } + + var id, machine, ownerMarker string + err = pg.QueryRowContext(ctx, ` + SELECT id, machine, owner_marker + FROM sessions + WHERE id IN ($1, $2) + `, nativeID, canonicalID).Scan(&id, &machine, &ownerMarker) + require.NoError(t, err) + assert.Equal(t, canonicalID, id) + assert.Equal(t, originA, machine) + assert.Equal(t, artifactOwnerMarkerPrefix+originA, ownerMarker) + var project, content string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT s.project, m.content + FROM sessions s + JOIN messages m ON m.session_id = s.id + WHERE s.id = $1 AND m.ordinal = 0 + `, canonicalID).Scan(&project, &content)) + assert.Equal(t, "fresh-origin", project) + assert.Equal(t, "fresh", content) + var parent string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT parent_session_id FROM sessions WHERE id = $1 + `, canonicalChildID).Scan(&parent)) + assert.Equal(t, canonicalID, parent, + "artifact relationships must resolve through the canonical origin identity") + + for _, pusher := range []*Sync{importerSync, originSync} { + result, pushErr := pusher.Push(ctx, true, nil) + require.NoError(t, pushErr) + assert.Zero(t, result.Errors) + assert.Zero(t, result.SkippedConflicts) + } + var count int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM sessions WHERE id IN ($1, $2) + `, nativeID, canonicalID).Scan(&count)) + assert.Equal(t, 1, count, + "native and imported copies must keep one PG row across repeated pushes") + }) + } +} + +func TestPushSSHShapedSessionDoesNotAdoptArtifactOwnership(t *testing.T) { + pgURL := testPGURL(t) + ctx := context.Background() + const schema = "agentsview_artifact_provenance_guard_test" + const remoteOrigin = "origin-a1b2c3" + const localOrigin = "origin-b4c5d6" + const sessionID = remoteOrigin + "~native-id" + + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + t.Cleanup(func() { require.NoError(t, pg.Close()) }) + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, localDB.Close()) }) + require.NoError(t, artifact.AdoptOrigin(localDB, localOrigin)) + require.NoError(t, localDB.UpsertSession(db.Session{ + ID: sessionID, Project: "ssh-copy", Machine: remoteOrigin, Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, localDB.ReplaceSessionMessages(sessionID, []db.Message{{ + SessionID: sessionID, Ordinal: 0, Role: "user", + Content: "ssh copy", ContentLength: 8, + }})) + + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions ( + id, machine, owner_marker, project, agent, created_at + ) VALUES ($1, $2, $3, $4, $5, NOW()) + `, sessionID, remoteOrigin, artifactOwnerMarkerPrefix+remoteOrigin, + "genuine-artifact", "claude") + require.NoError(t, err, "seed genuine artifact-owned row") + + syncer := &Sync{ + pg: pg, local: localDB, machine: "ssh-importer", + schema: schema, schemaDone: true, + } + result, err := syncer.Push(ctx, false, nil) + require.NoError(t, err) + assert.Equal(t, 1, result.SkippedConflicts, + "an SSH-shaped row without artifact import provenance must retain legacy collision protection") + assert.Zero(t, result.SessionsPushed) + + var project, ownerMarker string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT project, owner_marker FROM sessions WHERE id = $1 + `, sessionID).Scan(&project, &ownerMarker)) + assert.Equal(t, "genuine-artifact", project) + assert.Equal(t, artifactOwnerMarkerPrefix+remoteOrigin, ownerMarker) +} + +func TestPushArtifactOriginAdoptionReusesLegacyBareRows(t *testing.T) { + pgURL := testPGURL(t) + ctx := context.Background() + const schema = "agentsview_artifact_origin_upgrade_test" + const originA = "origin-a1b2c3" + const originB = "origin-b4c5d6" + const parentID = "native-parent" + const childID = "native-child" + const canonicalParentID = originA + "~" + parentID + const canonicalChildID = originA + "~" + childID + const stateScope = "work" + + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + t.Cleanup(func() { require.NoError(t, pg.Close()) }) + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + originDB, err := db.Open(filepath.Join(t.TempDir(), "origin.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, originDB.Close()) }) + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: parentID, Project: "alpha", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, originDB.ReplaceSessionMessages(parentID, []db.Message{{ + SessionID: parentID, Ordinal: 0, Role: "user", + Content: "parent", ContentLength: 6, SourceUUID: "parent-source", + }})) + parentRef := parentID + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: childID, Project: "alpha", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:01Z", ParentSessionID: &parentRef, + })) + require.NoError(t, originDB.ReplaceSessionMessages(childID, []db.Message{{ + SessionID: childID, Ordinal: 0, Role: "user", + Content: "child", ContentLength: 5, SourceUUID: "child-source", + }})) + + originSync := &Sync{ + pg: pg, local: originDB, machine: "host-a", + schema: schema, schemaDone: true, + syncState: newScopedSyncStateStore(originDB, stateScope, false), + syncStateTarget: stateScope, + } + first, err := originSync.Push(ctx, false, nil) + require.NoError(t, err) + require.Equal(t, 2, first.SessionsPushed) + watermark, err := originDB.GetSyncState("last_push_at:" + stateScope) + require.NoError(t, err) + require.NotEmpty(t, watermark, "precondition: initial push established incremental state") + // Simulate a pusher upgraded from a build predating identity-mode state. + require.NoError(t, originDB.SetSyncState( + "pg_artifact_identity_v1:"+stateScope, "", + )) + + _, err = pg.ExecContext(ctx, + `UPDATE sessions SET display_name = 'PG title' WHERE id = $1`, + parentID, + ) + require.NoError(t, err, "seed PG-local display name") + _, err = pg.ExecContext(ctx, + `INSERT INTO starred_sessions (session_id) VALUES ($1)`, + parentID, + ) + require.NoError(t, err, "seed PG-local star") + _, err = pg.ExecContext(ctx, ` + INSERT INTO pinned_messages ( + session_id, message_id, ordinal, source_uuid, note + ) + SELECT $1, ordinal, ordinal, COALESCE(source_uuid, ''), 'PG pin' + FROM messages WHERE session_id = $1 AND ordinal = 0 + `, parentID) + require.NoError(t, err, "seed PG-local pin") + + before, err := originDB.GetSessionFull(ctx, parentID) + require.NoError(t, err) + require.NotNil(t, before) + require.NoError(t, artifact.AdoptOrigin(originDB, originA)) + afterAdopt, err := originDB.GetSessionFull(ctx, parentID) + require.NoError(t, err) + require.NotNil(t, afterAdopt) + assert.Equal(t, before.LocalModifiedAt, afterAdopt.LocalModifiedAt, + "origin adoption must not need to mutate session timestamps") + + upgraded, err := originSync.Push(ctx, false, nil) + require.NoError(t, err) + assert.Equal(t, 2, upgraded.SessionsPushed, + "identity-mode change must revisit unchanged sessions") + assert.Zero(t, upgraded.Errors) + assert.Zero(t, upgraded.SkippedConflicts) + + var stableBareRows int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM sessions + WHERE id IN ($1, $2) AND machine = $3 AND owner_marker = $4 + `, parentID, childID, originA, artifactOwnerMarkerPrefix+originA).Scan(&stableBareRows)) + assert.Equal(t, 2, stableBareRows, + "same-owner legacy bare rows must upgrade in place to stable artifact ownership") + var canonicalRows int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM sessions WHERE id IN ($1, $2) + `, canonicalParentID, canonicalChildID).Scan(&canonicalRows)) + assert.Zero(t, canonicalRows, "upgrade must not migrate primary keys or duplicate rows") + assert.Equal(t, parentID, pgParentSessionID(t, ctx, pg, childID)) + assertPGArtifactUpgradeCuration(t, ctx, pg, parentID) + mode, err := originDB.GetSyncState("pg_artifact_identity_v1:" + stateScope) + require.NoError(t, err) + assert.Equal(t, artifactOwnerMarkerPrefix+originA, mode, + "successful push must persist the target-scoped identity mode") + + artifactRepository := openArtifactTestRepository(t, ctx) + exportResult, err := artifact.ExportToStore(ctx, originDB, artifactRepository.Content(), artifact.ExportOptions{ + Origin: originA, + Full: true, + }) + require.NoError(t, err) + importerDB, err := db.Open(filepath.Join(t.TempDir(), "importer.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, importerDB.Close()) }) + require.NoError(t, artifact.AdoptOrigin(importerDB, originB)) + artifactImportResult := importExportedArtifactCheckpoint( + t, ctx, importerDB, artifactRepository, originA, originB, exportResult, + ) + require.Equal(t, 2, artifactImportResult.Sessions) + require.Equal(t, 2, artifactImportResult.Messages) + + importerSync := &Sync{ + pg: pg, local: importerDB, machine: "host-b", + schema: schema, schemaDone: true, + syncState: newScopedSyncStateStore(importerDB, stateScope, false), + syncStateTarget: stateScope, + } + importResult, err := importerSync.Push(ctx, false, nil) + require.NoError(t, err) + assert.Zero(t, importResult.Errors) + assert.Zero(t, importResult.SkippedConflicts) + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM sessions + WHERE id IN ($1, $2, $3, $4) + `, parentID, childID, canonicalParentID, canonicalChildID).Scan(&stableBareRows)) + assert.Equal(t, 2, stableBareRows, + "an imported copy must reuse the stable legacy bare aliases") + assert.Equal(t, parentID, pgParentSessionID(t, ctx, pg, childID)) + assertPGArtifactUpgradeCuration(t, ctx, pg, parentID) + + _, err = pg.ExecContext(ctx, ` + UPDATE sessions SET parent_session_id = $1 WHERE id = $2 + `, canonicalParentID, childID) + require.NoError(t, err, "seed stale canonical relationship") + time.Sleep(5 * time.Millisecond) + name := "Imported child" + require.NoError(t, importerDB.RenameSession(canonicalChildID, &name)) + replicaResult, err := importerSync.Push(ctx, false, nil) + require.NoError(t, err) + assert.Zero(t, replicaResult.SessionsPushed, + "an imported replica must not update an existing canonical alias") + assert.Zero(t, replicaResult.SkippedConflicts) + assert.Equal(t, canonicalParentID, pgParentSessionID(t, ctx, pg, childID), + "a stale replica must leave the canonical row unchanged") + + originName := "Origin child" + require.NoError(t, originDB.RenameSession(childID, &originName)) + fallbackResult, err := originSync.Push(ctx, false, nil) + require.NoError(t, err) + assert.Equal(t, 1, fallbackResult.SessionsPushed, + "the modified origin child should use committed-PG relationship fallback") + assert.Zero(t, fallbackResult.SkippedConflicts) + assert.Equal(t, parentID, pgParentSessionID(t, ctx, pg, childID), + "origin relationship fallback must resolve the stable bare parent alias") + assertPGArtifactUpgradeCuration(t, ctx, pg, parentID) +} + +func TestPushArtifactOriginAdoptionConvergesImporterFirst(t *testing.T) { + pgURL := testPGURL(t) + ctx := context.Background() + const schema = "agentsview_artifact_origin_upgrade_importer_first_test" + const originA = "origin-a1b2c3" + const originB = "origin-b4c5d6" + const parentID = "native-parent" + const childID = "native-child" + const canonicalParentID = originA + "~" + parentID + const canonicalChildID = originA + "~" + childID + const referenceID = "pg-reference-holder" + const stateScope = "work" + const sourceDisplayName = "Source title" + const localDisplayName = "PG title" + const pinCreatedAt = "2026-01-04T00:00:00Z" + const updatedSourceDisplay = "Origin renamed child" + const canonicalChildDeletedAt = "2026-01-05T00:00:00Z" + + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + t.Cleanup(func() { require.NoError(t, pg.Close()) }) + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + originDB, err := db.Open(filepath.Join(t.TempDir(), "origin.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, originDB.Close()) }) + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: parentID, Project: "alpha", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, originDB.ReplaceSessionMessages(parentID, []db.Message{{ + SessionID: parentID, Ordinal: 0, Role: "user", + Content: "parent", ContentLength: 6, SourceUUID: "parent-source", + }})) + parentRef := parentID + require.NoError(t, originDB.UpsertSession(db.Session{ + ID: childID, Project: "alpha", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 0, HasToolCalls: true, + CreatedAt: "2026-01-01T00:00:01Z", + ParentSessionID: &parentRef, + SourceSessionID: parentID, + })) + require.NoError(t, originDB.ReplaceSessionMessages(childID, []db.Message{{ + SessionID: childID, Ordinal: 0, Role: "assistant", + Content: "child", ContentLength: 5, SourceUUID: "child-source", + HasToolUse: true, + ToolCalls: []db.ToolCall{{ + ToolName: "subagent", Category: "Task", ToolUseID: "call-parent", + SubagentSessionID: parentID, + ResultEvents: []db.ToolResultEvent{{ + ToolUseID: "call-parent", SubagentSessionID: parentID, + Source: "tool_result", Status: "completed", Content: "done", + ContentLength: 4, + }}, + }}, + }})) + + originSync := &Sync{ + pg: pg, local: originDB, machine: "host-a", + schema: schema, schemaDone: true, + syncState: newScopedSyncStateStore(originDB, stateScope, false), + syncStateTarget: stateScope, + } + legacyResult, err := originSync.Push(ctx, false, nil) + require.NoError(t, err) + require.Equal(t, 2, legacyResult.SessionsPushed) + require.NoError(t, originDB.SetSyncState( + "pg_artifact_identity_v1:"+stateScope, "", + )) + + _, err = pg.ExecContext(ctx, ` + UPDATE sessions + SET display_name = $1, + source_display_name = $2 + WHERE id = $3 + `, localDisplayName, sourceDisplayName, parentID) + require.NoError(t, err, "seed PG-local session curation") + _, err = pg.ExecContext(ctx, ` + INSERT INTO starred_sessions (session_id, created_at) + VALUES ($1, '2026-01-04T00:00:00Z'::timestamptz) + `, parentID) + require.NoError(t, err, "seed PG-local star") + _, err = pg.ExecContext(ctx, ` + INSERT INTO pinned_messages ( + session_id, message_id, ordinal, source_uuid, note, created_at + ) + VALUES ($1, 0, 0, 'parent-source', 'PG pin', $2::timestamptz) + `, parentID, pinCreatedAt) + require.NoError(t, err, "seed PG-local pin") + + require.NoError(t, artifact.AdoptOrigin(originDB, originA)) + artifactRepository := openArtifactTestRepository(t, ctx) + exportResult, err := artifact.ExportToStore(ctx, originDB, artifactRepository.Content(), artifact.ExportOptions{ + Origin: originA, + Full: true, + }) + require.NoError(t, err) + importerDB, err := db.Open(filepath.Join(t.TempDir(), "importer.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, importerDB.Close()) }) + require.NoError(t, artifact.AdoptOrigin(importerDB, originB)) + importResult := importExportedArtifactCheckpoint( + t, ctx, importerDB, artifactRepository, originA, originB, exportResult, + ) + require.Equal(t, 2, importResult.Sessions) + require.Equal(t, 2, importResult.Messages) + + importerSync := &Sync{ + pg: pg, local: importerDB, machine: "host-b", + schema: schema, schemaDone: true, + syncState: newScopedSyncStateStore(importerDB, stateScope, false), + syncStateTarget: stateScope, + } + importerResult, err := importerSync.Push(ctx, false, nil) + require.NoError(t, err) + assert.Zero(t, importerResult.Errors) + assert.Zero(t, importerResult.SkippedConflicts) + assertPGArtifactNativeRows(t, ctx, pg, parentID, canonicalParentID, 2) + assertPGArtifactNativeRows(t, ctx, pg, childID, canonicalChildID, 2) + _, err = pg.ExecContext(ctx, ` + UPDATE sessions + SET deleted_at = $1::timestamptz, + source_deleted_at = NULL + WHERE id = $2 + `, canonicalChildDeletedAt, canonicalChildID) + require.NoError(t, err, "seed newer canonical PG curation") + require.NoError(t, originDB.RenameSession( + childID, new(updatedSourceDisplay), + )) + require.NoError(t, originDB.SoftDeleteSession(parentID)) + updatedParent, err := originDB.GetSessionFull(ctx, parentID) + require.NoError(t, err) + require.NotNil(t, updatedParent) + require.NotNil(t, updatedParent.DeletedAt) + updatedSourceDeletedAt, ok := ParseSQLiteTimestamp(*updatedParent.DeletedAt) + require.True(t, ok, "parse source soft-delete timestamp") + wantSourceDeletedAt := updatedSourceDeletedAt.UTC().Format(time.RFC3339) + + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions ( + id, machine, owner_marker, project, agent, created_at, + parent_session_id, source_session_id + ) VALUES ($1, 'pg-only', 'pg-only-owner', 'alpha', 'claude', NOW(), $2, $2) + `, referenceID, parentID) + require.NoError(t, err, "seed PG-only session relationships") + _, err = pg.ExecContext(ctx, ` + INSERT INTO messages (session_id, ordinal, role, content) + VALUES ($1, 0, 'assistant', 'reference') + `, referenceID) + require.NoError(t, err, "seed PG-only message") + _, err = pg.ExecContext(ctx, ` + INSERT INTO tool_calls ( + session_id, tool_name, category, call_index, tool_use_id, + subagent_session_id, message_ordinal + ) VALUES ($1, 'subagent', 'Task', 0, 'pg-call', $2, 0) + `, referenceID, parentID) + require.NoError(t, err, "seed PG-only tool call relationship") + _, err = pg.ExecContext(ctx, ` + INSERT INTO tool_result_events ( + session_id, tool_call_message_ordinal, call_index, + tool_use_id, subagent_session_id, source, status, + content, content_length, event_index + ) VALUES ($1, 0, 0, 'pg-call', $2, 'tool_result', + 'completed', 'done', 4, 0) + `, referenceID, parentID) + require.NoError(t, err, "seed PG-only tool result relationship") + + upgradeResult, err := originSync.Push(ctx, false, nil) + require.NoError(t, err) + assert.Equal(t, 2, upgradeResult.SessionsPushed, + "identity-mode change must revisit unchanged origin sessions") + assert.Zero(t, upgradeResult.Errors) + assert.Zero(t, upgradeResult.SkippedConflicts) + assertPGArtifactNativeRows(t, ctx, pg, parentID, canonicalParentID, 1) + assertPGArtifactNativeRows(t, ctx, pg, childID, canonicalChildID, 1) + assertPGArtifactStableOwner(t, ctx, pg, canonicalParentID, originA) + assertPGArtifactStableOwner(t, ctx, pg, canonicalChildID, originA) + assertPGArtifactUpgradeCurationDetails( + t, ctx, pg, canonicalParentID, + localDisplayName, sourceDisplayName, + wantSourceDeletedAt, wantSourceDeletedAt, pinCreatedAt, + ) + assertPGArtifactUpdatedSourceDisplay( + t, ctx, pg, canonicalChildID, + updatedSourceDisplay, canonicalChildDeletedAt, + ) + assertPGArtifactUpgradeRelationships( + t, ctx, pg, canonicalParentID, canonicalChildID, referenceID, + ) + + for _, pusher := range []*Sync{importerSync, originSync} { + repeated, pushErr := pusher.Push(ctx, true, nil) + require.NoError(t, pushErr) + assert.Zero(t, repeated.Errors) + assert.Zero(t, repeated.SkippedConflicts) + } + assertPGArtifactNativeRows(t, ctx, pg, parentID, canonicalParentID, 1) + assertPGArtifactNativeRows(t, ctx, pg, childID, canonicalChildID, 1) + assertPGArtifactStableOwner(t, ctx, pg, canonicalParentID, originA) + assertPGArtifactStableOwner(t, ctx, pg, canonicalChildID, originA) + assertPGArtifactUpgradeCurationAfterRepeatedPushes( + t, ctx, pg, canonicalParentID, + localDisplayName, wantSourceDeletedAt, pinCreatedAt, + ) + assertPGArtifactUpgradeRelationships( + t, ctx, pg, canonicalParentID, canonicalChildID, referenceID, + ) +} + +func TestPushArtifactOriginAdoptionIgnoresForeignBareAlias(t *testing.T) { + pgURL := testPGURL(t) + ctx := context.Background() + const schema = "agentsview_artifact_origin_upgrade_proof_test" + const origin = "origin-a1b2c3" + const nativeID = "native-id" + const canonicalID = origin + "~" + nativeID + + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + t.Cleanup(func() { require.NoError(t, pg.Close()) }) + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, localDB.Close()) }) + require.NoError(t, artifact.AdoptOrigin(localDB, origin)) + require.NoError(t, localDB.UpsertSession(db.Session{ + ID: nativeID, Project: "local-project", Machine: "local", Agent: "claude", + MessageCount: 1, UserMessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + })) + require.NoError(t, localDB.ReplaceSessionMessages(nativeID, []db.Message{{ + SessionID: nativeID, Ordinal: 0, Role: "user", + Content: "local", ContentLength: 5, + }})) + + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions ( + id, machine, owner_marker, project, agent, created_at + ) VALUES + ($1, $2, $3, 'canonical-before', 'claude', NOW()), + ($4, 'legacy-host', 'different-random-marker', + 'legacy-before', 'claude', NOW()) + `, canonicalID, origin, artifactOwnerMarkerPrefix+origin, nativeID) + require.NoError(t, err, "seed unproven duplicate pair") + + syncer := &Sync{ + pg: pg, local: localDB, machine: "host-a", + schema: schema, schemaDone: true, + } + result, err := syncer.Push(ctx, true, nil) + require.NoError(t, err) + assert.Zero(t, result.Errors) + assert.Zero(t, result.SkippedConflicts) + assert.Equal(t, 1, result.SessionsPushed) + + rows, err := pg.QueryContext(ctx, ` + SELECT id, project FROM sessions + WHERE id IN ($1, $2) ORDER BY id + `, nativeID, canonicalID) + require.NoError(t, err) + defer rows.Close() + projects := map[string]string{} + for rows.Next() { + var id, project string + require.NoError(t, rows.Scan(&id, &project)) + projects[id] = project + } + require.NoError(t, rows.Err()) + assert.Equal(t, map[string]string{ + nativeID: "legacy-before", + canonicalID: "local-project", + }, projects, "a foreign bare collision must not block the stable canonical row") +} + +func assertPGArtifactNativeRows( + t *testing.T, ctx context.Context, pg *sql.DB, + bareID, canonicalID string, want int, +) { + t.Helper() + var count int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM sessions WHERE id IN ($1, $2) + `, bareID, canonicalID).Scan(&count)) + assert.Equal(t, want, count) +} + +func assertPGArtifactStableOwner( + t *testing.T, ctx context.Context, pg *sql.DB, + sessionID, origin string, +) { + t.Helper() + var machine, ownerMarker string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT machine, owner_marker FROM sessions WHERE id = $1 + `, sessionID).Scan(&machine, &ownerMarker)) + assert.Equal(t, origin, machine) + assert.Equal(t, artifactOwnerMarkerPrefix+origin, ownerMarker) +} + +func assertPGArtifactUpgradeCurationDetails( + t *testing.T, ctx context.Context, pg *sql.DB, sessionID, + wantDisplay, wantSourceDisplay, wantDeleted, wantSourceDeleted, + wantPinCreated string, +) { + t.Helper() + assertPGArtifactSessionCuration( + t, ctx, pg, sessionID, + wantDisplay, wantSourceDisplay, wantDeleted, wantSourceDeleted, + ) + + var stars int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM starred_sessions WHERE session_id = $1 + `, sessionID).Scan(&stars)) + assert.Equal(t, 1, stars) + + var ordinal, messageID int + var sourceUUID, note string + var createdAt time.Time + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT message_id, ordinal, source_uuid, note, created_at + FROM pinned_messages WHERE session_id = $1 + `, sessionID).Scan( + &messageID, &ordinal, &sourceUUID, ¬e, &createdAt, + )) + assert.Equal(t, 0, messageID) + assert.Equal(t, 0, ordinal) + assert.Equal(t, "parent-source", sourceUUID) + assert.Equal(t, "PG pin", note) + assert.Equal(t, wantPinCreated, createdAt.UTC().Format(time.RFC3339)) +} + +func assertPGArtifactSessionCuration( + t *testing.T, ctx context.Context, pg *sql.DB, sessionID, + wantDisplay, wantSourceDisplay, wantDeleted, wantSourceDeleted string, +) { + t.Helper() + var displayName, sourceDisplayName string + var deletedAt, sourceDeletedAt sql.NullTime + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT display_name, source_display_name, deleted_at, source_deleted_at + FROM sessions WHERE id = $1 + `, sessionID).Scan( + &displayName, &sourceDisplayName, &deletedAt, &sourceDeletedAt, + )) + assert.Equal(t, wantDisplay, displayName) + assert.Equal(t, wantSourceDisplay, sourceDisplayName) + if assert.True(t, deletedAt.Valid, "deleted_at") { + assert.Equal(t, wantDeleted, deletedAt.Time.UTC().Format(time.RFC3339)) + } + if assert.True(t, sourceDeletedAt.Valid, "source_deleted_at") { + assert.Equal(t, wantSourceDeleted, + sourceDeletedAt.Time.UTC().Format(time.RFC3339)) + } +} + +func assertPGArtifactUpdatedSourceDisplay( + t *testing.T, ctx context.Context, pg *sql.DB, sessionID, + wantDisplay, wantDeleted string, +) { + t.Helper() + var displayName, sourceDisplayName sql.NullString + var deletedAt time.Time + var sourceDeletedAt sql.NullTime + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT display_name, source_display_name, deleted_at, source_deleted_at + FROM sessions WHERE id = $1 + `, sessionID).Scan( + &displayName, &sourceDisplayName, &deletedAt, &sourceDeletedAt, + )) + if assert.True(t, displayName.Valid, "display_name") { + assert.Equal(t, wantDisplay, displayName.String) + } + if assert.True(t, sourceDisplayName.Valid, "source_display_name") { + assert.Equal(t, wantDisplay, sourceDisplayName.String) + } + assert.Equal(t, wantDeleted, deletedAt.UTC().Format(time.RFC3339)) + assert.False(t, sourceDeletedAt.Valid, + "canonical delete override must retain the current source baseline") +} + +func assertPGArtifactUpgradeCurationAfterRepeatedPushes( + t *testing.T, ctx context.Context, pg *sql.DB, sessionID, + wantDisplay, wantDeleted, wantPinCreated string, +) { + t.Helper() + var displayName string + var deletedAt time.Time + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT display_name, deleted_at + FROM sessions WHERE id = $1 + `, sessionID).Scan(&displayName, &deletedAt)) + assert.Equal(t, wantDisplay, displayName) + assert.Equal(t, wantDeleted, deletedAt.UTC().Format(time.RFC3339)) + + var stars int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM starred_sessions WHERE session_id = $1 + `, sessionID).Scan(&stars)) + assert.Equal(t, 1, stars) + + var ordinal, messageID int + var sourceUUID, note string + var createdAt time.Time + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT message_id, ordinal, source_uuid, note, created_at + FROM pinned_messages WHERE session_id = $1 + `, sessionID).Scan( + &messageID, &ordinal, &sourceUUID, ¬e, &createdAt, + )) + assert.Equal(t, 0, messageID) + assert.Equal(t, 0, ordinal) + assert.Equal(t, "parent-source", sourceUUID) + assert.Equal(t, "PG pin", note) + assert.Equal(t, wantPinCreated, createdAt.UTC().Format(time.RFC3339)) +} + +func assertPGArtifactUpgradeRelationships( + t *testing.T, ctx context.Context, pg *sql.DB, + canonicalParentID, canonicalChildID, referenceID string, +) { + t.Helper() + var parentID, sourceID string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT parent_session_id, source_session_id + FROM sessions WHERE id = $1 + `, canonicalChildID).Scan(&parentID, &sourceID)) + assert.Equal(t, canonicalParentID, parentID) + assert.Equal(t, canonicalParentID, sourceID) + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT parent_session_id, source_session_id + FROM sessions WHERE id = $1 + `, referenceID).Scan(&parentID, &sourceID)) + assert.Equal(t, canonicalParentID, parentID) + assert.Equal(t, canonicalParentID, sourceID) + + var toolCallSubagentID string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT subagent_session_id FROM tool_calls WHERE session_id = $1 + `, referenceID).Scan(&toolCallSubagentID)) + assert.Equal(t, canonicalParentID, toolCallSubagentID, "tool_calls") + var toolResultSubagentID string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT subagent_session_id + FROM tool_result_events WHERE session_id = $1 + `, referenceID).Scan(&toolResultSubagentID)) + assert.Equal(t, canonicalParentID, toolResultSubagentID, "tool_result_events") +} + +func pgParentSessionID( + t *testing.T, ctx context.Context, pg *sql.DB, sessionID string, +) string { + t.Helper() + var parentID string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COALESCE(parent_session_id, '') FROM sessions WHERE id = $1 + `, sessionID).Scan(&parentID)) + return parentID +} + +func assertPGArtifactUpgradeCuration( + t *testing.T, ctx context.Context, pg *sql.DB, sessionID string, +) { + t.Helper() + var displayName string + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COALESCE(display_name, '') FROM sessions WHERE id = $1 + `, sessionID).Scan(&displayName)) + assert.Equal(t, "PG title", displayName) + var stars, pins int + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM starred_sessions WHERE session_id = $1 + `, sessionID).Scan(&stars)) + require.NoError(t, pg.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM pinned_messages + WHERE session_id = $1 AND note = 'PG pin' + `, sessionID).Scan(&pins)) + assert.Equal(t, 1, stars) + assert.Equal(t, 1, pins) +} + // TestPushSessionGuardsAgainstCrossMachineCollision verifies that when two // machines share the same session ID (from dotfile sync, directory restore, etc.), // the second machine's push is skipped if the session is already owned by a @@ -84,7 +983,10 @@ func TestPushSessionGuardsAgainstCrossMachineCollision(t *testing.T) { // Execute pushSession. tx, err := pg.BeginTx(ctx, nil) require.NoError(t, err, "BeginTx") - err = sync.pushSession(ctx, tx, sess, markerID, nil) + err = sync.pushSession(ctx, tx, sess, pushedSessionIdentity{ + ID: sess.ID, + Machine: sess.Machine, + }, markerID, nil) require.ErrorIs(t, err, errSessionOwnershipConflict, "pushSession should return ownership conflict sentinel") require.NoError(t, tx.Commit(), "Commit") @@ -156,7 +1058,10 @@ func TestPushSessionAllowsMachineRenameForSameOwnerMarker(t *testing.T) { tx, err := pg.BeginTx(ctx, nil) require.NoError(t, err, "BeginTx") - require.NoError(t, sync.pushSession(ctx, tx, sess, markerID, nil), "pushSession") + require.NoError(t, sync.pushSession(ctx, tx, sess, pushedSessionIdentity{ + ID: sess.ID, + Machine: "renamed-host", + }, markerID, nil), "pushSession") require.NoError(t, tx.Commit(), "Commit") var machine, ownerMarker string @@ -168,6 +1073,214 @@ func TestPushSessionAllowsMachineRenameForSameOwnerMarker(t *testing.T) { assert.Equal(t, markerID, ownerMarker) } +// TestPushResolvesRelationshipIDsToPrefixedTargets verifies that when a +// referenced session is pushed under a collision-avoidance prefix, the +// relationship ids pointing at it -- source_session_id, parent_session_id, and +// a tool-call subagent_session_id -- are rewritten to the prefixed id so child +// and subagent rows link to the right PG session instead of a foreign machine's +// row or a dangling id. A non-colliding parent keeps its bare id. +func TestPushResolvesRelationshipIDsToPrefixedTargets(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_relationship_resolution_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "machine-b", + schema: schema, + schemaDone: true, + } + + // A foreign machine already owns the bare "shared" id, so machine-b's + // "shared" session must be pushed under the "machine-b~shared" prefix. + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions ( + id, machine, owner_marker, project, agent, created_at + ) VALUES ($1, $2, $3, $4, $5, NOW()) + `, "shared", "machine-a", "foreign-owner", "test-proj", "claude") + require.NoError(t, err, "insert foreign-owned shared session") + + sharedParent := "shared" + plainParent := "plain" + sessions := []db.Session{ + {ID: "shared", Project: "test-proj", Machine: "machine-b", + Agent: "claude", MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + {ID: "plain", Project: "test-proj", Machine: "machine-b", + Agent: "claude", MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + {ID: "child-shared", Project: "test-proj", Machine: "machine-b", + Agent: "claude", MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + SourceSessionID: "shared", ParentSessionID: &sharedParent}, + {ID: "child-plain", Project: "test-proj", Machine: "machine-b", + Agent: "claude", MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + ParentSessionID: &plainParent}, + } + for _, s := range sessions { + require.NoError(t, localDB.UpsertSession(s), "UpsertSession "+s.ID) + } + + // child-shared references the colliding "shared" id from a tool call too. + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: "child-shared", Ordinal: 0, Role: "assistant", + Content: "spawning", HasToolUse: true, + ToolCalls: []db.ToolCall{{ + ToolName: "subagent", Category: "Task", + SubagentSessionID: "shared", + }}, + }}), "InsertMessages child-shared") + for _, id := range []string{"shared", "plain", "child-plain"} { + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: id, Ordinal: 0, Role: "user", + Content: "hi", ContentLength: 2, + }}), "InsertMessages "+id) + } + + _, err = sync.Push(ctx, false, nil) + require.NoError(t, err, "Push") + + const prefixedShared = "machine-b~shared" + var n int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE id = $1 AND machine = $2`, + prefixedShared, "machine-b").Scan(&n), "count prefixed shared") + assert.Equal(t, 1, n, "machine-b's shared session stored under prefixed id") + + var source, parent string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT source_session_id, parent_session_id + FROM sessions WHERE id = $1`, + "child-shared").Scan(&source, &parent), "read child-shared relations") + assert.Equal(t, prefixedShared, source, "source_session_id resolved") + assert.Equal(t, prefixedShared, parent, "parent_session_id resolved") + + var subagent string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT subagent_session_id FROM tool_calls WHERE session_id = $1`, + "child-shared").Scan(&subagent), "read child-shared subagent link") + assert.Equal(t, prefixedShared, subagent, "subagent_session_id resolved") + + var plainParentGot string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT parent_session_id FROM sessions WHERE id = $1`, + "child-plain").Scan(&plainParentGot), "read child-plain parent") + assert.Equal(t, "plain", plainParentGot, "non-colliding parent stays bare") +} + +// TestPushRepairsStaleSubagentLinkOnIncrementalPush verifies that an +// incremental push repairs a PG tool-call subagent link left at its unprefixed +// local id by a push that predated the collision: the parent's tool call was +// pushed while "sub-1" was still unclaimed, and the subagent session only +// later collided with a foreign owner and moved under "machine-b~sub-1". The +// local rows never change, so both the session candidacy fingerprint and the +// message fast path would otherwise skip the parent and never rewrite the +// link; only the resolved subagent id distinguishes it from PG. +func TestPushRepairsStaleSubagentLinkOnIncrementalPush(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_stale_subagent_repair_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "machine-b", + schema: schema, + schemaDone: true, + } + + // The parent references "sub-1" before any session by that id exists in + // PG or locally, so the first push writes the tool-call link at its bare + // local id. + require.NoError(t, localDB.UpsertSession(db.Session{ + ID: "parent-1", Project: "proj", Machine: "machine-b", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + }), "UpsertSession parent-1") + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: "parent-1", Ordinal: 0, Role: "assistant", + Content: "spawning", HasToolUse: true, + ToolCalls: []db.ToolCall{{ + ToolName: "subagent", Category: "Task", SubagentSessionID: "sub-1", + }}, + }}), "InsertMessages parent-1") + + _, err = sync.Push(ctx, false, nil) + require.NoError(t, err, "first Push") + + subagent := func() string { + var s string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT subagent_session_id FROM tool_calls WHERE session_id = $1`, + "parent-1").Scan(&s), "read parent-1 subagent link") + return s + } + require.Equal(t, "sub-1", subagent(), + "first push keeps the bare link while the id is unclaimed") + + // A foreign machine claims the bare "sub-1" id, then machine-b's subagent + // session appears locally and is pushed under the "machine-b~sub-1" + // prefix. The parent's local row is unchanged, so its PG link now points + // at the foreign machine's session. + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions ( + id, machine, owner_marker, project, agent, created_at + ) VALUES ($1, $2, $3, $4, $5, NOW()) + `, "sub-1", "machine-a", "foreign-owner", "proj", "claude") + require.NoError(t, err, "insert foreign-owned subagent session") + + require.NoError(t, localDB.UpsertSession(db.Session{ + ID: "sub-1", Project: "proj", Machine: "machine-b", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + }), "UpsertSession sub-1") + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: "sub-1", Ordinal: 0, Role: "user", Content: "hi", ContentLength: 2, + }}), "InsertMessages sub-1") + + const prefixedSub = "machine-b~sub-1" + _, err = sync.Push(ctx, false, nil) + require.NoError(t, err, "second Push") + var n int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE id = $1 AND machine = $2`, + prefixedSub, "machine-b").Scan(&n), "count prefixed sub-1") + require.Equal(t, 1, n, "machine-b's subagent stored under prefixed id") + require.Equal(t, "sub-1", subagent(), + "precondition: parent was not a candidate, so its link is stale") + + // Re-list the parent without touching its message content, so only the + // resolved subagent id distinguishes it from what was last pushed. + require.NoError(t, localDB.BumpLocalModifiedAt("parent-1"), + "mark parent-1 modified") + + res, err := sync.Push(ctx, false, nil) + require.NoError(t, err, "third Push") + assert.Zero(t, res.Errors, "third push should report no failures") + assert.Equal(t, prefixedSub, subagent(), + "incremental push repairs the stale subagent link") +} + func TestPushSessionAdoptsLegacyLocalSentinelRow(t *testing.T) { pgURL := testPGURL(t) @@ -215,7 +1328,10 @@ func TestPushSessionAdoptsLegacyLocalSentinelRow(t *testing.T) { require.NoError(t, err, "BeginTx") markerID, err := sync.pushMarkerID() require.NoError(t, err, "pushMarkerID") - require.NoError(t, sync.pushSession(ctx, tx, sess, markerID, nil), "pushSession") + require.NoError(t, sync.pushSession(ctx, tx, sess, pushedSessionIdentity{ + ID: sess.ID, + Machine: "host-a", + }, markerID, nil), "pushSession") require.NoError(t, tx.Commit(), "Commit") var machine, ownerMarker string @@ -226,3 +1342,352 @@ func TestPushSessionAdoptsLegacyLocalSentinelRow(t *testing.T) { assert.Equal(t, "host-a", machine) assert.Equal(t, markerID, ownerMarker) } + +// TestResolveOwnedPushIDReusesLegacyPrefixAfterRename verifies that the shared +// id resolver (used by both resolvePushedSessionIdentity and +// relationshipResolver.lookup) reuses a row owned under a prior machine prefix. +// A pusher that once stored a colliding session under "old-host~id" and later +// renamed to "new-host" must resolve back to "old-host~id" instead of minting a +// duplicate "new-host~id". +func TestResolveOwnedPushIDReusesLegacyPrefixAfterRename(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_legacy_prefix_resolve_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "new-host", + schema: schema, + schemaDone: true, + } + markerID, err := sync.pushMarkerID() + require.NoError(t, err, "pushMarkerID") + + // A foreign machine owns the bare id, which is why this pusher stored its + // session under the old machine prefix before the rename. + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + "sess-x", "foreign", "foreign-owner", "proj", "claude") + require.NoError(t, err, "insert foreign bare row") + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + "old-host~sess-x", "old-host", markerID, "proj", "claude") + require.NoError(t, err, "insert owned legacy-prefixed row") + + identity := pushedSessionIdentity{Machine: "new-host"} + got, err := sync.resolveOwnedPushIdentityID( + ctx, "sess-x", identity, markerID, []string{"old-host"}, + ) + require.NoError(t, err) + assert.Equal(t, "old-host~sess-x", got, + "a renamed pusher must reuse the row it owns under the old machine prefix") + + // Without the old machine name, the resolver cannot find the owned row and + // would mint a duplicate under the new prefix -- the blind spot being fixed. + got, err = sync.resolveOwnedPushIdentityID(ctx, "sess-x", identity, markerID, nil) + require.NoError(t, err) + assert.Equal(t, "new-host~sess-x", got, + "precondition: without the legacy machine the resolver duplicates the row") +} + +// TestPushReusesLegacyPrefixedRowAfterRename verifies the end-to-end push: +// after a machine rename, a session previously stored under the old machine +// prefix is updated in place rather than duplicated under the new prefix. +func TestPushReusesLegacyPrefixedRowAfterRename(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_legacy_prefix_push_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "new-host", + schema: schema, + schemaDone: true, + } + markerID, err := sync.pushMarkerID() + require.NoError(t, err, "pushMarkerID") + + // Record that this marker last pushed as "old-host", so the rename push + // treats "old-host" as a legacy machine prefix. + _, err = pg.ExecContext(ctx, ` + INSERT INTO sync_metadata (key, value) VALUES ($1, $2) + ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value`, + pushMarkerKeyPrefix+markerID, "old-host") + require.NoError(t, err, "seed push marker machine") + + // A foreign machine owns the bare id; this pusher's session lives under the + // old machine prefix from before the rename. + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + "sess-x", "foreign", "foreign-owner", "proj", "claude") + require.NoError(t, err, "insert foreign bare row") + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + "old-host~sess-x", "old-host", markerID, "proj", "claude") + require.NoError(t, err, "insert owned legacy-prefixed row") + + sess := db.Session{ + ID: "sess-x", Project: "proj", Machine: "local", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", + } + require.NoError(t, localDB.UpsertSession(sess), "UpsertSession") + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: "sess-x", Ordinal: 0, Role: "user", Content: "hi", ContentLength: 2, + }}), "InsertMessages") + + _, err = sync.Push(ctx, false, nil) + require.NoError(t, err, "Push") + + var dup int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE id = $1`, "new-host~sess-x").Scan(&dup), + "count new-prefix duplicate") + assert.Equal(t, 0, dup, "rename must not create a duplicate under the new prefix") + + // The legacy-prefixed row was updated in place: the machine column reflects + // the rename and the pushed message landed there. + var machine string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT machine FROM sessions WHERE id = $1`, "old-host~sess-x").Scan(&machine), + "read reused row machine") + assert.Equal(t, "new-host", machine, "owned legacy row updated in place") + var msgs int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE session_id = $1`, "old-host~sess-x").Scan(&msgs), + "count reused row messages") + assert.Equal(t, 1, msgs, "message pushed to the reused legacy row") + + var foreignOwner string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT owner_marker FROM sessions WHERE id = $1`, "sess-x").Scan(&foreignOwner), + "read foreign bare row") + assert.Equal(t, "foreign-owner", foreignOwner, "foreign bare row not adopted") +} + +// TestPushSkipsPrefixedConflictWithoutAbortingBatch verifies that a single +// per-session ownership conflict on the current-machine prefixed id is skipped +// and reported, while unrelated sessions in the same push still go through. The +// conflict must not fail the whole push from identity pre-resolution. +func TestPushSkipsPrefixedConflictWithoutAbortingBatch(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_prefixed_conflict_skip_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "new-host", + schema: schema, + schemaDone: true, + } + markerID, err := sync.pushMarkerID() + require.NoError(t, err, "pushMarkerID") + + // A different owner holds both the bare id and the current-machine prefixed + // id, so "conf-1" has nowhere to land and must be skipped as a conflict. + for _, id := range []string{"conf-1", "new-host~conf-1"} { + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + id, "other-host", "other-owner", "proj", "claude") + require.NoError(t, err, "insert foreign "+id) + } + + for _, s := range []db.Session{ + {ID: "conf-1", Project: "proj", Machine: "new-host", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + {ID: "clean-1", Project: "proj", Machine: "new-host", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + } { + require.NoError(t, localDB.UpsertSession(s), "UpsertSession "+s.ID) + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: s.ID, Ordinal: 0, Role: "user", Content: "hi", ContentLength: 2, + }}), "InsertMessages "+s.ID) + } + + res, err := sync.Push(ctx, false, nil) + require.NoError(t, err, "a per-session conflict must not fail the whole push") + assert.Equal(t, 1, res.SkippedConflicts, "the conflicting session is reported as skipped") + + var clean int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE id = $1 AND owner_marker = $2`, + "clean-1", markerID).Scan(&clean), "count clean session") + assert.Equal(t, 1, clean, "unrelated session pushed despite the conflict") + + var owner string + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT owner_marker FROM sessions WHERE id = $1`, "new-host~conf-1").Scan(&owner), + "read foreign prefixed row") + assert.Equal(t, "other-owner", owner, "foreign prefixed row not overwritten") +} + +func TestPushSkipsRelationshipsToPrefixedOwnershipConflict(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_prefixed_conflict_relationship_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "new-host", + schema: schema, + schemaDone: true, + } + + for _, id := range []string{"conf-1", "new-host~conf-1"} { + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + id, "other-host", "other-owner", "proj", "claude") + require.NoError(t, err, "insert foreign "+id) + } + + sourceID := "conf-1" + for _, s := range []db.Session{ + {ID: "conf-1", Project: "proj", Machine: "new-host", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + {ID: "child-1", Project: "proj", Machine: "new-host", Agent: "claude", + SourceSessionID: sourceID, MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + } { + require.NoError(t, localDB.UpsertSession(s), "UpsertSession "+s.ID) + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: s.ID, Ordinal: 0, Role: "user", Content: "hi", ContentLength: 2, + }}), "InsertMessages "+s.ID) + } + + res, err := sync.Push(ctx, false, nil) + require.NoError(t, err, "relationship to a per-session conflict must not fail the whole push") + assert.Equal(t, 2, res.SkippedConflicts, + "the conflicted session and the dependent relationship session are skipped") + assert.Zero(t, res.Errors) + + var childRows int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE id = $1`, "child-1").Scan(&childRows), + "count child session") + assert.Equal(t, 0, childRows, "dependent session must not be pushed with a foreign source link") + + var foreignRefs int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE source_session_id = $1`, + "new-host~conf-1").Scan(&foreignRefs), "count references to foreign prefixed row") + assert.Equal(t, 0, foreignRefs, "no pushed row may point at the foreign prefixed session") +} + +func TestPushSkipsRelationshipsToAlreadyPrefixedOwnershipConflict(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_already_prefixed_conflict_relationship_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB, err := db.Open(filepath.Join(t.TempDir(), "local.db")) + require.NoError(t, err, "db.Open") + defer localDB.Close() + + sync := &Sync{ + pg: pg, + local: localDB, + machine: "new-host", + schema: schema, + schemaDone: true, + } + + const conflictedID = "new-host~conf-1" + _, err = pg.ExecContext(ctx, ` + INSERT INTO sessions (id, machine, owner_marker, project, agent, created_at) + VALUES ($1, $2, $3, $4, $5, NOW())`, + conflictedID, "other-host", "other-owner", "proj", "claude") + require.NoError(t, err, "insert foreign already-prefixed row") + + for _, s := range []db.Session{ + {ID: conflictedID, Project: "proj", Machine: "new-host", Agent: "claude", + MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + {ID: "child-1", Project: "proj", Machine: "new-host", Agent: "claude", + SourceSessionID: conflictedID, MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z"}, + } { + require.NoError(t, localDB.UpsertSession(s), "UpsertSession "+s.ID) + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: s.ID, Ordinal: 0, Role: "user", Content: "hi", ContentLength: 2, + }}), "InsertMessages "+s.ID) + } + + res, err := sync.Push(ctx, false, nil) + require.NoError(t, err, "already-prefixed relationship conflict must not fail the whole push") + assert.Equal(t, 2, res.SkippedConflicts, + "the conflicted already-prefixed session and its dependent are skipped") + assert.Zero(t, res.Errors) + + var childRows int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE id = $1`, "child-1").Scan(&childRows), + "count child session") + assert.Equal(t, 0, childRows, "dependent session must not be pushed with a foreign source link") + + var foreignRefs int + require.NoError(t, pg.QueryRowContext(ctx, + `SELECT COUNT(*) FROM sessions WHERE source_session_id = $1`, + conflictedID).Scan(&foreignRefs), "count references to foreign already-prefixed row") + assert.Equal(t, 0, foreignRefs, "no pushed row may point at the foreign already-prefixed session") +} diff --git a/internal/postgres/curation.go b/internal/postgres/curation.go index 9d1b96ebf..1bb59fddd 100644 --- a/internal/postgres/curation.go +++ b/internal/postgres/curation.go @@ -3,6 +3,7 @@ package postgres import ( "context" "database/sql" + "errors" "fmt" "time" @@ -43,15 +44,20 @@ func (s *Store) StarSession(sessionID string) (bool, error) { } // UnstarSession removes a session star from the shared PG dashboard -// metadata. -func (s *Store) UnstarSession(sessionID string) error { - if _, err := s.pg.Exec( +// metadata and reports whether a row was removed. +func (s *Store) UnstarSession(sessionID string) (bool, error) { + res, err := s.pg.Exec( `DELETE FROM starred_sessions WHERE session_id = $1`, sessionID, - ); err != nil { - return fmt.Errorf("unstarring session %s: %w", sessionID, err) + ) + if err != nil { + return false, fmt.Errorf("unstarring session %s: %w", sessionID, err) } - return nil + n, err := res.RowsAffected() + if err != nil { + return false, fmt.Errorf("checking unstar result for %s: %w", sessionID, err) + } + return n > 0, nil } // ListStarredSessionIDs returns shared PG-starred session IDs. @@ -80,35 +86,60 @@ func (s *Store) ListStarredSessionIDs( // BulkStarSessions stars multiple existing sessions in one transaction. // Unknown session IDs are skipped. -func (s *Store) BulkStarSessions(sessionIDs []string) error { +func (s *Store) BulkStarSessions(sessionIDs []string) ([]string, error) { if len(sessionIDs) == 0 { - return nil + return nil, nil } tx, err := s.pg.Begin() if err != nil { - return fmt.Errorf("beginning star transaction: %w", err) + return nil, fmt.Errorf("beginning star transaction: %w", err) } defer func() { _ = tx.Rollback() }() - stmt, err := tx.Prepare(` + // Check existence separately from the insert so stale IDs are skipped and + // the caller learns which sessions were actually starred, mirroring the + // SQLite backend. + exists, err := tx.Prepare(`SELECT 1 FROM sessions WHERE id = $1`) + if err != nil { + return nil, fmt.Errorf("preparing existence statement: %w", err) + } + defer exists.Close() + insert, err := tx.Prepare(` INSERT INTO starred_sessions (session_id) - SELECT $1 WHERE EXISTS ( - SELECT 1 FROM sessions WHERE id = $1 - ) + VALUES ($1) ON CONFLICT (session_id) DO NOTHING`) if err != nil { - return fmt.Errorf("preparing star statement: %w", err) + return nil, fmt.Errorf("preparing star statement: %w", err) } - defer stmt.Close() + defer insert.Close() + starred := make([]string, 0, len(sessionIDs)) for _, id := range sessionIDs { - if _, err := stmt.Exec(id); err != nil { - return fmt.Errorf("starring session %s: %w", id, err) + var one int + switch err := exists.QueryRow(id).Scan(&one); { + case errors.Is(err, sql.ErrNoRows): + continue + case err != nil: + return nil, fmt.Errorf("checking session %s: %w", id, err) + } + res, err := insert.Exec(id) + if err != nil { + return nil, fmt.Errorf("starring session %s: %w", id, err) + } + rowsAffected, err := res.RowsAffected() + if err != nil { + return nil, fmt.Errorf("checking star insert result for %s: %w", id, err) + } + if rowsAffected > 0 { + starred = append(starred, id) } } - return tx.Commit() + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("committing star transaction: %w", err) + } + return starred, nil } // PinMessage creates or updates a shared PG pin for a message. PG diff --git a/internal/postgres/curation_pgtest_test.go b/internal/postgres/curation_pgtest_test.go index 784924c3a..c424903d5 100644 --- a/internal/postgres/curation_pgtest_test.go +++ b/internal/postgres/curation_pgtest_test.go @@ -69,9 +69,12 @@ func TestStoreStarsAndPins(t *testing.T) { ok, err = store.StarSession("missing") require.NoError(t, err, "StarSession missing") assert.False(t, ok, "StarSession missing") - require.NoError(t, store.BulkStarSessions( - []string{"cur-star-2", "missing"}, - ), "BulkStarSessions") + bulkStarred, err := store.BulkStarSessions([]string{"cur-star-2", "missing"}) + require.NoError(t, err, "BulkStarSessions") + assert.Equal(t, []string{"cur-star-2"}, bulkStarred, "starred ids returned") + bulkStarred, err = store.BulkStarSessions([]string{"cur-star-1", "cur-star-2"}) + require.NoError(t, err, "BulkStarSessions already starred") + assert.Empty(t, bulkStarred, "already-starred ids should not be returned") ids, err := store.ListStarredSessionIDs(ctx) require.NoError(t, err, "ListStarredSessionIDs") @@ -83,7 +86,12 @@ func TestStoreStarsAndPins(t *testing.T) { for _, id := range ids { assert.True(t, wantStars[id], "unexpected starred id %q in %v", id, ids) } - require.NoError(t, store.UnstarSession("cur-star-1"), "UnstarSession") + removed, err := store.UnstarSession("cur-star-1") + require.NoError(t, err, "UnstarSession") + require.True(t, removed, "UnstarSession") + removed, err = store.UnstarSession("cur-star-1") + require.NoError(t, err, "UnstarSession no-op") + require.False(t, removed, "UnstarSession no-op") ids, err = store.ListStarredSessionIDs(ctx) require.NoError(t, err, "ListStarredSessionIDs after unstar") require.Len(t, ids, 1) @@ -247,6 +255,109 @@ func TestPushPreservesMultiplePGPinsBySourceUUID(t *testing.T) { assert.Equal(t, 3, pin.Ordinal) } +func TestPushBackfillsLegacyPinSourceUUIDBeforeOrdinalShift(t *testing.T) { + pgURL := testPGURL(t) + cleanPGSchema(t, pgURL) + t.Cleanup(func() { cleanPGSchema(t, pgURL) }) + + local := testDB(t) + ps, err := New( + pgURL, "agentsview", local, + "curation-machine", true, + SyncOptions{}, + ) + require.NoError(t, err, "New sync") + defer ps.Close() + + ctx := context.Background() + require.NoError(t, ps.EnsureSchema(ctx), "EnsureSchema") + + sess := db.Session{ + ID: "pg-legacy-pin-shift", + Project: "proj-curation", + Machine: "local", + Agent: "codex", + MessageCount: 2, + UserMessageCount: 1, + CreatedAt: "2026-05-01T00:00:00Z", + } + require.NoError(t, local.UpsertSession(sess), "UpsertSession first") + require.NoError(t, local.InsertMessages([]db.Message{ + { + SessionID: "pg-legacy-pin-shift", + Ordinal: 0, + Role: "user", + Content: "question", + SourceUUID: "uuid-question", + }, + { + SessionID: "pg-legacy-pin-shift", + Ordinal: 1, + Role: "assistant", + Content: "answer", + SourceUUID: "uuid-answer", + }, + }), "InsertMessages first") + _, err = ps.Push(ctx, false, nil) + require.NoError(t, err, "Push first") + + store, err := NewStore(pgURL, "agentsview", true) + require.NoError(t, err, "NewStore") + defer store.Close() + + note := "legacy pin" + _, err = store.PinMessage("pg-legacy-pin-shift", 1, ¬e) + require.NoError(t, err, "PinMessage") + _, err = ps.pg.ExecContext(ctx, ` + UPDATE pinned_messages + SET source_uuid = '' + WHERE session_id = $1 AND message_id = $2`, + "pg-legacy-pin-shift", 1, + ) + require.NoError(t, err, "clear source_uuid to simulate legacy pin") + + sess.MessageCount = 3 + require.NoError(t, local.UpsertSession(sess), "UpsertSession second") + require.NoError(t, local.ReplaceSessionMessages( + "pg-legacy-pin-shift", + []db.Message{ + { + SessionID: "pg-legacy-pin-shift", + Ordinal: 0, + Role: "user", + Content: "question", + SourceUUID: "uuid-question", + }, + { + SessionID: "pg-legacy-pin-shift", + Ordinal: 1, + Role: "user", + Content: "[compact]", + SourceUUID: "uuid-boundary", + IsCompactBoundary: true, + }, + { + SessionID: "pg-legacy-pin-shift", + Ordinal: 2, + Role: "assistant", + Content: "answer", + SourceUUID: "uuid-answer", + }, + }, + ), "ReplaceSessionMessages") + + _, err = ps.Push(ctx, true, nil) + require.NoError(t, err, "Push rewrite") + + pins, err := store.ListPinnedMessages(ctx, "pg-legacy-pin-shift", "") + require.NoError(t, err, "ListPinnedMessages") + require.Len(t, pins, 1) + assert.Equal(t, int64(2), pins[0].MessageID) + assert.Equal(t, 2, pins[0].Ordinal) + require.NotNil(t, pins[0].Note) + assert.Equal(t, note, *pins[0].Note) +} + func TestReconcilePinnedMessagesPrefersCurrentTargetPin(t *testing.T) { pgURL := testPGURL(t) diff --git a/internal/postgres/integration_test.go b/internal/postgres/integration_test.go index 84e6f7b63..c905ca0c7 100644 --- a/internal/postgres/integration_test.go +++ b/internal/postgres/integration_test.go @@ -163,7 +163,7 @@ func TestPushSecretFindingsReportsChange(t *testing.T) { pushOnce := func() bool { tx, err := ps.pg.BeginTx(ctx, nil) require.NoError(t, err, "begin tx") - changed, err := ps.pushSecretFindings(ctx, tx, sessID) + changed, err := ps.pushSecretFindings(ctx, tx, sessID, sessID) if err != nil { _ = tx.Rollback() t.Fatalf("pushSecretFindings: %v", err) diff --git a/internal/postgres/messages.go b/internal/postgres/messages.go index d80291b7a..13dfff31f 100644 --- a/internal/postgres/messages.go +++ b/internal/postgres/messages.go @@ -2,6 +2,8 @@ package postgres import ( "context" + "database/sql" + "errors" "fmt" "slices" "strings" @@ -280,6 +282,28 @@ func (s *Store) GetResumeModelCounts( return counts, nil } +// GetMessageForMetadataPin returns the stable message identity fields used by +// metadata pin events. PostgreSQL message IDs exposed through the API are +// ordinals, matching PinMessage. +func (s *Store) GetMessageForMetadataPin( + ctx context.Context, sessionID string, messageID int64, +) (*db.Message, error) { + var msg db.Message + err := s.pg.QueryRowContext(ctx, ` + SELECT ordinal, session_id, ordinal, COALESCE(source_uuid, '') + FROM messages + WHERE session_id = $1 AND ordinal = $2`, sessionID, messageID).Scan( + &msg.ID, &msg.SessionID, &msg.Ordinal, &msg.SourceUUID, + ) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("querying message metadata for pin: %w", err) + } + return &msg, nil +} + // SearchSession performs ILIKE substring search within a single // session's messages, returning matching ordinals. func (s *Store) SearchSession( diff --git a/internal/postgres/metadata.go b/internal/postgres/metadata.go new file mode 100644 index 000000000..6d052f618 --- /dev/null +++ b/internal/postgres/metadata.go @@ -0,0 +1,22 @@ +package postgres + +import ( + "context" + + "go.kenn.io/agentsview/internal/db" +) + +// ListMetadataConflicts returns no rows for PostgreSQL read mode because the +// local artifact metadata ledger is not part of the shared SQL mirror. +func (s *Store) ListMetadataConflicts( + context.Context, + []string, +) ([]db.MetadataConflict, error) { + return []db.MetadataConflict{}, nil +} + +// CountMetadataConflicts returns zero for PostgreSQL read mode because the local +// artifact metadata ledger is not part of the shared SQL mirror. +func (s *Store) CountMetadataConflicts(context.Context) (int, error) { + return 0, nil +} diff --git a/internal/postgres/name_source_pgtest_test.go b/internal/postgres/name_source_pgtest_test.go index 3641f0d96..3643acbbb 100644 --- a/internal/postgres/name_source_pgtest_test.go +++ b/internal/postgres/name_source_pgtest_test.go @@ -62,7 +62,10 @@ func TestPushSessionNameRoundTrip(t *testing.T) { // Push via pushSession directly. tx, err := pg.BeginTx(ctx, nil) require.NoError(t, err, "BeginTx") - if err := sync.pushSession(ctx, tx, sess, markerID, nil); err != nil { + if err := sync.pushSession(ctx, tx, sess, pushedSessionIdentity{ + ID: sess.ID, + Machine: sess.Machine, + }, markerID, nil); err != nil { _ = tx.Rollback() t.Fatalf("pushSession: %v", err) } @@ -105,7 +108,10 @@ func TestPushSessionNameRoundTrip(t *testing.T) { tx2, err := pg.BeginTx(ctx, nil) require.NoError(t, err, "BeginTx (second)") - if err := sync.pushSession(ctx, tx2, sess, markerID, nil); err != nil { + if err := sync.pushSession(ctx, tx2, sess, pushedSessionIdentity{ + ID: sess.ID, + Machine: sess.Machine, + }, markerID, nil); err != nil { _ = tx2.Rollback() t.Fatalf("pushSession (second): %v", err) } diff --git a/internal/postgres/push.go b/internal/postgres/push.go index 6a7c581f5..bdd17eed7 100644 --- a/internal/postgres/push.go +++ b/internal/postgres/push.go @@ -15,8 +15,10 @@ import ( "sort" "strconv" "strings" + "sync" "time" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/export" ) @@ -27,6 +29,9 @@ const ( sessionAliasBackfillStateKey = "pg_session_alias_backfill_v1" projectIdentityPublicationStateKey = "project_identity_publication_revision_v2" transcriptRevisionBackfillStateKey = "pg_transcript_revision_backfill_v1" + artifactIdentityModeStateKey = "pg_artifact_identity_v1" + artifactOwnerMarkerPrefix = "artifact-origin:" + legacyArtifactIdentityMode = "legacy" ) // pushMarkerIDStateKey names the local sync-state entry holding this DB's @@ -40,6 +45,7 @@ const ( var errSessionOwnershipConflict = errors.New("session ownership conflict") var errSessionExcluded = errors.New("session excluded") +var errArtifactReplicaExists = errors.New("artifact replica already exists") type pushBoundaryState struct { Cutoff string `json:"cutoff"` @@ -170,6 +176,26 @@ func (s *Sync) Push( if err != nil { return result, err } + localArtifactOrigin, err := artifact.StoredOrigin(s.local) + if err != nil { + return result, err + } + artifactIdentityMode := currentArtifactIdentityMode(localArtifactOrigin) + storedArtifactIdentityMode, err := state.GetSyncState( + artifactIdentityModeStateKey, + ) + if err != nil { + return result, fmt.Errorf( + "reading %s: %w", artifactIdentityModeStateKey, err, + ) + } + if lastPush != "" && storedArtifactIdentityMode != artifactIdentityMode { + log.Printf( + "pgsync: artifact identity mode changed; forcing full push", + ) + lastPush = "" + full = true + } markerMachine, markerMachineAliases, markerExists, err := s.pgPushMarkerMachineState(ctx, markerID) if err != nil { return result, err @@ -298,6 +324,7 @@ func (s *Sync) Push( var priorFingerprints map[string]string sessionFingerprints := make(map[string]string, len(sessionByID)) + sessionIdentities := make(map[string]pushedSessionIdentity, len(sessionByID)) if !full { var bErr error priorFingerprints, _, _, bErr = readBoundaryAndFingerprints( @@ -308,8 +335,32 @@ func (s *Sync) Push( } } + artifactImportedSessions, err := artifact.ImportedSessionIDs( + s.local, mapKeys(sessionByID), + ) + if err != nil { + return result, err + } + ctx, err = s.preloadPGSessionOwners(ctx, s.pushIdentityOwnerCandidateIDs( + sessionByID, artifactImportedSessions, localArtifactOrigin, + markerID, legacyMarkerMachines, + )) + if err != nil { + return result, err + } + for id, sess := range sessionByID { + _, artifactImported := artifactImportedSessions[sess.ID] + identity, err := s.resolvePushedSessionIdentity( + ctx, sess, localArtifactOrigin, artifactImported, + markerID, legacyMarkerMachines, + ) + if err != nil { + return result, err + } + sessionIdentities[id] = identity + } if err := purgePGExcludedPushSessions( - ctx, s.pg, sessionByID, + ctx, s.pg, sessionByID, sessionIdentities, ); err != nil { return result, err } @@ -322,10 +373,14 @@ func (s *Sync) Push( "computing local usage event fingerprints: %w", err, ) } - // The fingerprint loop issues several local queries per candidate - // session; on a full push that covers every session and runs for - // minutes, so it reports its own progress phase rather than sitting - // silent until the first batch lands. + if err := s.markRelationshipConflicts( + ctx, sessionByID, sessionIdentities, markerID, legacyMarkerMachines, + ); err != nil { + return result, err + } + // Fingerprint preparation resolves relationship targets and reads local + // dependency state in bounded batches. Full pushes can spend minutes here, + // so report this phase rather than staying silent until writes begin. log.Printf("pgsync: computing push fingerprints for %d candidate session(s)", len(sessionByID)) reportPrepare := func(done int) { @@ -349,6 +404,36 @@ func (s *Sync) Push( return result, err } for _, id := range chunk { + sess := sessionByID[id] + identity := sessionIdentities[id] + prepared++ + if identity.Conflict { + if prepared%pushPrepareProgressStride == 0 { + reportPrepare(prepared) + } + continue + } + // Resolve relationship ids (source/parent) to the ids their target + // sessions are stored under in PG once every identity is known, so the + // fingerprint and the written row agree and child rows link correctly + // even when a target was pushed under a collision-avoidance prefix. + resolver := s.newRelationshipResolver( + sessionIdentities, identity, + markerID, legacyMarkerMachines, + ) + if sess.SourceSessionID != "" || sess.ParentSessionID != nil { + resolvedSource, err := resolver.resolve(ctx, sess.SourceSessionID) + if err != nil { + return result, err + } + sess.SourceSessionID = resolvedSource + resolvedParent, err := resolver.resolvePtr(ctx, sess.ParentSessionID) + if err != nil { + return result, err + } + sess.ParentSessionID = resolvedParent + sessionByID[id] = sess + } usageFP, usageKnown := usageFingerprints[id] dependencyFP, err := depState.dependencyFingerprint( s.local, id, usageFP, usageKnown, @@ -359,12 +444,14 @@ func (s *Sync) Push( id, err, ) } - sess := sessionByID[id] + subagentFP, err := resolver.resolvedSubagentLinkFingerprint(ctx, id) + if err != nil { + return result, err + } sessionFingerprints[id] = sessionPushFingerprint( - sess, pushedSessionMachine(sess, s.machine), - usageFP, markerID, dependencyFP, + sess, identity.ID, identity.Machine, + usageFP, identity.effectiveOwnerMarker(markerID), dependencyFP+subagentFP, ) - prepared++ if prepared%pushPrepareProgressStride == 0 { reportPrepare(prepared) } @@ -374,6 +461,9 @@ func (s *Sync) Push( if len(priorFingerprints) > 0 { for id := range sessionByID { + if sessionIdentities[id].Conflict { + continue + } if priorFingerprints[id] == sessionFingerprints[id] { delete(sessionByID, id) } @@ -432,6 +522,11 @@ func (s *Sync) Push( if err := s.syncProjectIdentityObservations(ctx, full); err != nil { return result, err } + if err := completeArtifactIdentityMode( + state, artifactIdentityMode, result, + ); err != nil { + return result, err + } result.Vectors, err = s.runVectorPushPhase(ctx, full, nil, onProgress) if err != nil { return result, err @@ -440,61 +535,32 @@ func (s *Sync) Push( return result, nil } - var pushed []db.Session - // Sessions whose individual retry also failed: their PG sessions/messages - // rows are stale or absent, so the vector phase must not push their newer - // local vectors ahead of them. - var failedSessions map[string]struct{} - const batchSize = 50 - for i := 0; i < len(sessions); i += batchSize { - end := min(i+batchSize, len(sessions)) - batch := sessions[i:end] - - batchResult, err := s.pushBatch( - ctx, batch, full, markerID, legacyMarkerMachines, - usageFingerprints, &pushed, - ) - if err != nil { - return result, err - } - if batchResult.ok { - result.SessionsPushed += batchResult.sessions - result.MessagesPushed += batchResult.messages - result.SkippedConflicts += batchResult.skippedConflicts - } else { - // Batch failed — retry each session individually - // so one bad session doesn't block the rest. - for _, sess := range batch { - sr, retryErr := s.pushBatch( - ctx, []db.Session{sess}, - full, markerID, legacyMarkerMachines, - usageFingerprints, &pushed, - ) - if retryErr != nil { - return result, retryErr - } - if sr.ok { - result.SessionsPushed += sr.sessions - result.MessagesPushed += sr.messages - result.SkippedConflicts += sr.skippedConflicts - } else { - result.Errors++ - if failedSessions == nil { - failedSessions = make(map[string]struct{}) - } - failedSessions[sess.ID] = struct{}{} - } - } - } - if onProgress != nil { - onProgress(PushProgress{ - SessionsDone: end, - SessionsTotal: len(sessions), - MessagesDone: result.MessagesPushed, - SkippedConflicts: result.SkippedConflicts, - Errors: result.Errors, - }) - } + sink := pgSessionSink{ + sync: s, + full: full, + markerID: markerID, + legacyMarkerMachines: legacyMarkerMachines, + sessionUsageFingerprints: usageFingerprints, + identities: sessionIdentities, + } + pushed, err := drainSessionBatches( + ctx, sessions, sink, &result, onProgress, + ) + if err != nil { + return result, err + } + // Sessions not written by the session phase include failed retries and + // ownership conflicts. Their vectors must not advance ahead of the + // sessions/messages rows they depend on. + failedSessions := make(map[string]struct{}, len(sessions)-len(pushed)) + for _, sess := range sessions { + failedSessions[sess.ID] = struct{}{} + } + for _, sess := range pushed { + delete(failedSessions, sess.ID) + } + if len(failedSessions) == 0 { + failedSessions = nil } if s.isFiltered() { @@ -552,6 +618,11 @@ func (s *Sync) Push( result.Errors, ) } + if err := completeArtifactIdentityMode( + state, artifactIdentityMode, result, + ); err != nil { + return result, err + } result.Vectors, err = s.runVectorPushPhase(ctx, full, failedSessions, onProgress) if err != nil { return result, err @@ -724,6 +795,99 @@ func filterProjectIdentityObservations( return out } +// sessionBatchSink persists batches of sessions to a target during a push. +// Extracting this seam keeps the batch/retry/progress orchestration +// (drainSessionBatches) free of any SQL: PostgreSQL is one sink today, and the +// artifact exporter can become another without duplicating the loop. +type sessionBatchSink interface { + // writeBatch persists batch atomically and appends successfully written + // sessions to *pushed. It returns ok=false (and no error) when the batch + // failed for a reason the caller should recover from by retrying each + // session individually. A non-nil error is fatal and aborts the push. + writeBatch( + ctx context.Context, batch []db.Session, pushed *[]db.Session, + ) (batchResult, error) +} + +// pgSessionSink writes batches to PostgreSQL via Sync.pushBatch. It binds the +// per-push parameters (full mode, push marker identity) so the orchestration +// does not need to thread them through. +type pgSessionSink struct { + sync *Sync + full bool + markerID string + legacyMarkerMachines []string + sessionUsageFingerprints map[string]string + identities map[string]pushedSessionIdentity +} + +func (p pgSessionSink) writeBatch( + ctx context.Context, batch []db.Session, pushed *[]db.Session, +) (batchResult, error) { + return p.sync.pushBatch( + ctx, batch, p.full, p.markerID, p.legacyMarkerMachines, + p.sessionUsageFingerprints, pushed, p.identities, + ) +} + +// drainSessionBatches pushes sessions through sink in fixed-size batches, +// accumulating counts into result and reporting progress after each batch. +// When a batch fails (ok=false) it retries each session individually so one +// bad session does not block the rest; sessions that still fail are counted as +// errors. It returns the sessions successfully written, in push order. +func drainSessionBatches( + ctx context.Context, + sessions []db.Session, + sink sessionBatchSink, + result *PushResult, + onProgress func(PushProgress), +) ([]db.Session, error) { + var pushed []db.Session + const batchSize = 50 + for i := 0; i < len(sessions); i += batchSize { + end := min(i+batchSize, len(sessions)) + batch := sessions[i:end] + + batchResult, err := sink.writeBatch(ctx, batch, &pushed) + if err != nil { + return pushed, err + } + if batchResult.ok { + result.SessionsPushed += batchResult.sessions + result.MessagesPushed += batchResult.messages + result.SkippedConflicts += batchResult.skippedConflicts + } else { + // Batch failed — retry each session individually + // so one bad session doesn't block the rest. + for _, sess := range batch { + sr, retryErr := sink.writeBatch( + ctx, []db.Session{sess}, &pushed, + ) + if retryErr != nil { + return pushed, retryErr + } + if sr.ok { + result.SessionsPushed += sr.sessions + result.MessagesPushed += sr.messages + result.SkippedConflicts += sr.skippedConflicts + } else { + result.Errors++ + } + } + } + if onProgress != nil { + onProgress(PushProgress{ + SessionsDone: end, + SessionsTotal: len(sessions), + MessagesDone: result.MessagesPushed, + SkippedConflicts: result.SkippedConflicts, + Errors: result.Errors, + }) + } + } + return pushed, nil +} + // pgPushMarkerMachineState reports whether this host's push marker is present // in PG and returns the current machine plus legacy machine aliases stored with // the marker. @@ -982,11 +1146,12 @@ func (s *Sync) pushBatch( legacyMarkerMachines []string, sessionUsageFingerprints map[string]string, pushed *[]db.Session, + identities map[string]pushedSessionIdentity, ) (batchResult, error) { preloadComparisons := len(batch) > 0 && !full result, err := s.pushBatchAttempt( ctx, batch, full, markerID, legacyMarkerMachines, - sessionUsageFingerprints, pushed, preloadComparisons, + sessionUsageFingerprints, pushed, identities, preloadComparisons, ) if err == nil || !errors.Is(err, errPushComparisonPreload) { return result, err @@ -998,7 +1163,7 @@ func (s *Sync) pushBatch( ) return s.pushBatchAttempt( ctx, batch, full, markerID, legacyMarkerMachines, - sessionUsageFingerprints, pushed, false, + sessionUsageFingerprints, pushed, identities, false, ) } @@ -1010,6 +1175,7 @@ func (s *Sync) pushBatchAttempt( legacyMarkerMachines []string, sessionUsageFingerprints map[string]string, pushed *[]db.Session, + identities map[string]pushedSessionIdentity, preloadComparisons bool, ) (batchResult, error) { tx, err := s.pg.BeginTx(ctx, nil) @@ -1024,7 +1190,17 @@ func (s *Sync) pushBatchAttempt( skippedConflicts := 0 sessionIDs := make([]string, 0, len(batch)) for _, sess := range batch { - sessionIDs = append(sessionIDs, sess.ID) + identity := identities[sess.ID] + if identity.Conflict { + continue + } + if identity.ID == "" { + identity = pushedSessionIdentity{ + ID: sess.ID, + Machine: pushedSessionMachine(sess, s.machine), + } + } + sessionIDs = append(sessionIDs, identity.ID) } comparisons := (*pushMessageComparison)(nil) if preloadComparisons && len(sessionIDs) > 0 { @@ -1041,9 +1217,23 @@ func (s *Sync) pushBatchAttempt( } for _, sess := range batch { + identity := identities[sess.ID] + if identity.Conflict { + skippedConflicts++ + continue + } + if identity.ID == "" { + identity = pushedSessionIdentity{ + ID: sess.ID, + Machine: pushedSessionMachine(sess, s.machine), + } + } if err := s.pushSession( - ctx, tx, sess, markerID, legacyMarkerMachines, + ctx, tx, sess, identity, markerID, legacyMarkerMachines, ); err != nil { + if errors.Is(err, errArtifactReplicaExists) { + continue + } if errors.Is(err, errSessionOwnershipConflict) { skippedConflicts++ continue @@ -1060,8 +1250,11 @@ func (s *Sync) pushBatchAttempt( return batchResult{}, nil } + resolver := s.newRelationshipResolver( + identities, identity, markerID, legacyMarkerMachines, + ) msgCount, err := s.pushMessages( - ctx, tx, sess.ID, full, + ctx, tx, sess.ID, identity.ID, resolver, full, sessionUsageFingerprints, comparisons, ) if err != nil { @@ -1074,7 +1267,9 @@ func (s *Sync) pushBatchAttempt( return batchResult{}, nil } - findingsChanged, err := s.pushSecretFindings(ctx, tx, sess.ID) + findingsChanged, err := s.pushSecretFindings( + ctx, tx, sess.ID, identity.ID, + ) if err != nil { log.Printf( "pgsync: secret findings %s: %v", @@ -1094,11 +1289,11 @@ func (s *Sync) pushBatchAttempt( UPDATE sessions SET updated_at = NOW() WHERE id = $1`, - sess.ID, + identity.ID, ); err != nil { log.Printf( "pgsync: bumping updated_at %s: %v", - sess.ID, err, + identity.ID, err, ) _ = tx.Rollback() *pushed = (*pushed)[:len(*pushed)-n] @@ -1285,6 +1480,25 @@ func completeTranscriptRevisionBackfill( return markTranscriptRevisionBackfillDone(local) } +func currentArtifactIdentityMode(origin string) string { + if origin == "" { + return legacyArtifactIdentityMode + } + return artifactOwnerMarkerPrefix + origin +} + +func completeArtifactIdentityMode( + local syncStateStore, mode string, result PushResult, +) error { + if result.Errors > 0 { + return nil + } + if err := local.SetSyncState(artifactIdentityModeStateKey, mode); err != nil { + return fmt.Errorf("updating %s: %w", artifactIdentityModeStateKey, err) + } + return nil +} + func persistPushTargetFingerprint( local syncStateStore, fingerprint string, @@ -1449,12 +1663,19 @@ func pgExcludedSessionIDsQuery(ids []string) (string, []any) { } func purgePGExcludedPushSessions( - ctx context.Context, pg *sql.DB, sessionByID map[string]db.Session, + ctx context.Context, + pg *sql.DB, + sessionByID map[string]db.Session, + identities map[string]pushedSessionIdentity, ) error { tombstoneIDsBySession := make(map[string][]string, len(sessionByID)) candidateIDs := []string{} for id, sess := range sessionByID { - tombstoneIDs := pgSessionTombstoneIDs(sess) + identity := identities[id] + if identity.Conflict || identity.ID == "" { + continue + } + tombstoneIDs := pgSessionTombstoneIDsForPushedID(sess, identity.ID) tombstoneIDsBySession[id] = tombstoneIDs candidateIDs = append(candidateIDs, tombstoneIDs...) } @@ -1519,9 +1740,9 @@ func deletePGExcludedSessionRows( } func deletePGSessionIfExcluded( - ctx context.Context, tx *sql.Tx, sess db.Session, + ctx context.Context, tx *sql.Tx, sess db.Session, pushedID string, ) (bool, error) { - ids := pgSessionTombstoneIDs(sess) + ids := pgSessionTombstoneIDsForPushedID(sess, pushedID) excluded, err := readPGExcludedSessionIDs(ctx, tx, ids) if err != nil { return false, err @@ -1538,17 +1759,25 @@ func deletePGSessionIfExcluded( return true, nil } +func pgSessionTombstoneIDsForPushedID(sess db.Session, pushedID string) []string { + if pushedID != "" { + sess.ID = pushedID + } + return pgSessionTombstoneIDs(sess) +} + // sessionPushFingerprint builds the change-detection fingerprint for a -// session. pushedMachine is the value pushSession actually writes to PG -// (pushedSessionMachine), not the raw sess.Machine: a "local"/empty sentinel -// row is written under the fallback machine, so the fingerprint must track the -// fallback to force a re-push when s.machine changes. +// session. pushedID and pushedMachine are the values pushSession actually +// writes to PG. They may differ from sess.ID/sess.Machine when a native +// session ID collides across machines or a "local"/empty sentinel row is +// written under the fallback machine, so the fingerprint must track them to +// force a re-push when the resolved PG identity changes. func sessionPushFingerprint( - sess db.Session, pushedMachine, + sess db.Session, pushedID, pushedMachine, usageEventFingerprint, ownerMarker, dependencyFingerprint string, ) string { fields := []string{ - sess.ID, + pushedID, sess.Project, pushedMachine, ownerMarker, @@ -1613,21 +1842,681 @@ func sessionPushFingerprint( sess.SecretsRulesVersion, usageEventFingerprint, } - var b strings.Builder - for _, f := range fields { - fmt.Fprintf(&b, "%d:%s", len(f), f) + var b strings.Builder + for _, f := range fields { + fmt.Fprintf(&b, "%d:%s", len(f), f) + } + return b.String() +} + +// pushedSessionMachine resolves the machine field for a PG row. Old rows +// pushed before this fix with machine="local" will be repaired gradually as +// each session is modified (message count change, etc.) and re-fingerprinted. +func pushedSessionMachine(sess db.Session, fallbackMachine string) string { + if sess.Machine != "" && sess.Machine != "local" { + return sess.Machine + } + return fallbackMachine +} + +func artifactPushIdentity( + sess db.Session, localOrigin string, artifactImported bool, +) (id, machine, ownerMarker string, ok bool) { + origin := "" + switch { + case localOrigin != "" && (sess.Machine == "" || sess.Machine == "local"): + origin = localOrigin + case localOrigin != "" && artifactImported && + sess.Machine != "" && sess.Machine != "local" && + strings.HasPrefix(sess.ID, sess.Machine+"~"): + origin = sess.Machine + default: + return "", "", "", false + } + return prefixedSessionID(origin, sess.ID), origin, + artifactOwnerMarkerPrefix + origin, true +} + +type pushedSessionIdentity struct { + ID string + Machine string + OwnerMarker string + LegacyOwnerMarkers []string + ArtifactReplica bool + AliasIDs []string + LegacyDuplicateID string + Conflict bool +} + +func (i pushedSessionIdentity) effectiveOwnerMarker(fallback string) string { + if i.OwnerMarker != "" { + return i.OwnerMarker + } + return fallback +} + +func (s *Sync) markRelationshipConflicts( + ctx context.Context, + sessionByID map[string]db.Session, + identities map[string]pushedSessionIdentity, + markerID string, + legacyMarkerMachines []string, +) error { + changed := true + for changed { + changed = false + for id, sess := range sessionByID { + identity := identities[id] + if identity.Conflict { + continue + } + resolver := s.newRelationshipResolver( + identities, identity, markerID, legacyMarkerMachines, + ) + if sess.SourceSessionID != "" { + if _, err := resolver.resolve(ctx, sess.SourceSessionID); err != nil { + if errors.Is(err, errSessionOwnershipConflict) { + identity.Conflict = true + identities[id] = identity + changed = true + continue + } + return err + } + } + if sess.ParentSessionID != nil && *sess.ParentSessionID != "" { + if _, err := resolver.resolve(ctx, *sess.ParentSessionID); err != nil { + if errors.Is(err, errSessionOwnershipConflict) { + identity.Conflict = true + identities[id] = identity + changed = true + continue + } + return err + } + } + if _, err := resolver.sessionNeedsSubagentRewrite(ctx, id); err != nil { + if errors.Is(err, errSessionOwnershipConflict) { + identity.Conflict = true + identities[id] = identity + changed = true + continue + } + return err + } + } + } + return nil +} + +func initialPushedSessionIdentity( + sess db.Session, + fallbackMachine string, + localArtifactOrigin string, + artifactImported bool, + markerID string, +) (pushedSessionIdentity, bool) { + identity := pushedSessionIdentity{ + ID: sess.ID, + Machine: pushedSessionMachine(sess, fallbackMachine), + } + if id, machine, ownerMarker, ok := artifactPushIdentity( + sess, localArtifactOrigin, artifactImported, + ); ok { + identity.ID = id + identity.Machine = machine + identity.OwnerMarker = ownerMarker + identity.LegacyOwnerMarkers = []string{markerID} + identity.ArtifactReplica = artifactImported + identity.AliasIDs = artifactPushAliasIDs(sess, id, machine) + return identity, true + } + return identity, false +} + +// resolvePushedSessionIdentity decides the PG id a local session is stored +// under. A session this sync owns -- by matching push marker, or an adoptable +// legacy/ownerless row (see sameSessionOwner) -- is updated in place: an +// existing row under the current or any prior machine prefix is reused, so +// machine renames and marker adoption keep updating the same row instead of +// creating a duplicate. Only a bare id already held by a different owner +// collides; that session is stored under the current machine prefix so both +// rows coexist instead of ping-ponging (issue 655). A collision is not rejected +// here: pushSession skips the conflicting row, so one conflicting session never +// fails the whole push. +func (s *Sync) resolvePushedSessionIdentity( + ctx context.Context, + sess db.Session, + localArtifactOrigin string, + artifactImported bool, + markerID string, + legacyMarkerMachines []string, +) (pushedSessionIdentity, error) { + identity, artifactIdentity := initialPushedSessionIdentity( + sess, s.machine, localArtifactOrigin, artifactImported, markerID, + ) + canonicalID := identity.ID + id, err := s.resolveOwnedPushIdentityID( + ctx, identity.ID, identity, markerID, legacyMarkerMachines, + ) + if err != nil { + if errors.Is(err, errSessionOwnershipConflict) { + identity.ID = "" + identity.Conflict = true + return identity, nil + } + return pushedSessionIdentity{}, err + } + identity.ID = id + if artifactIdentity && id == canonicalID { + legacyDuplicateID, conflict, resolveErr := + s.artifactLegacyDuplicateCandidate( + ctx, canonicalID, identity, markerID, + ) + if resolveErr != nil { + return pushedSessionIdentity{}, resolveErr + } + if conflict { + identity.ID = "" + identity.Conflict = true + return identity, nil + } + identity.LegacyDuplicateID = legacyDuplicateID + } + return identity, nil +} + +func (s *Sync) pushIdentityOwnerCandidateIDs( + sessionByID map[string]db.Session, + artifactImportedSessions map[string]struct{}, + localArtifactOrigin string, + markerID string, + legacyMarkerMachines []string, +) []string { + ids := make(map[string]struct{}, len(sessionByID)*3) + for _, sess := range sessionByID { + _, artifactImported := artifactImportedSessions[sess.ID] + identity, _ := initialPushedSessionIdentity( + sess, s.machine, localArtifactOrigin, artifactImported, markerID, + ) + ids[identity.ID] = struct{}{} + for _, machine := range pushIDMachinePrefixes( + identity.Machine, legacyMarkerMachines, + ) { + candidateID := prefixedSessionID(machine, identity.ID) + if candidateID != identity.ID { + ids[candidateID] = struct{}{} + } + } + for _, aliasID := range uniqueNonEmptyStrings(identity.AliasIDs) { + ids[aliasID] = struct{}{} + } + } + result := make([]string, 0, len(ids)) + for id := range ids { + if id != "" { + result = append(result, id) + } + } + sort.Strings(result) + return result +} + +// artifactLegacyDuplicateCandidate recognizes the narrow upgrade state where +// an importer already created the stable artifact id while this origin still +// owns its pre-artifact bare row. Both rows must already have the exact owners +// expected for that history. A foreign-owned bare alias is an unrelated id +// collision and is ignored; only a current-marker alias with invalid canonical +// ownership is surfaced as a conflict rather than guessed safe to merge. +func (s *Sync) artifactLegacyDuplicateCandidate( + ctx context.Context, + canonicalID string, + identity pushedSessionIdentity, + markerID string, +) (legacyDuplicateID string, conflict bool, err error) { + if canonicalID == "" || identity.OwnerMarker == "" || markerID == "" { + return "", false, nil + } + canonicalMachine, canonicalOwnerMarker, canonicalExists, canonicalErr := + s.pgSessionOwner(ctx, canonicalID) + if canonicalErr != nil { + return "", false, canonicalErr + } + if !canonicalExists { + return "", false, nil + } + for _, aliasID := range uniqueNonEmptyStrings(identity.AliasIDs) { + if aliasID == canonicalID { + continue + } + _, aliasOwnerMarker, aliasExists, aliasErr := s.pgSessionOwner( + ctx, aliasID, + ) + if aliasErr != nil { + return "", false, aliasErr + } + if !aliasExists { + continue + } + if aliasOwnerMarker != markerID { + continue + } + if canonicalMachine != identity.Machine || + canonicalOwnerMarker != identity.OwnerMarker { + return "", true, nil + } + return aliasID, false, nil + } + return "", false, nil +} + +func (s *Sync) resolveOwnedPushIdentityID( + ctx context.Context, + localID string, + identity pushedSessionIdentity, + markerID string, + legacyMarkerMachines []string, +) (string, error) { + machine := identity.Machine + currentPrefixedID := prefixedSessionID(machine, localID) + currentPrefixConflict := false + for _, candidate := range pushIDMachinePrefixes(machine, legacyMarkerMachines) { + prefixedID := prefixedSessionID(candidate, localID) + if prefixedID == localID { + continue + } + existingMachine, ownerMarker, ok, err := s.pgSessionOwner(ctx, prefixedID) + if err != nil { + return "", err + } + if ok && samePushedSessionOwner( + ownerMarker, existingMachine, identity, markerID, legacyMarkerMachines, + ) { + return prefixedID, nil + } + if ok && prefixedID == currentPrefixedID { + currentPrefixConflict = true + } + } + existingMachine, ownerMarker, ok, err := s.pgSessionOwner(ctx, localID) + if err != nil { + return "", err + } + if ok && samePushedSessionOwner( + ownerMarker, existingMachine, identity, markerID, legacyMarkerMachines, + ) { + return localID, nil + } + for _, aliasID := range uniqueNonEmptyStrings(identity.AliasIDs) { + if aliasID == localID { + continue + } + aliasMachine, aliasOwnerMarker, aliasExists, aliasErr := s.pgSessionOwner( + ctx, aliasID, + ) + if aliasErr != nil { + return "", aliasErr + } + if aliasExists && samePushedSessionOwner( + aliasOwnerMarker, aliasMachine, identity, + markerID, legacyMarkerMachines, + ) { + return aliasID, nil + } + } + if ok && !samePushedSessionOwner( + ownerMarker, existingMachine, identity, markerID, legacyMarkerMachines, + ) { + if localID == currentPrefixedID || currentPrefixConflict { + return "", errSessionOwnershipConflict + } + return prefixedSessionID(machine, localID), nil + } + return localID, nil +} + +func artifactPushAliasIDs( + sess db.Session, canonicalID, origin string, +) []string { + aliasID := sess.ID + if sess.Machine != "" && sess.Machine != "local" { + aliasID = strings.TrimPrefix(sess.ID, origin+"~") + } + if aliasID == "" || aliasID == canonicalID { + return nil + } + return []string{aliasID} +} + +func artifactCanonicalAliasIDs(origin, canonicalID string) []string { + prefix := origin + "~" + if origin == "" || !strings.HasPrefix(canonicalID, prefix) { + return nil + } + aliasID := strings.TrimPrefix(canonicalID, prefix) + if aliasID == "" || aliasID == canonicalID { + return nil + } + return []string{aliasID} +} + +// pushIDMachinePrefixes lists the machine names whose id prefixes identify rows +// this owner may already hold: the current machine first, then prior machine +// names (legacy marker machines) the same marker pushed under before a rename. +// The current machine is not repeated when it also appears in the legacy set. +func pushIDMachinePrefixes(machine string, legacyMarkerMachines []string) []string { + prefixes := make([]string, 0, len(legacyMarkerMachines)+1) + if machine != "" { + prefixes = append(prefixes, machine) + } + for _, m := range legacyMarkerMachines { + if m == machine || m == "" { + continue + } + prefixes = append(prefixes, m) + } + return prefixes +} + +type pgSessionOwnerRecord struct { + machine string + ownerMarker string + exists bool +} + +type pgSessionOwnerCache struct { + mu sync.Mutex + entries map[string]pgSessionOwnerRecord +} + +type pgSessionOwnerCacheContextKey struct{} + +func (c *pgSessionOwnerCache) get(id string) (pgSessionOwnerRecord, bool) { + c.mu.Lock() + defer c.mu.Unlock() + record, ok := c.entries[id] + return record, ok +} + +func (c *pgSessionOwnerCache) put(id string, record pgSessionOwnerRecord) { + c.mu.Lock() + defer c.mu.Unlock() + c.entries[id] = record +} + +// preloadPGSessionOwners resolves a candidate set in one PG round trip and +// records both hits and misses. pgSessionOwner reuses this cache and memoizes +// any relationship targets discovered later in the same push. +func (s *Sync) preloadPGSessionOwners( + ctx context.Context, ids []string, +) (context.Context, error) { + unique := uniqueNonEmptyStrings(ids) + cache := &pgSessionOwnerCache{ + entries: make(map[string]pgSessionOwnerRecord, len(unique)), + } + for _, id := range unique { + cache.entries[id] = pgSessionOwnerRecord{} + } + if len(unique) == 0 { + return context.WithValue(ctx, pgSessionOwnerCacheContextKey{}, cache), nil + } + rows, err := s.pg.QueryContext(ctx, ` + SELECT id, machine, owner_marker + FROM sessions + WHERE id = ANY($1)`, unique) + if err != nil { + return ctx, fmt.Errorf("preloading pg session owners: %w", err) + } + defer rows.Close() + for rows.Next() { + var id, machine string + var ownerMarker sql.NullString + if err := rows.Scan(&id, &machine, &ownerMarker); err != nil { + return ctx, fmt.Errorf("scanning pg session owner: %w", err) + } + cache.entries[id] = pgSessionOwnerRecord{ + machine: machine, ownerMarker: ownerMarker.String, exists: true, + } + } + if err := rows.Err(); err != nil { + return ctx, fmt.Errorf("iterating pg session owners: %w", err) + } + return context.WithValue(ctx, pgSessionOwnerCacheContextKey{}, cache), nil +} + +// pgSessionOwner returns the machine and owner_marker of a PG session row, and +// whether it exists. owner_marker is empty for legacy rows pushed before the +// marker model. +func (s *Sync) pgSessionOwner( + ctx context.Context, + id string, +) (string, string, bool, error) { + cache, _ := ctx.Value(pgSessionOwnerCacheContextKey{}).(*pgSessionOwnerCache) + if cache != nil { + if record, ok := cache.get(id); ok { + return record.machine, record.ownerMarker, record.exists, nil + } + } + var machine string + var ownerMarker sql.NullString + err := s.pg.QueryRowContext(ctx, + `SELECT machine, owner_marker FROM sessions WHERE id = $1`, + id, + ).Scan(&machine, &ownerMarker) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + if cache != nil { + cache.put(id, pgSessionOwnerRecord{}) + } + return "", "", false, nil + } + return "", "", false, fmt.Errorf( + "reading pg session owner for %s: %w", + id, err, + ) + } + if cache != nil { + cache.put(id, pgSessionOwnerRecord{ + machine: machine, ownerMarker: ownerMarker.String, exists: true, + }) + } + return machine, ownerMarker.String, true, nil +} + +func prefixedSessionID(machine, id string) string { + if machine == "" || id == "" { + return id + } + prefix := machine + "~" + if strings.HasPrefix(id, prefix) { + return id + } + return prefix + id +} + +// relationshipResolver maps the local session ids that appear in a session's +// relationship fields (source/parent) and in tool-call subagent links to the +// ids those target sessions are stored under in PG. When a target session was +// pushed under a collision-avoidance prefix (machine~id), the unprefixed local +// id would dangle or point at a different machine's session; resolving keeps +// child and subagent rows linked to the correct PG row. +// +// A child session and the sessions it references (parent, source, subagents) +// originate on the same machine, so the resolver is scoped to one machine and +// mirrors resolvePushedSessionIdentity's ownership rule. +type relationshipResolver struct { + sync *Sync + identities map[string]pushedSessionIdentity + identity pushedSessionIdentity + markerID string + legacyMarkerMachines []string + cache map[string]string +} + +func (s *Sync) newRelationshipResolver( + identities map[string]pushedSessionIdentity, + identity pushedSessionIdentity, + markerID string, + legacyMarkerMachines []string, +) relationshipResolver { + return relationshipResolver{ + sync: s, + identities: identities, + identity: identity, + markerID: markerID, + legacyMarkerMachines: legacyMarkerMachines, + cache: make(map[string]string), + } +} + +// resolve maps a local session id to the id it is stored under in PG. It +// prefers the in-run identity map (which already reflects collision prefixing +// for every session in this push) and falls back to committed PG state for +// targets outside the push window. +func (r relationshipResolver) resolve( + ctx context.Context, localID string, +) (string, error) { + if localID == "" { + return "", nil + } + if identity, ok := r.identities[localID]; ok { + if identity.Conflict { + return "", errSessionOwnershipConflict + } + if identity.ID != "" { + return identity.ID, nil + } + } + if resolved, ok := r.cache[localID]; ok { + return resolved, nil + } + resolved, err := r.lookup(ctx, localID) + if err != nil { + return "", err + } + r.cache[localID] = resolved + return resolved, nil +} + +// resolvePtr resolves a relationship id held behind a pointer, preserving nil +// and returning the original pointer when the value is unchanged. +func (r relationshipResolver) resolvePtr( + ctx context.Context, localID *string, +) (*string, error) { + if localID == nil || *localID == "" { + return localID, nil + } + resolved, err := r.resolve(ctx, *localID) + if err != nil { + return nil, err + } + if resolved == *localID { + return localID, nil + } + return &resolved, nil +} + +// lookup resolves an id absent from the in-run identity map by consulting +// committed PG state, sharing resolveOwnedPushIdentityID with identity +// resolution: it reuses a row owned under the current or any legacy machine +// prefix (so a target pushed before a rename still resolves), returns the +// current prefix when the bare id is held by another owner, and otherwise keeps +// the bare id. +func (r relationshipResolver) lookup( + ctx context.Context, localID string, +) (string, error) { + identity := r.identity + if r.identity.OwnerMarker != "" { + localID = prefixedSessionID(r.identity.Machine, localID) + identity.AliasIDs = artifactCanonicalAliasIDs( + r.identity.Machine, localID, + ) + } + return r.sync.resolveOwnedPushIdentityID( + ctx, localID, identity, r.markerID, r.legacyMarkerMachines, + ) +} + +// rewriteSubagentIDs resolves the subagent session link on each tool call and +// result event in msgs in place so they reference the PG ids of the subagent +// sessions rather than their unprefixed local ids. +func (r relationshipResolver) rewriteSubagentIDs( + ctx context.Context, msgs []db.Message, +) error { + for i := range msgs { + for j := range msgs[i].ToolCalls { + tc := &msgs[i].ToolCalls[j] + resolved, err := r.resolve(ctx, tc.SubagentSessionID) + if err != nil { + return err + } + tc.SubagentSessionID = resolved + for k := range tc.ResultEvents { + ev := &tc.ResultEvents[k] + resolvedEv, err := r.resolve(ctx, ev.SubagentSessionID) + if err != nil { + return err + } + ev.SubagentSessionID = resolvedEv + } + } + } + return nil +} + +// sessionNeedsSubagentRewrite reports whether any subagent link in the local +// session resolves to a different PG id than its stored local value. The push +// fast path compares local and PG fingerprints built from the unprefixed local +// ids, so a row already in PG with a stale unprefixed subagent id would match +// and skip the rewrite; this check forces message replacement in that case. +func (r relationshipResolver) sessionNeedsSubagentRewrite( + ctx context.Context, localSessionID string, +) (bool, error) { + ids, err := r.sync.local.SessionSubagentSessionIDs(localSessionID) + if err != nil { + return false, fmt.Errorf( + "reading subagent session ids for %s: %w", localSessionID, err, + ) + } + for _, id := range ids { + resolved, err := r.resolve(ctx, id) + if err != nil { + return false, err + } + if resolved != id { + return true, nil + } } - return b.String() + return false, nil } -// pushedSessionMachine resolves the machine field for a PG row. Old rows -// pushed before this fix with machine="local" will be repaired gradually as -// each session is modified (message count change, etc.) and re-fingerprinted. -func pushedSessionMachine(sess db.Session, fallbackMachine string) string { - if sess.Machine != "" && sess.Machine != "local" { - return sess.Machine +// resolvedSubagentLinkFingerprint fingerprints the PG ids the session's +// subagent links resolve to. The dependency fingerprint is built from the +// local rows, which keep their unprefixed ids, so a collision alias that +// appears after the session was last pushed would leave the stored fingerprint +// unchanged and the fast path would skip the session with its PG tool-call +// links stale; folding the resolved ids in forces that re-push. Sessions +// without subagent links contribute an empty string so their fingerprints are +// unaffected. +func (r relationshipResolver) resolvedSubagentLinkFingerprint( + ctx context.Context, localSessionID string, +) (string, error) { + ids, err := r.sync.local.SessionSubagentSessionIDs(localSessionID) + if err != nil { + return "", fmt.Errorf( + "reading subagent session ids for %s: %w", localSessionID, err, + ) } - return fallbackMachine + var b strings.Builder + for _, id := range ids { + resolved, err := r.resolve(ctx, id) + if err != nil { + return "", err + } + b.WriteString("\x1f") + b.WriteString(resolved) + } + return b.String(), nil } func sameSessionOwner( @@ -1649,6 +2538,29 @@ func sameSessionOwner( return existingMachine == pushedMachine } +func samePushedSessionOwner( + existingOwnerMarker, existingMachine string, + identity pushedSessionIdentity, + markerID string, + legacyMarkerMachines []string, +) bool { + if identity.OwnerMarker == "" { + return sameSessionOwner( + existingOwnerMarker, existingMachine, markerID, + identity.Machine, legacyMarkerMachines, + ) + } + if existingOwnerMarker != "" { + return existingOwnerMarker == identity.OwnerMarker || + slices.Contains(identity.LegacyOwnerMarkers, existingOwnerMarker) + } + // Ownerless artifact rows are legacy-compatible only when PG's marker + // history proves this pusher previously wrote under that machine name. + // Matching the artifact origin alone would let an importer adopt a row it + // did not create. + return slices.Contains(legacyMarkerMachines, existingMachine) +} + func stringValue(value *string) string { if value == nil { return "" @@ -1709,34 +2621,48 @@ func nilStrTS(s *string) any { // backend-neutral transcript_revision column so PG readers can observe // transcript content changes without depending on local sync metadata. func (s *Sync) pushSession( - ctx context.Context, tx *sql.Tx, sess db.Session, markerID string, + ctx context.Context, tx *sql.Tx, sess db.Session, + identity pushedSessionIdentity, markerID string, legacyMarkerMachines []string, ) error { + if identity.LegacyDuplicateID != "" { + if err := verifyArtifactLegacyDuplicateOwnership( + ctx, tx, identity, markerID, + ); err != nil { + return err + } + } createdAt, _ := ParseSQLiteTimestamp(sess.CreatedAt) isAutomated := sess.IsAutomated - pushedMachine := pushedSessionMachine(sess, s.machine) + ownerMarker := identity.effectiveOwnerMarker(markerID) var existingMachine sql.NullString var existingOwnerMarker sql.NullString checkErr := tx.QueryRowContext(ctx, - `SELECT machine, owner_marker FROM sessions WHERE id = $1`, sess.ID, + `SELECT machine, owner_marker FROM sessions WHERE id = $1`, + identity.ID, ).Scan(&existingMachine, &existingOwnerMarker) if checkErr != nil && !errors.Is(checkErr, sql.ErrNoRows) { - return fmt.Errorf("checking session ownership %s: %w", sess.ID, checkErr) + return fmt.Errorf( + "checking session ownership %s: %w", identity.ID, checkErr, + ) } - if checkErr == nil && !sameSessionOwner( + if checkErr == nil && !samePushedSessionOwner( existingOwnerMarker.String, existingMachine.String, + identity, markerID, - pushedMachine, legacyMarkerMachines, ) { log.Printf( "pgsync: session %s: skipping — already owned by machine %q, "+ "this pusher is %q; sync from the origin machine to update", - sess.ID, existingMachine.String, pushedMachine, + identity.ID, existingMachine.String, identity.Machine, ) return errSessionOwnershipConflict } + if checkErr == nil && identity.ArtifactReplica { + return errArtifactReplicaExists + } if legacyMarkerMachines == nil { legacyMarkerMachines = []string{} } @@ -1744,6 +2670,10 @@ func (s *Sync) pushSession( if err != nil { return fmt.Errorf("encoding legacy marker machines: %w", err) } + legacyOwnerMarkersJSON, err := json.Marshal(identity.LegacyOwnerMarkers) + if err != nil { + return fmt.Errorf("encoding legacy owner markers: %w", err) + } result, err := tx.ExecContext(ctx, ` INSERT INTO sessions ( id, machine, owner_marker, project, agent, @@ -1881,7 +2811,10 @@ func (s *Sync) pushSession( SELECT jsonb_array_elements_text($63::jsonb) )) ) - OR sessions.owner_marker = EXCLUDED.owner_marker) + OR sessions.owner_marker = EXCLUDED.owner_marker + OR sessions.owner_marker IN ( + SELECT jsonb_array_elements_text($64::jsonb) + )) AND NOT EXISTS ( SELECT 1 FROM excluded_sessions WHERE id = EXCLUDED.id @@ -1946,7 +2879,7 @@ func (s *Sync) pushSession( OR sessions.duplicate_prompt_count IS DISTINCT FROM EXCLUDED.duplicate_prompt_count OR sessions.no_code_context_count IS DISTINCT FROM EXCLUDED.no_code_context_count OR sessions.runaway_tool_loop_count IS DISTINCT FROM EXCLUDED.runaway_tool_loop_count)`, - sess.ID, pushedMachine, markerID, + identity.ID, identity.Machine, ownerMarker, sanitizePG(sess.Project), sess.Agent, nilStr(sess.FirstMessage), @@ -1990,12 +2923,13 @@ func (s *Sync) pushSession( sanitizePG(sess.AgentLabel), sanitizePG(sess.Entrypoint), string(legacyMarkerMachinesJSON), + string(legacyOwnerMarkersJSON), ) if err != nil { return err } if rowsAffected, rowsErr := result.RowsAffected(); rowsErr == nil && rowsAffected == 0 { - excluded, excludedErr := deletePGSessionIfExcluded(ctx, tx, sess) + excluded, excludedErr := deletePGSessionIfExcluded(ctx, tx, sess, identity.ID) if excludedErr != nil { return excludedErr } @@ -2003,7 +2937,8 @@ func (s *Sync) pushSession( return errSessionExcluded } refreshErr := tx.QueryRowContext(ctx, - `SELECT machine, owner_marker FROM sessions WHERE id = $1`, sess.ID, + `SELECT machine, owner_marker FROM sessions WHERE id = $1`, + identity.ID, ).Scan(&existingMachine, &existingOwnerMarker) if refreshErr != nil { // The guarded upsert changed no rows and we cannot @@ -2017,30 +2952,234 @@ func (s *Sync) pushSession( sess.ID, refreshErr, ) } - if !sameSessionOwner( + if !samePushedSessionOwner( existingOwnerMarker.String, existingMachine.String, - markerID, pushedMachine, legacyMarkerMachines, + identity, markerID, legacyMarkerMachines, ) { log.Printf( "pgsync: session %s: skipping — already owned by machine %q, this pusher is %q; sync from the origin machine to update", - sess.ID, existingMachine.String, pushedMachine, + identity.ID, existingMachine.String, identity.Machine, ) return errSessionOwnershipConflict } } - excluded, excludedErr := deletePGSessionIfExcluded(ctx, tx, sess) + excluded, excludedErr := deletePGSessionIfExcluded(ctx, tx, sess, identity.ID) if excludedErr != nil { return excludedErr } if excluded { return errSessionExcluded } - if err := replacePGSessionAliases(ctx, tx, sess); err != nil { + if identity.LegacyDuplicateID != "" { + if err := consolidateArtifactLegacyDuplicate( + ctx, tx, identity, + ); err != nil { + return err + } + } + aliasSession := sess + aliasSession.ID = identity.ID + if err := replacePGSessionAliases(ctx, tx, aliasSession); err != nil { return err } return nil } +type pgSessionOwnership struct { + machine string + ownerMarker string +} + +// verifyArtifactLegacyDuplicateOwnership locks both sides of an artifact +// upgrade merge before pushSession changes either row. The canonical row must +// have the stable artifact owner and machine, while the bare row must still be +// owned by this pusher's exact pre-artifact random marker. +func verifyArtifactLegacyDuplicateOwnership( + ctx context.Context, + tx *sql.Tx, + identity pushedSessionIdentity, + markerID string, +) error { + if identity.ID == "" || identity.LegacyDuplicateID == "" || + identity.ID == identity.LegacyDuplicateID || + identity.OwnerMarker == "" || markerID == "" { + return fmt.Errorf( + "%w: invalid artifact duplicate consolidation identity", + errSessionOwnershipConflict, + ) + } + + rows, err := tx.QueryContext(ctx, ` + SELECT id, machine, owner_marker + FROM sessions + WHERE id IN ($1, $2) + ORDER BY id + FOR UPDATE + `, identity.ID, identity.LegacyDuplicateID) + if err != nil { + return fmt.Errorf( + "locking artifact duplicate sessions: %w", err, + ) + } + owners := make(map[string]pgSessionOwnership, 2) + for rows.Next() { + var id string + var owner pgSessionOwnership + if err := rows.Scan(&id, &owner.machine, &owner.ownerMarker); err != nil { + _ = rows.Close() + return fmt.Errorf( + "reading artifact duplicate ownership: %w", err, + ) + } + owners[id] = owner + } + if err := rows.Close(); err != nil { + return fmt.Errorf( + "closing artifact duplicate ownership rows: %w", err, + ) + } + if err := rows.Err(); err != nil { + return fmt.Errorf( + "reading artifact duplicate ownership rows: %w", err, + ) + } + + canonical, canonicalOK := owners[identity.ID] + legacy, legacyOK := owners[identity.LegacyDuplicateID] + if !canonicalOK || !legacyOK || + canonical.machine != identity.Machine || + canonical.ownerMarker != identity.OwnerMarker || + legacy.ownerMarker != markerID { + return fmt.Errorf( + "%w: artifact duplicate ownership changed before consolidation", + errSessionOwnershipConflict, + ) + } + return nil +} + +// consolidateArtifactLegacyDuplicate transfers PG-local state to the stable +// canonical artifact row, rewrites references that would otherwise dangle, +// and removes the proven legacy duplicate. The caller holds row locks from +// verifyArtifactLegacyDuplicateOwnership for the duration of this transaction. +func consolidateArtifactLegacyDuplicate( + ctx context.Context, + tx *sql.Tx, + identity pushedSessionIdentity, +) error { + result, err := tx.ExecContext(ctx, ` + UPDATE sessions AS canonical + SET display_name = CASE + WHEN legacy.display_name IS DISTINCT FROM + legacy.source_display_name + THEN legacy.display_name + ELSE canonical.display_name + END, + source_display_name = CASE + WHEN legacy.display_name IS DISTINCT FROM + legacy.source_display_name + THEN legacy.source_display_name + ELSE canonical.source_display_name + END, + deleted_at = CASE + WHEN legacy.deleted_at IS DISTINCT FROM + legacy.source_deleted_at + THEN legacy.deleted_at + ELSE canonical.deleted_at + END, + source_deleted_at = CASE + WHEN legacy.deleted_at IS DISTINCT FROM + legacy.source_deleted_at + THEN legacy.source_deleted_at + ELSE canonical.source_deleted_at + END + FROM sessions AS legacy + WHERE canonical.id = $1 AND legacy.id = $2 + `, identity.ID, identity.LegacyDuplicateID) + if err != nil { + return fmt.Errorf("copying artifact duplicate session curation: %w", err) + } + if rowsAffected, rowsErr := result.RowsAffected(); rowsErr != nil { + return fmt.Errorf("counting artifact duplicate curation rows: %w", rowsErr) + } else if rowsAffected != 1 { + return fmt.Errorf( + "copying artifact duplicate session curation affected %d rows", + rowsAffected, + ) + } + + if _, err := tx.ExecContext(ctx, ` + INSERT INTO starred_sessions (session_id, created_at) + SELECT $1, created_at + FROM starred_sessions + WHERE session_id = $2 + ON CONFLICT (session_id) DO UPDATE SET + created_at = EXCLUDED.created_at + `, identity.ID, identity.LegacyDuplicateID); err != nil { + return fmt.Errorf("copying artifact duplicate star: %w", err) + } + if _, err := tx.ExecContext(ctx, ` + INSERT INTO pinned_messages ( + session_id, message_id, ordinal, source_uuid, note, created_at + ) + SELECT $1, message_id, ordinal, source_uuid, note, created_at + FROM pinned_messages + WHERE session_id = $2 + ON CONFLICT (session_id, message_id) DO UPDATE SET + ordinal = EXCLUDED.ordinal, + source_uuid = EXCLUDED.source_uuid, + note = EXCLUDED.note, + created_at = EXCLUDED.created_at + `, identity.ID, identity.LegacyDuplicateID); err != nil { + return fmt.Errorf("copying artifact duplicate pins: %w", err) + } + + if _, err := tx.ExecContext(ctx, ` + UPDATE sessions + SET parent_session_id = $1 + WHERE parent_session_id = $2 + `, identity.ID, identity.LegacyDuplicateID); err != nil { + return fmt.Errorf("rewriting artifact duplicate parent links: %w", err) + } + if _, err := tx.ExecContext(ctx, ` + UPDATE sessions + SET source_session_id = $1 + WHERE source_session_id = $2 + `, identity.ID, identity.LegacyDuplicateID); err != nil { + return fmt.Errorf("rewriting artifact duplicate source links: %w", err) + } + if _, err := tx.ExecContext(ctx, ` + UPDATE tool_calls + SET subagent_session_id = $1 + WHERE subagent_session_id = $2 + `, identity.ID, identity.LegacyDuplicateID); err != nil { + return fmt.Errorf("rewriting artifact duplicate tool-call links: %w", err) + } + if _, err := tx.ExecContext(ctx, ` + UPDATE tool_result_events + SET subagent_session_id = $1 + WHERE subagent_session_id = $2 + `, identity.ID, identity.LegacyDuplicateID); err != nil { + return fmt.Errorf("rewriting artifact duplicate tool-result links: %w", err) + } + + result, err = tx.ExecContext(ctx, ` + DELETE FROM sessions WHERE id = $1 + `, identity.LegacyDuplicateID) + if err != nil { + return fmt.Errorf("deleting artifact legacy duplicate: %w", err) + } + if rowsAffected, rowsErr := result.RowsAffected(); rowsErr != nil { + return fmt.Errorf("counting deleted artifact duplicate rows: %w", rowsErr) + } else if rowsAffected != 1 { + return fmt.Errorf( + "deleting artifact legacy duplicate affected %d rows", + rowsAffected, + ) + } + return nil +} + // pushMessages replaces a session's messages and tool calls // in PG. It skips the replacement when the PG message count // already matches the local count, avoiding redundant work @@ -2048,12 +3187,14 @@ func (s *Sync) pushSession( func (s *Sync) pushMessages( ctx context.Context, tx *sql.Tx, - sessionID string, + localSessionID string, + pgSessionID string, + resolver relationshipResolver, full bool, sessionUsageFingerprints map[string]string, comparisons *pushMessageComparison, ) (int, error) { - localCount, err := s.local.MessageCount(sessionID) + localCount, err := s.local.MessageCount(localSessionID) if err != nil { return 0, fmt.Errorf( "counting local messages: %w", err, @@ -2062,7 +3203,7 @@ func (s *Sync) pushMessages( if localCount == 0 { if _, err := tx.ExecContext(ctx, `DELETE FROM tool_result_events WHERE session_id = $1`, - sessionID, + pgSessionID, ); err != nil { return 0, fmt.Errorf( "deleting stale pg tool_result_events: %w", err, @@ -2070,7 +3211,7 @@ func (s *Sync) pushMessages( } if _, err := tx.ExecContext(ctx, `DELETE FROM tool_calls WHERE session_id = $1`, - sessionID, + pgSessionID, ); err != nil { return 0, fmt.Errorf( "deleting stale pg tool_calls: %w", err, @@ -2078,7 +3219,7 @@ func (s *Sync) pushMessages( } if _, err := tx.ExecContext(ctx, `DELETE FROM messages WHERE session_id = $1`, - sessionID, + pgSessionID, ); err != nil { return 0, fmt.Errorf( "deleting stale pg messages: %w", err, @@ -2089,11 +3230,13 @@ func (s *Sync) pushMessages( // state.db-only session) with zero messages. Sync them here // too so their cost reaches PG instead of being dropped with // the rest of the message-replace path below. - if err := s.replaceUsageEvents(ctx, tx, sessionID); err != nil { + if err := s.replaceUsageEvents( + ctx, tx, localSessionID, pgSessionID, + ); err != nil { return 0, err } if err := reconcilePinnedMessages( - ctx, tx, sessionID, + ctx, tx, pgSessionID, ); err != nil { return 0, err } @@ -2101,7 +3244,7 @@ func (s *Sync) pushMessages( } pgAgg, pgToolAgg, hasPreloadedComparisons := comparisonAggregates( - sessionID, comparisons, + pgSessionID, comparisons, ) if !hasPreloadedComparisons { if err := tx.QueryRowContext(ctx, @@ -2116,7 +3259,7 @@ func (s *Sync) pushMessages( ) FROM messages WHERE session_id = $1`, - sessionID, + pgSessionID, ).Scan( &pgAgg.Count, &pgAgg.Sum, &pgAgg.Max, &pgAgg.Min, @@ -2131,7 +3274,7 @@ func (s *Sync) pushMessages( COALESCE(SUM(result_content_length), 0) FROM tool_calls WHERE session_id = $1`, - sessionID, + pgSessionID, ).Scan(&pgToolAgg.Count, &pgToolAgg.Sum); err != nil { return 0, fmt.Errorf( "counting pg tool_calls: %w", err, @@ -2140,182 +3283,198 @@ func (s *Sync) pushMessages( } if !full && pgAgg.Count == localCount && pgAgg.Count > 0 { - localFP := pushLocalMessageFingerprint{} - - localFP.Sum, localFP.Max, localFP.Min, err = s.local.MessageContentFingerprint( - sessionID, - ) - if err != nil { - return 0, fmt.Errorf( - "computing local content fingerprint: %w", - err, - ) - } - localFP.ContentHashFP, err = s.local.MessageContentHashFingerprint( - sessionID, - ) - if err != nil { - return 0, fmt.Errorf( - "computing local content hash fingerprint: %w", - err, - ) - } - localFP.RoleTimeFP, err = localMessageRoleTimePGFingerprint( - s.local, sessionID, - ) - if err != nil { - return 0, fmt.Errorf( - "computing local role/time fingerprint: %w", - err, - ) - } - localFP.FlagsFP, err = s.local.MessageFlagsFingerprint(sessionID) - if err != nil { - return 0, fmt.Errorf( - "computing local message flags fingerprint: %w", - err, - ) - } - localFP.SystemFP, err = s.local.SystemMessageFingerprint(sessionID) - if err != nil { - return 0, fmt.Errorf( - "computing local system message fingerprint: %w", err, - ) - } - localFP.ToolCallCount, err = s.local.ToolCallCount(sessionID) - if err != nil { - return 0, fmt.Errorf( - "counting local tool_calls: %w", err, - ) - } - localFP.ToolCallSum, err = s.local.ToolCallContentFingerprint( - sessionID, - ) - if err != nil { - return 0, fmt.Errorf( - "computing local tool_call content fingerprint: %w", - err, - ) - } - localFP.ToolCallFP, err = s.local.ToolCallFingerprint(sessionID) - if err != nil { - return 0, fmt.Errorf( - "computing local tool_call fingerprint: %w", err, - ) - } - localFP.ToolResultFP, err = localToolResultEventPGFingerprint( - s.local, sessionID, + // A row already in PG with a stale unprefixed subagent id matches the + // local fingerprint (both unprefixed), so the skip below would never + // repair it. Force replacement when any subagent link now resolves to a + // different pushed id. + subagentRewrite, err := resolver.sessionNeedsSubagentRewrite( + ctx, localSessionID, ) if err != nil { - return 0, fmt.Errorf( - "computing local tool_result_event fingerprint: %w", err, - ) - } - localFP.TokenFP, err = s.local.MessageTokenFingerprint(sessionID) - if err != nil { - return 0, fmt.Errorf( - "computing local token fingerprint: %w", - err, - ) + return 0, err } + if !subagentRewrite { + localFP := pushLocalMessageFingerprint{} - usageFromMap := false - if sessionUsageFingerprints != nil { - var ok bool - localFP.UsageEventFP, ok = sessionUsageFingerprints[sessionID] - usageFromMap = ok - } - if !usageFromMap { - localFP.UsageEventFP, err = s.local.UsageEventFingerprint(sessionID) + localFP.Sum, localFP.Max, localFP.Min, err = s.local.MessageContentFingerprint( + localSessionID, + ) if err != nil { return 0, fmt.Errorf( - "computing local usage event fingerprint: %w", + "computing local content fingerprint: %w", err, ) } - } - - if comparisons == nil { - pgContentHashFP, err := pgMessageContentHashFingerprint( - ctx, tx, sessionID, + localFP.ContentHashFP, err = s.local.MessageContentHashFingerprint( + localSessionID, ) if err != nil { return 0, fmt.Errorf( - "computing pg content hash fingerprint: %w", + "computing local content hash fingerprint: %w", err, ) } - pgRoleTimeFP, err := pgMessageRoleTimeFingerprint( - ctx, tx, sessionID, + localFP.RoleTimeFP, err = localMessageRoleTimePGFingerprint( + s.local, localSessionID, ) if err != nil { return 0, fmt.Errorf( - "computing pg role/time fingerprint: %w", + "computing local role/time fingerprint: %w", err, ) } - pgFlagsFP, err := pgMessageFlagsFingerprint(ctx, tx, sessionID) + localFP.FlagsFP, err = s.local.MessageFlagsFingerprint(localSessionID) if err != nil { return 0, fmt.Errorf( - "computing pg message flags fingerprint: %w", + "computing local message flags fingerprint: %w", err, ) } - pgTokenFP, err := pgMessageTokenFingerprint(ctx, tx, sessionID) + localFP.SystemFP, err = s.local.SystemMessageFingerprint(localSessionID) if err != nil { return 0, fmt.Errorf( - "computing pg token fingerprint: %w", - err, + "computing local system message fingerprint: %w", err, ) } - pgTCFP, err := pgToolCallFingerprint(ctx, tx, sessionID) + localFP.ToolCallCount, err = s.local.ToolCallCount(localSessionID) if err != nil { return 0, fmt.Errorf( - "computing pg tool_call fingerprint: %w", - err, + "counting local tool_calls: %w", err, ) } - pgResultFP, err := pgToolResultEventFingerprint(ctx, tx, sessionID) + localFP.ToolCallSum, err = s.local.ToolCallContentFingerprint( + localSessionID, + ) if err != nil { return 0, fmt.Errorf( - "computing pg tool_result_event fingerprint: %w", + "computing local tool_call content fingerprint: %w", err, ) } - pgUsageFP, err := pgUsageEventFingerprint(ctx, tx, sessionID) + localFP.ToolCallFP, err = s.local.ToolCallFingerprint(localSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing local tool_call fingerprint: %w", err, + ) + } + localFP.ToolResultFP, err = localToolResultEventPGFingerprint( + s.local, localSessionID, + ) + if err != nil { + return 0, fmt.Errorf( + "computing local tool_result_event fingerprint: %w", err, + ) + } + localFP.TokenFP, err = s.local.MessageTokenFingerprint(localSessionID) if err != nil { return 0, fmt.Errorf( - "computing pg usage event fingerprint: %w", + "computing local token fingerprint: %w", err, ) } - if localFP.Sum == pgAgg.Sum && - localFP.Max == pgAgg.Max && - localFP.Min == pgAgg.Min && - localFP.ContentHashFP == pgContentHashFP && - localFP.RoleTimeFP == pgRoleTimeFP && - localFP.FlagsFP == pgFlagsFP && - localFP.SystemFP == pgAgg.SysFP && - localFP.ToolCallCount == pgToolAgg.Count && - localFP.ToolCallSum == pgToolAgg.Sum && - localFP.ToolCallFP == pgTCFP && - localFP.ToolResultFP == pgResultFP && - localFP.TokenFP == pgTokenFP && - localFP.UsageEventFP == pgUsageFP { + usageFromMap := false + if sessionUsageFingerprints != nil { + var ok bool + localFP.UsageEventFP, ok = sessionUsageFingerprints[localSessionID] + usageFromMap = ok + } + if !usageFromMap { + localFP.UsageEventFP, err = s.local.UsageEventFingerprint(localSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing local usage event fingerprint: %w", + err, + ) + } + } + + if comparisons == nil { + pgContentHashFP, err := pgMessageContentHashFingerprint( + ctx, tx, pgSessionID, + ) + if err != nil { + return 0, fmt.Errorf( + "computing pg content hash fingerprint: %w", + err, + ) + } + pgRoleTimeFP, err := pgMessageRoleTimeFingerprint( + ctx, tx, pgSessionID, + ) + if err != nil { + return 0, fmt.Errorf( + "computing pg role/time fingerprint: %w", + err, + ) + } + pgFlagsFP, err := pgMessageFlagsFingerprint(ctx, tx, pgSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing pg message flags fingerprint: %w", + err, + ) + } + pgTokenFP, err := pgMessageTokenFingerprint(ctx, tx, pgSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing pg token fingerprint: %w", + err, + ) + } + pgTCFP, err := pgToolCallFingerprint(ctx, tx, pgSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing pg tool_call fingerprint: %w", + err, + ) + } + pgResultFP, err := pgToolResultEventFingerprint(ctx, tx, pgSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing pg tool_result_event fingerprint: %w", + err, + ) + } + pgUsageFP, err := pgUsageEventFingerprint(ctx, tx, pgSessionID) + if err != nil { + return 0, fmt.Errorf( + "computing pg usage event fingerprint: %w", + err, + ) + } + + if localFP.Sum == pgAgg.Sum && + localFP.Max == pgAgg.Max && + localFP.Min == pgAgg.Min && + localFP.ContentHashFP == pgContentHashFP && + localFP.RoleTimeFP == pgRoleTimeFP && + localFP.FlagsFP == pgFlagsFP && + localFP.SystemFP == pgAgg.SysFP && + localFP.ToolCallCount == pgToolAgg.Count && + localFP.ToolCallSum == pgToolAgg.Sum && + localFP.ToolCallFP == pgTCFP && + localFP.ToolResultFP == pgResultFP && + localFP.TokenFP == pgTokenFP && + localFP.UsageEventFP == pgUsageFP { + return 0, nil + } + } else if shouldSkipSessionMessages( + pgSessionID, localCount, localFP, full, comparisons, + ) { return 0, nil } - } else if shouldSkipSessionMessages( - sessionID, localCount, localFP, full, comparisons, - ) { - return 0, nil } } + if err := backfillLegacyPinnedMessageSourceUUIDs(ctx, tx, pgSessionID); err != nil { + return 0, err + } + if _, err := tx.ExecContext(ctx, ` DELETE FROM tool_result_events WHERE session_id = $1 - `, sessionID); err != nil { + `, pgSessionID); err != nil { return 0, fmt.Errorf( "deleting pg tool_result_events: %w", err, ) @@ -2323,7 +3482,7 @@ func (s *Sync) pushMessages( if _, err := tx.ExecContext(ctx, ` DELETE FROM tool_calls WHERE session_id = $1 - `, sessionID); err != nil { + `, pgSessionID); err != nil { return 0, fmt.Errorf( "deleting pg tool_calls: %w", err, ) @@ -2331,12 +3490,14 @@ func (s *Sync) pushMessages( if _, err := tx.ExecContext(ctx, ` DELETE FROM messages WHERE session_id = $1 - `, sessionID); err != nil { + `, pgSessionID); err != nil { return 0, fmt.Errorf( "deleting pg messages: %w", err, ) } - if err := s.replaceUsageEvents(ctx, tx, sessionID); err != nil { + if err := s.replaceUsageEvents( + ctx, tx, localSessionID, pgSessionID, + ); err != nil { return 0, err } @@ -2344,7 +3505,7 @@ func (s *Sync) pushMessages( startOrdinal := 0 for { msgs, err := s.local.GetMessages( - ctx, sessionID, startOrdinal, + ctx, localSessionID, startOrdinal, db.MaxMessageLimit, true, ) if err != nil { @@ -2361,23 +3522,29 @@ func (s *Sync) pushMessages( return count, fmt.Errorf( "pushMessages %s: ordinal did not "+ "advance (start=%d, last=%d)", - sessionID, startOrdinal, + localSessionID, startOrdinal, msgs[len(msgs)-1].Ordinal, ) } + if err := resolver.rewriteSubagentIDs(ctx, msgs); err != nil { + return count, fmt.Errorf( + "resolving subagent session ids: %w", err, + ) + } + if err := bulkInsertMessages( - ctx, tx, sessionID, msgs, + ctx, tx, pgSessionID, msgs, ); err != nil { return count, err } if err := bulkInsertToolCalls( - ctx, tx, sessionID, msgs, + ctx, tx, pgSessionID, msgs, ); err != nil { return count, err } if err := bulkInsertToolResultEvents( - ctx, tx, sessionID, msgs, + ctx, tx, pgSessionID, msgs, ); err != nil { return count, err } @@ -2385,7 +3552,7 @@ func (s *Sync) pushMessages( startOrdinal = nextOrdinal } - if err := reconcilePinnedMessages(ctx, tx, sessionID); err != nil { + if err := reconcilePinnedMessages(ctx, tx, pgSessionID); err != nil { return count, err } @@ -2399,19 +3566,22 @@ func (s *Sync) pushMessages( // zero-message and the normal message-replace paths in pushMessages call // this so a session's cost always reaches PG. func (s *Sync) replaceUsageEvents( - ctx context.Context, tx *sql.Tx, sessionID string, + ctx context.Context, tx *sql.Tx, + localSessionID string, pgSessionID string, ) error { if _, err := tx.ExecContext(ctx, ` DELETE FROM usage_events WHERE session_id = $1 - `, sessionID); err != nil { + `, pgSessionID); err != nil { return fmt.Errorf("deleting pg usage_events: %w", err) } - usageEvents, err := s.local.GetUsageEvents(ctx, sessionID) + usageEvents, err := s.local.GetUsageEvents(ctx, localSessionID) if err != nil { return fmt.Errorf("reading local usage events: %w", err) } - if err := bulkInsertUsageEvents(ctx, tx, usageEvents); err != nil { + if err := bulkInsertUsageEvents( + ctx, tx, pgSessionID, usageEvents, + ); err != nil { return err } return nil @@ -2420,20 +3590,8 @@ func (s *Sync) replaceUsageEvents( func reconcilePinnedMessages( ctx context.Context, tx *sql.Tx, sessionID string, ) error { - if _, err := tx.ExecContext(ctx, ` - UPDATE pinned_messages p - SET source_uuid = m.source_uuid - FROM messages m - WHERE p.session_id = $1 - AND m.session_id = p.session_id - AND m.ordinal = p.message_id - AND p.source_uuid = '' - AND m.source_uuid <> ''`, - sessionID, - ); err != nil { - return fmt.Errorf( - "backfilling pg pin source_uuid: %w", err, - ) + if err := backfillLegacyPinnedMessageSourceUUIDs(ctx, tx, sessionID); err != nil { + return err } // Move shifted source-backed pins out of the real ordinal range @@ -2590,6 +3748,27 @@ func reconcilePinnedMessages( return nil } +func backfillLegacyPinnedMessageSourceUUIDs( + ctx context.Context, tx *sql.Tx, sessionID string, +) error { + if _, err := tx.ExecContext(ctx, ` + UPDATE pinned_messages p + SET source_uuid = m.source_uuid + FROM messages m + WHERE p.session_id = $1 + AND m.session_id = p.session_id + AND m.ordinal = p.message_id + AND p.source_uuid = '' + AND m.source_uuid <> ''`, + sessionID, + ); err != nil { + return fmt.Errorf( + "backfilling pg pin source_uuid: %w", err, + ) + } + return nil +} + func pgMessageTokenFingerprint( ctx context.Context, tx *sql.Tx, sessionID string, ) (string, error) { @@ -2958,7 +4137,8 @@ func bulkInsertMessages( } func bulkInsertUsageEvents( - ctx context.Context, tx *sql.Tx, events []db.UsageEvent, + ctx context.Context, tx *sql.Tx, + sessionID string, events []db.UsageEvent, ) error { if len(events) == 0 { return nil @@ -3001,7 +4181,7 @@ func bulkInsertUsageEvents( cost = *ev.CostUSD } args = append(args, - ev.SessionID, + sessionID, ordinal, sanitizePG(ev.Source), sanitizePG(ev.Model), @@ -3237,31 +4417,32 @@ func bulkInsertToolResultEvents( // sessions.secrets_rules_version is pushed by pushSession alongside // the rest of the session columns. func (s *Sync) pushSecretFindings( - ctx context.Context, tx *sql.Tx, sessionID string, + ctx context.Context, tx *sql.Tx, + localSessionID string, pgSessionID string, ) (bool, error) { res, err := tx.ExecContext(ctx, `DELETE FROM secret_findings WHERE session_id = $1`, - sessionID, + pgSessionID, ) if err != nil { return false, fmt.Errorf( "deleting pg secret_findings for %s: %w", - sessionID, err, + pgSessionID, err, ) } deleted, err := res.RowsAffected() if err != nil { return false, fmt.Errorf( "counting deleted secret_findings for %s: %w", - sessionID, err, + pgSessionID, err, ) } - findings, err := s.local.SessionSecretFindings(ctx, sessionID) + findings, err := s.local.SessionSecretFindings(ctx, localSessionID) if err != nil { return false, fmt.Errorf( "reading local secret_findings for %s: %w", - sessionID, err, + localSessionID, err, ) } if len(findings) == 0 { @@ -3294,7 +4475,7 @@ func (s *Sync) pushSecretFindings( p+5, p+6, p+7, p+8, p+9, p+10, p+11, ) args = append(args, - f.SessionID, f.RuleName, f.Confidence, + pgSessionID, f.RuleName, f.Confidence, f.LocationKind, f.MessageOrdinal, f.CallIndex, f.EventIndex, f.MatchStart, f.MatchEnd, f.MatchIndex, @@ -3306,7 +4487,7 @@ func (s *Sync) pushSecretFindings( ); err != nil { return false, fmt.Errorf( "bulk inserting secret_findings for %s: %w", - sessionID, err, + pgSessionID, err, ) } } diff --git a/internal/postgres/push_pgtest_test.go b/internal/postgres/push_pgtest_test.go index 8c5ae2d28..180491a50 100644 --- a/internal/postgres/push_pgtest_test.go +++ b/internal/postgres/push_pgtest_test.go @@ -638,7 +638,10 @@ func TestPushSessionTerminationStatus(t *testing.T) { t.Helper() tx, err := pg.BeginTx(ctx, nil) require.NoError(t, err, "BeginTx") - if err := sync.pushSession(ctx, tx, s, markerID, nil); err != nil { + if err := sync.pushSession(ctx, tx, s, pushedSessionIdentity{ + ID: s.ID, + Machine: s.Machine, + }, markerID, nil); err != nil { _ = tx.Rollback() t.Fatalf("pushSession: %v", err) } @@ -704,7 +707,12 @@ func TestPushSessionPreservesSourceMachine(t *testing.T) { require.NoError(t, err, "BeginTx") markerID, err := sync.pushMarkerID() require.NoError(t, err, "pushMarkerID") - require.NoError(t, sync.pushSession(ctx, tx, remoteSession, markerID, nil), "pushSession") + require.NoError(t, sync.pushSession( + ctx, tx, remoteSession, pushedSessionIdentity{ + ID: remoteSession.ID, + Machine: remoteSession.Machine, + }, markerID, nil, + ), "pushSession") require.NoError(t, tx.Commit(), "Commit") var got string @@ -1181,7 +1189,7 @@ func TestPushSessionSkipsPGExcludedSession(t *testing.T) { Machine: "test-machine", Agent: "claude", CreatedAt: "2026-01-01T00:00:00Z", - }, markerID, nil) + }, pushedSessionIdentity{ID: sessionID, Machine: "test-machine"}, markerID, nil) require.ErrorIs(t, err, errSessionExcluded) require.NoError(t, tx.Rollback(), "Rollback") @@ -1942,6 +1950,272 @@ func TestPushIncrementalWithOnlyForeignMachineSessions(t *testing.T) { "second incremental push must not rewrite the session") } +func TestPushPreservesMixedLocalAndImportedMachinesForPGStore(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_push_mixed_machine_attribution_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + localDB := testDB(t) + sync := &Sync{ + pg: pg, + local: localDB, + machine: "laptop-a1b2c3", + schema: schema, + schemaDone: true, + } + + localSess := db.Session{ + ID: "local-session-001", + Project: "local-proj", + Machine: "local", + Agent: "claude", + MessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + } + importedSess := db.Session{ + ID: "desktop-d4e5f6~foreign-session-001", + Project: "foreign-proj", + Machine: "desktop-d4e5f6", + Agent: "codex", + MessageCount: 1, + CreatedAt: "2026-01-01T00:01:00Z", + } + for _, sess := range []db.Session{localSess, importedSess} { + require.NoError(t, localDB.UpsertSession(sess), "UpsertSession %s", sess.ID) + require.NoError(t, localDB.InsertMessages([]db.Message{{ + SessionID: sess.ID, + Ordinal: 0, + Role: "assistant", + Content: "hello from " + sess.Machine, + ContentLength: len("hello from " + sess.Machine), + }}), "InsertMessages %s", sess.ID) + } + + res, err := sync.Push(ctx, false, nil) + require.NoError(t, err, "first Push") + assert.Zero(t, res.Errors, "first push should report no failures") + assert.Equal(t, 2, res.SessionsPushed) + + rows, err := pg.Query(`SELECT id, machine FROM sessions`) + require.NoError(t, err, "querying pushed machines") + defer rows.Close() + gotMachines := map[string]string{} + for rows.Next() { + var id, machine string + require.NoError(t, rows.Scan(&id, &machine), "scanning pushed machine") + gotMachines[id] = machine + } + require.NoError(t, rows.Err(), "iterating pushed machines") + assert.Equal(t, "laptop-a1b2c3", gotMachines[localSess.ID]) + assert.Equal(t, "desktop-d4e5f6", gotMachines[importedSess.ID]) + + store, err := NewStore(pgURL, schema, true) + require.NoError(t, err, "NewStore") + defer store.Close() + machines, err := store.GetMachines(ctx, false, false) + require.NoError(t, err, "GetMachines") + assert.ElementsMatch(t, []string{"desktop-d4e5f6", "laptop-a1b2c3"}, machines) + + importedPage, err := store.ListSessions(ctx, db.SessionFilter{ + Machine: "desktop-d4e5f6", + Limit: 10, + }) + require.NoError(t, err, "ListSessions imported machine") + require.Equal(t, 1, importedPage.Total) + require.Len(t, importedPage.Sessions, 1) + assert.Equal(t, importedSess.ID, importedPage.Sessions[0].ID) + assert.Equal(t, "desktop-d4e5f6", importedPage.Sessions[0].Machine) + + localPage, err := store.ListSessions(ctx, db.SessionFilter{ + Machine: "laptop-a1b2c3", + Limit: 10, + }) + require.NoError(t, err, "ListSessions local machine") + require.Equal(t, 1, localPage.Total) + require.Len(t, localPage.Sessions, 1) + assert.Equal(t, localSess.ID, localPage.Sessions[0].ID) + assert.Equal(t, "laptop-a1b2c3", localPage.Sessions[0].Machine) + + res, err = sync.Push(ctx, false, nil) + require.NoError(t, err, "second Push") + assert.Zero(t, res.Errors, "second push should report no failures") + assert.Zero(t, res.SessionsPushed, "unchanged mixed-machine sessions should not be re-pushed") +} + +func TestPushKeepsSameNativeIDDistinctAcrossMachines(t *testing.T) { + pgURL := testPGURL(t) + + const schema = "agentsview_push_native_id_collision_test" + pg, err := Open(pgURL, schema, true) + require.NoError(t, err, "Open") + defer pg.Close() + + ctx := context.Background() + _, err = pg.Exec(`DROP SCHEMA IF EXISTS ` + schema + ` CASCADE`) + require.NoError(t, err, "drop schema") + require.NoError(t, EnsureSchema(ctx, pg, schema), "EnsureSchema") + + const nativeID = "shared-native-session-001" + const machineA = "laptop-a1b2c3" + const machineB = "desktop-d4e5f6" + const projectA = "laptop-proj" + const projectB = "desktop-proj" + pgIDA := nativeID + pgIDB := prefixedSessionID(machineB, nativeID) + + localA := testDB(t) + localB := testDB(t) + syncA := &Sync{ + pg: pg, + local: localA, + machine: machineA, + schema: schema, + schemaDone: true, + } + syncB := &Sync{ + pg: pg, + local: localB, + machine: machineB, + schema: schema, + schemaDone: true, + } + + seed := func(local *db.DB, project, content string) db.Session { + t.Helper() + sess := db.Session{ + ID: nativeID, + Project: project, + Machine: "local", + Agent: "claude", + MessageCount: 1, + CreatedAt: "2026-01-01T00:00:00Z", + } + require.NoError(t, local.UpsertSession(sess), "UpsertSession") + require.NoError(t, local.InsertMessages([]db.Message{{ + SessionID: nativeID, + Ordinal: 0, + Role: "assistant", + Content: content, + ContentLength: len(content), + }}), "InsertMessages") + return sess + } + sessA := seed(localA, projectA, "from laptop") + sessB := seed(localB, projectB, "from desktop") + + res, err := syncA.Push(ctx, false, nil) + require.NoError(t, err, "Push A") + assert.Zero(t, res.Errors, "first A push should report no failures") + assert.Equal(t, 1, res.SessionsPushed) + + res, err = syncB.Push(ctx, false, nil) + require.NoError(t, err, "Push B") + assert.Zero(t, res.Errors, "first B push should report no failures") + assert.Equal(t, 1, res.SessionsPushed) + + rows, err := pg.Query( + `SELECT id, machine, project FROM sessions ORDER BY id`, + ) + require.NoError(t, err, "querying pushed collision rows") + defer rows.Close() + got := map[string]db.Session{} + for rows.Next() { + var row db.Session + require.NoError(t, rows.Scan( + &row.ID, &row.Machine, &row.Project, + ), "scanning collision row") + got[row.ID] = row + } + require.NoError(t, rows.Err(), "iterating collision rows") + require.Len(t, got, 2) + assert.Equal(t, machineA, got[pgIDA].Machine) + assert.Equal(t, projectA, got[pgIDA].Project) + assert.Equal(t, machineB, got[pgIDB].Machine) + assert.Equal(t, projectB, got[pgIDB].Project) + + var contentA, contentB string + require.NoError(t, pg.QueryRow( + `SELECT content FROM messages WHERE session_id = $1 AND ordinal = 0`, + pgIDA, + ).Scan(&contentA), "reading A message") + require.NoError(t, pg.QueryRow( + `SELECT content FROM messages WHERE session_id = $1 AND ordinal = 0`, + pgIDB, + ).Scan(&contentB), "reading B message") + assert.Equal(t, "from laptop", contentA) + assert.Equal(t, "from desktop", contentB) + + store, err := NewStore(pgURL, schema, true) + require.NoError(t, err, "NewStore") + defer store.Close() + pageA, err := store.ListSessions(ctx, db.SessionFilter{ + Machine: machineA, + Limit: 10, + }) + require.NoError(t, err, "ListSessions A") + require.Equal(t, 1, pageA.Total) + require.Len(t, pageA.Sessions, 1) + assert.Equal(t, pgIDA, pageA.Sessions[0].ID) + assert.Equal(t, projectA, pageA.Sessions[0].Project) + pageB, err := store.ListSessions(ctx, db.SessionFilter{ + Machine: machineB, + Limit: 10, + }) + require.NoError(t, err, "ListSessions B") + require.Equal(t, 1, pageB.Total) + require.Len(t, pageB.Sessions, 1) + assert.Equal(t, pgIDB, pageB.Sessions[0].ID) + assert.Equal(t, projectB, pageB.Sessions[0].Project) + + res, err = syncA.Push(ctx, false, nil) + require.NoError(t, err, "second Push A") + assert.Zero(t, res.Errors, "second A push should report no failures") + assert.Zero(t, res.SessionsPushed) + res, err = syncB.Push(ctx, false, nil) + require.NoError(t, err, "second Push B") + assert.Zero(t, res.Errors, "second B push should report no failures") + assert.Zero(t, res.SessionsPushed) + + var rowCount int + require.NoError(t, pg.QueryRow( + `SELECT COUNT(*) FROM sessions WHERE id IN ($1, $2)`, + pgIDA, pgIDB, + ).Scan(&rowCount), "counting collision rows") + assert.Equal(t, 2, rowCount) + + updatedB := sessB + updatedB.Project = "desktop-proj-updated" + require.NoError(t, localB.UpsertSession(updatedB), "update B session") + // UpsertSession does not write local_modified_at, so mark the session + // modified explicitly; otherwise the incremental push will not re-list it. + require.NoError(t, localB.BumpLocalModifiedAt(nativeID), + "mark updated B session modified") + res, err = syncB.Push(ctx, false, nil) + require.NoError(t, err, "updated Push B") + assert.Zero(t, res.Errors, "updated B push should report no failures") + assert.Equal(t, 1, res.SessionsPushed) + + var projectAfterA, projectAfterB string + require.NoError(t, pg.QueryRow( + `SELECT project FROM sessions WHERE id = $1`, + pgIDA, + ).Scan(&projectAfterA), "reading A project after B update") + require.NoError(t, pg.QueryRow( + `SELECT project FROM sessions WHERE id = $1`, + pgIDB, + ).Scan(&projectAfterB), "reading B project after update") + assert.Equal(t, sessA.Project, projectAfterA) + assert.Equal(t, updatedB.Project, projectAfterB) +} + // TestPushDetectsResetWhenCompetingMachineRowsExist verifies that a PG reset is // detected even when another pusher has repopulated rows under a machine value // this host also writes. The local session carries Machine "remote-host" (as a @@ -2451,11 +2725,18 @@ func TestPushReportsSkippedConflicts(t *testing.T) { schemaDone: true, } - const sessID = "conflict-001" + // A same-id collision against a different live machine is resolved by + // storing the session under a machine-prefixed id, not skipped. A skipped + // conflict now arises only when the id cannot be re-prefixed: an imported + // foreign-origin session already carries its origin's prefix, so when PG + // holds that id under a different owner marker -- e.g. a third machine + // imported and pushed the same session first -- this push must leave the + // row to its owner and report it as a skipped conflict. + const sessID = "machine-a~conflict-001" require.NoError(t, localDB.UpsertSession(db.Session{ ID: sessID, Project: "proj", - Machine: "machine-b", + Machine: "machine-a", Agent: "claude", MessageCount: 1, CreatedAt: "2026-01-01T00:00:00Z", @@ -2480,4 +2761,18 @@ func TestPushReportsSkippedConflicts(t *testing.T) { assert.Zero(t, res.Errors, "push should not report failed sessions") assert.Zero(t, res.SessionsPushed, "conflicting session should not be counted as pushed") assert.Equal(t, 1, res.SkippedConflicts, "skipped conflicts should be observable in PushResult") + + // The conflicting row is left untouched and no doubly-prefixed row is + // created in this pusher's namespace. + var ownerMarker string + require.NoError(t, pg.QueryRow( + `SELECT owner_marker FROM sessions WHERE id = $1`, sessID, + ).Scan(&ownerMarker), "reading conflicting row owner") + assert.Equal(t, "other-owner", ownerMarker) + var doublyPrefixed int + require.NoError(t, pg.QueryRow( + `SELECT COUNT(*) FROM sessions WHERE id = $1`, + prefixedSessionID("machine-b", sessID), + ).Scan(&doublyPrefixed), "counting doubly-prefixed rows") + assert.Zero(t, doublyPrefixed) } diff --git a/internal/postgres/push_test.go b/internal/postgres/push_test.go index 1d351cf2d..e9347e4b6 100644 --- a/internal/postgres/push_test.go +++ b/internal/postgres/push_test.go @@ -6,8 +6,11 @@ import ( "database/sql/driver" "encoding/json" "errors" + "fmt" "io" "path/filepath" + "regexp" + "strconv" "strings" "sync" "testing" @@ -321,6 +324,23 @@ func TestTranscriptRevisionBackfillForcesOneFullPush(t *testing.T) { assert.False(t, needed) } +func TestArtifactIdentityModePersistsOnlyAfterSuccessfulPush(t *testing.T) { + store := &syncStateStoreStub{values: map[string]string{ + artifactIdentityModeStateKey: legacyArtifactIdentityMode, + }} + mode := artifactOwnerMarkerPrefix + "origin-a1b2c3" + + require.NoError(t, completeArtifactIdentityMode( + store, mode, PushResult{Errors: 1}, + )) + assert.Equal(t, legacyArtifactIdentityMode, + store.values[artifactIdentityModeStateKey], + "a partial push must leave the prior mode so the next run retries the transition") + + require.NoError(t, completeArtifactIdentityMode(store, mode, PushResult{})) + assert.Equal(t, mode, store.values[artifactIdentityModeStateKey]) +} + func TestCompleteSessionAliasBackfillMarksDoneUnlessErrors(t *testing.T) { for _, tc := range []struct { name string @@ -523,6 +543,125 @@ func TestSessionAliasBackfillKeysStayFilteredForPushState(t *testing.T) { assert.Empty(t, store.values["last_push_at:work"]) } +// scriptedSink is a sessionBatchSink that records calls and writes every +// session except those named in fail. It mirrors pushBatch semantics: a +// multi-session batch containing a failing session rolls back as a unit +// (ok=false) so the driver retries each session individually, and a failing +// single-session batch returns ok=false. A session named in fatal makes the +// batch containing it return a fatal error. +type scriptedSink struct { + fail map[string]bool + fatal string + calls [][]string +} + +func (s *scriptedSink) writeBatch( + _ context.Context, batch []db.Session, pushed *[]db.Session, +) (batchResult, error) { + ids := make([]string, len(batch)) + for i, sess := range batch { + ids[i] = sess.ID + } + s.calls = append(s.calls, ids) + for _, sess := range batch { + if sess.ID == s.fatal { + return batchResult{}, errors.New("fatal sink error") + } + if s.fail[sess.ID] { + // Whole batch rolls back without writing. + return batchResult{ok: false}, nil + } + } + msgs := 0 + for _, sess := range batch { + *pushed = append(*pushed, sess) + msgs += 2 + } + return batchResult{ok: true, sessions: len(batch), messages: msgs}, nil +} + +func sessionsWithIDs(ids ...string) []db.Session { + out := make([]db.Session, len(ids)) + for i, id := range ids { + out[i] = db.Session{ID: id} + } + return out +} + +func TestDrainSessionBatchesChunksAndReportsProgress(t *testing.T) { + var ids []string + for i := range 120 { + ids = append(ids, fmt.Sprintf("s%03d", i)) + } + sessions := sessionsWithIDs(ids...) + sink := &scriptedSink{} + + var result PushResult + var progress []PushProgress + pushed, err := drainSessionBatches( + context.Background(), sessions, sink, &result, + func(p PushProgress) { progress = append(progress, p) }, + ) + require.NoError(t, err) + + // Batched in chunks of 50: 50 + 50 + 20. + require.Len(t, sink.calls, 3) + assert.Len(t, sink.calls[0], 50) + assert.Len(t, sink.calls[1], 50) + assert.Len(t, sink.calls[2], 20) + + assert.Equal(t, 120, result.SessionsPushed) + assert.Equal(t, 240, result.MessagesPushed) + assert.Equal(t, 0, result.Errors) + assert.Len(t, pushed, 120) + + require.Len(t, progress, 3) + assert.Equal(t, PushProgress{SessionsDone: 50, SessionsTotal: 120, MessagesDone: 100}, progress[0]) + assert.Equal(t, PushProgress{SessionsDone: 100, SessionsTotal: 120, MessagesDone: 200}, progress[1]) + assert.Equal(t, PushProgress{SessionsDone: 120, SessionsTotal: 120, MessagesDone: 240}, progress[2]) +} + +func TestDrainSessionBatchesRetriesFailedBatchIndividually(t *testing.T) { + sessions := sessionsWithIDs("a", "b", "c") + sink := &scriptedSink{fail: map[string]bool{"b": true}} + + var result PushResult + pushed, err := drainSessionBatches( + context.Background(), sessions, sink, &result, nil, + ) + require.NoError(t, err) + + // First the whole batch (rolls back), then each session individually. + require.Len(t, sink.calls, 4) + assert.Equal(t, []string{"a", "b", "c"}, sink.calls[0]) + assert.Equal(t, []string{"a"}, sink.calls[1]) + assert.Equal(t, []string{"b"}, sink.calls[2]) + assert.Equal(t, []string{"c"}, sink.calls[3]) + + assert.Equal(t, 2, result.SessionsPushed) + assert.Equal(t, 4, result.MessagesPushed) + assert.Equal(t, 1, result.Errors) + pushedIDs := make([]string, len(pushed)) + for i, sess := range pushed { + pushedIDs[i] = sess.ID + } + assert.Equal(t, []string{"a", "c"}, pushedIDs) +} + +func TestDrainSessionBatchesFatalErrorAborts(t *testing.T) { + sessions := sessionsWithIDs("a", "b", "c") + sink := &scriptedSink{fatal: "b"} + + var result PushResult + pushed, err := drainSessionBatches( + context.Background(), sessions, sink, &result, nil, + ) + require.Error(t, err) + assert.Equal(t, 0, result.SessionsPushed) + assert.Equal(t, 0, result.Errors) + assert.Empty(t, pushed) +} + func TestReadPushBoundaryStateValidity(t *testing.T) { const cutoff = "2026-03-11T12:34:56.123Z" @@ -593,6 +732,97 @@ func TestPGExcludedSessionIDsQueryUsesSingleArrayParameter(t *testing.T) { ) } +func TestArtifactPushIdentityCanonicalizesNativeAndImportedCopies(t *testing.T) { + tests := []struct { + name string + session db.Session + localOrigin string + imported bool + wantID string + wantMachine string + wantOwner string + wantOK bool + }{ + { + name: "locally owned artifact session", + session: db.Session{ID: "native-id", Machine: "local"}, + localOrigin: "origin-a1b2c3", + wantID: "origin-a1b2c3~native-id", + wantMachine: "origin-a1b2c3", + wantOwner: "artifact-origin:origin-a1b2c3", + wantOK: true, + }, + { + name: "imported artifact session", + session: db.Session{ID: "origin-a1b2c3~native-id", Machine: "origin-a1b2c3"}, + localOrigin: "origin-b4c5d6", + imported: true, + wantID: "origin-a1b2c3~native-id", + wantMachine: "origin-a1b2c3", + wantOwner: "artifact-origin:origin-a1b2c3", + wantOK: true, + }, + { + name: "ssh-shaped session without artifact provenance stays legacy", + session: db.Session{ID: "origin-a1b2c3~native-id", Machine: "origin-a1b2c3"}, + localOrigin: "origin-b4c5d6", + wantOK: false, + }, + { + name: "no artifact origin preserves legacy collision behavior", + session: db.Session{ID: "native-id", Machine: "local"}, + wantOK: false, + }, + { + name: "prefixed foreign session without local artifact opt in stays legacy", + session: db.Session{ID: "remote-host~native-id", Machine: "remote-host"}, + wantOK: false, + }, + { + name: "unrelated foreign machine id is not treated as imported", + session: db.Session{ID: "native-id", Machine: "remote-host"}, + localOrigin: "origin-b4c5d6", + wantOK: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + id, machine, owner, ok := artifactPushIdentity( + tt.session, tt.localOrigin, tt.imported, + ) + assert.Equal(t, tt.wantOK, ok) + assert.Equal(t, tt.wantID, id) + assert.Equal(t, tt.wantMachine, machine) + assert.Equal(t, tt.wantOwner, owner) + }) + } +} + +func TestArtifactPushOwnershipAdoptsOnlyCurrentPushersLegacyMarker(t *testing.T) { + identity := pushedSessionIdentity{ + Machine: "origin-a1b2c3", + OwnerMarker: "artifact-origin:origin-a1b2c3", + LegacyOwnerMarkers: []string{"this-pusher-marker"}, + } + + assert.True(t, samePushedSessionOwner( + "artifact-origin:origin-a1b2c3", "origin-a1b2c3", + identity, "this-pusher-marker", nil, + )) + assert.True(t, samePushedSessionOwner( + "this-pusher-marker", "origin-a1b2c3", + identity, "this-pusher-marker", nil, + )) + assert.False(t, samePushedSessionOwner( + "another-pusher-marker", "origin-a1b2c3", + identity, "this-pusher-marker", nil, + )) + assert.False(t, samePushedSessionOwner( + "", "origin-a1b2c3", identity, "this-pusher-marker", nil, + ), "matching an artifact origin must not let an importer seize an ownerless row") +} + func TestDeletePGExcludedSessionRowsUsesSingleArrayParameter(t *testing.T) { execer := &capturePGExec{} @@ -643,6 +873,7 @@ func TestPushSessionRechecksExclusionAfterSuccessfulUpsert(t *testing.T) { Agent: "claude", CreatedAt: "2026-01-01T00:00:00Z", }, + pushedSessionIdentity{ID: "sess-race", Machine: "push-machine"}, "marker", nil, ) @@ -669,14 +900,16 @@ func TestPushSessionCarriesDeletionCauseInStableParameterOrder(t *testing.T) { Agent: "claude", CreatedAt: "2026-01-01T00:00:00Z", DeletedAt: &deletedAt, DeletionCause: &cause, }, + pushedSessionIdentity{ID: "session", Machine: "push-machine"}, "marker", nil, ) require.NoError(t, err) - require.Len(t, state.upsertArgs, 63) + require.Len(t, state.upsertArgs, 64) assert.IsType(t, time.Time{}, state.upsertArgs[12].Value) assert.IsType(t, time.Time{}, state.upsertArgs[13].Value) assert.Equal(t, cause, state.upsertArgs[14].Value) assert.Equal(t, "[]", state.upsertArgs[62].Value) + assert.Equal(t, "null", state.upsertArgs[63].Value) query := strings.ToLower(strings.Join(strings.Fields(state.upsertQuery), " ")) assert.Contains(t, query, @@ -695,8 +928,8 @@ func TestSessionPushFingerprintIncludesDeletionCause(t *testing.T) { withCause.DeletionCause = &cause assert.NotEqual(t, - sessionPushFingerprint(base, base.Machine, "", "", ""), - sessionPushFingerprint(withCause, withCause.Machine, "", "", ""), + sessionPushFingerprint(base, base.ID, base.Machine, "", "", ""), + sessionPushFingerprint(withCause, withCause.ID, withCause.Machine, "", "", ""), ) } @@ -722,6 +955,7 @@ func TestPushSessionStoresVibeFallbackAlias(t *testing.T) { CreatedAt: "2026-01-01T00:00:00Z", FilePath: &filePath, }, + pushedSessionIdentity{ID: "vibe:canonical-uuid", Machine: "push-machine"}, "marker", nil, ) @@ -733,6 +967,95 @@ func TestPushSessionStoresVibeFallbackAlias(t *testing.T) { require.NoError(t, tx.Rollback(), "Rollback") } +func TestPushSessionStoresAliasUnderResolvedPGID(t *testing.T) { + state := &pushSessionProbeState{aliases: map[string]string{}} + pg := newPushSessionProbeDB(t, state) + tx, err := pg.BeginTx(context.Background(), nil) + require.NoError(t, err, "BeginTx") + + sessionDir := filepath.Join( + t.TempDir(), + "session_20260616_083518_alias1", + ) + filePath := filepath.Join(sessionDir, "messages.jsonl") + syncer := &Sync{machine: "push-machine"} + err = syncer.pushSession( + context.Background(), tx, + db.Session{ + ID: "vibe:canonical-uuid", + Project: "proj", + Machine: "push-machine", + Agent: "vibe", + CreatedAt: "2026-01-01T00:00:00Z", + FilePath: &filePath, + }, + pushedSessionIdentity{ + ID: "push-machine~vibe:canonical-uuid", + Machine: "push-machine", + }, + "marker", nil, + ) + + require.NoError(t, err, "pushSession") + assert.Empty(t, state.aliases["vibe:canonical-uuid"]) + assert.Equal(t, + "push-machine~vibe:session_20260616_083518_alias1", + state.aliases["push-machine~vibe:canonical-uuid"], + ) + require.NoError(t, tx.Rollback(), "Rollback") +} + +func TestPushSessionIgnoresBareTombstoneForResolvedPGID(t *testing.T) { + state := &pushSessionProbeState{ + aliases: map[string]string{}, + existingExcluded: map[string]bool{ + "vibe:canonical-deleted": true, + }, + excludedIDs: map[string]bool{}, + } + pg := newPushSessionProbeDB(t, state) + tx, err := pg.BeginTx(context.Background(), nil) + require.NoError(t, err, "BeginTx") + + sessionDir := filepath.Join( + t.TempDir(), + "session_20260616_083518_bare01", + ) + filePath := filepath.Join(sessionDir, "messages.jsonl") + syncer := &Sync{machine: "push-machine"} + err = syncer.pushSession( + context.Background(), tx, + db.Session{ + ID: "vibe:canonical-deleted", + Project: "proj", + Machine: "push-machine", + Agent: "vibe", + CreatedAt: "2026-01-01T00:00:00Z", + FilePath: &filePath, + }, + pushedSessionIdentity{ + ID: "push-machine~vibe:canonical-deleted", + Machine: "push-machine", + }, + "marker", nil, + ) + + require.NoError(t, err, "pushSession") + assert.False(t, + state.excludedIDs["vibe:canonical-deleted"], + "resolved pushes must not adopt another owner's bare tombstone", + ) + assert.False(t, + state.deletedExcluded, + "resolved pushes must not delete another owner's bare row", + ) + assert.Equal(t, + "push-machine~vibe:session_20260616_083518_bare01", + state.aliases["push-machine~vibe:canonical-deleted"], + ) + require.NoError(t, tx.Rollback(), "Rollback") +} + func TestPushSessionExcludesVibeFallbackAliasWhenCanonicalExcluded(t *testing.T) { state := &pushSessionProbeState{ existingExcluded: map[string]bool{ @@ -760,6 +1083,7 @@ func TestPushSessionExcludesVibeFallbackAliasWhenCanonicalExcluded(t *testing.T) CreatedAt: "2026-01-01T00:00:00Z", FilePath: &filePath, }, + pushedSessionIdentity{ID: "vibe:canonical-deleted", Machine: "push-machine"}, "marker", nil, ) @@ -798,6 +1122,7 @@ func TestPushSessionSkipsVibeCanonicalWhenFallbackAliasExcluded(t *testing.T) { CreatedAt: "2026-01-01T00:00:00Z", FilePath: &filePath, }, + pushedSessionIdentity{ID: "vibe:canonical-active", Machine: "push-machine"}, "marker", nil, ) @@ -835,9 +1160,15 @@ func TestPurgePGExcludedPushSessionsChecksDerivedAliases(t *testing.T) { FilePath: &filePath, }, } + identities := map[string]pushedSessionIdentity{ + "vibe:canonical-unchanged": { + ID: "vibe:canonical-unchanged", + Machine: "push-machine", + }, + } err := purgePGExcludedPushSessions( - context.Background(), pg, sessionByID, + context.Background(), pg, sessionByID, identities, ) require.NoError(t, err, "purgePGExcludedPushSessions") @@ -852,8 +1183,55 @@ func TestPurgePGExcludedPushSessionsChecksDerivedAliases(t *testing.T) { assert.Equal(t, 0, state.upserts) } +func TestPurgePGExcludedPushSessionsUsesResolvedPGID(t *testing.T) { + state := &pushSessionProbeState{ + existingExcluded: map[string]bool{ + "vibe:canonical-foreign": true, + }, + excludedIDs: map[string]bool{}, + } + pg := newPushSessionProbeDB(t, state) + + sessionDir := filepath.Join( + t.TempDir(), + "session_20260616_083518_resolved", + ) + filePath := filepath.Join(sessionDir, "messages.jsonl") + sessionByID := map[string]db.Session{ + "vibe:canonical-foreign": { + ID: "vibe:canonical-foreign", + Project: "proj", + Machine: "push-machine", + Agent: "vibe", + CreatedAt: "2026-01-01T00:00:00Z", + FilePath: &filePath, + }, + } + identities := map[string]pushedSessionIdentity{ + "vibe:canonical-foreign": { + ID: "push-machine~vibe:canonical-foreign", + Machine: "push-machine", + }, + } + + err := purgePGExcludedPushSessions( + context.Background(), pg, sessionByID, identities, + ) + + require.NoError(t, err, "purgePGExcludedPushSessions") + assert.Contains(t, sessionByID, "vibe:canonical-foreign") + assert.Empty(t, state.excludedIDs) + assert.False(t, + state.deletedExcluded, + "resolved purge must not delete another owner's bare row", + ) + assert.Equal(t, 1, state.exclusionChecks) +} + type pushSessionProbeDriver struct{} +var postgresParameterPattern = regexp.MustCompile(`\$(\d+)`) + type pushSessionProbeConn struct { state *pushSessionProbeState } @@ -877,6 +1255,13 @@ type pushSessionProbeState struct { existingExcluded map[string]bool upsertQuery string upsertArgs []driver.NamedValue + ownerQueries int + owners map[string]pushSessionProbeOwner +} + +type pushSessionProbeOwner struct { + machine string + marker string } var ( @@ -946,6 +1331,9 @@ func (c *pushSessionProbeConn) ExecContext( switch { case strings.Contains(normalized, "insert into sessions"): + if err := validatePostgresParameterCount(query, len(args)); err != nil { + return nil, err + } c.state.upserts++ c.state.upsertQuery = query c.state.upsertArgs = append([]driver.NamedValue(nil), args...) @@ -987,6 +1375,24 @@ func (c *pushSessionProbeConn) ExecContext( } } +func validatePostgresParameterCount(query string, argumentCount int) error { + highestParameter := 0 + for _, match := range postgresParameterPattern.FindAllStringSubmatch(query, -1) { + parameter, err := strconv.Atoi(match[1]) + if err != nil { + return fmt.Errorf("parsing postgres parameter %q: %w", match[0], err) + } + highestParameter = max(highestParameter, parameter) + } + if highestParameter != argumentCount { + return fmt.Errorf( + "postgres query references $%d but received %d arguments", + highestParameter, argumentCount, + ) + } + return nil +} + func (c *pushSessionProbeConn) QueryContext( _ context.Context, query string, args []driver.NamedValue, ) (driver.Rows, error) { @@ -995,7 +1401,29 @@ func (c *pushSessionProbeConn) QueryContext( defer c.state.mu.Unlock() switch { + case strings.Contains(normalized, "select id, machine, owner_marker"): + c.state.ownerQueries++ + values := [][]driver.Value{} + for _, id := range namedValueStrings(args) { + if owner, ok := c.state.owners[id]; ok { + values = append(values, []driver.Value{id, owner.machine, owner.marker}) + } + } + return &pushSessionProbeRows{ + columns: []string{"id", "machine", "owner_marker"}, + values: values, + }, nil case strings.Contains(normalized, "select machine, owner_marker"): + c.state.ownerQueries++ + if len(args) > 0 { + id, _ := args[0].Value.(string) + if owner, ok := c.state.owners[id]; ok { + return &pushSessionProbeRows{ + columns: []string{"machine", "owner_marker"}, + values: [][]driver.Value{{owner.machine, owner.marker}}, + }, nil + } + } return &pushSessionProbeRows{ columns: []string{"machine", "owner_marker"}, }, nil @@ -1032,6 +1460,71 @@ func (c *pushSessionProbeConn) QueryContext( } } +func TestPreloadPGSessionOwnersUsesOneQueryAndCachesMisses(t *testing.T) { + state := &pushSessionProbeState{owners: map[string]pushSessionProbeOwner{ + "owned-a": {machine: "desk", marker: "marker-a"}, + "owned-b": {machine: "laptop", marker: "marker-b"}, + }} + sync := &Sync{pg: newPushSessionProbeDB(t, state)} + + ctx, err := sync.preloadPGSessionOwners( + context.Background(), []string{"owned-a", "owned-b", "missing"}, + ) + require.NoError(t, err) + for _, tc := range []struct { + id string + machine string + marker string + exists bool + }{ + {id: "owned-a", machine: "desk", marker: "marker-a", exists: true}, + {id: "owned-b", machine: "laptop", marker: "marker-b", exists: true}, + {id: "missing"}, + } { + machine, marker, exists, lookupErr := sync.pgSessionOwner(ctx, tc.id) + require.NoError(t, lookupErr) + assert.Equal(t, tc.machine, machine) + assert.Equal(t, tc.marker, marker) + assert.Equal(t, tc.exists, exists) + } + assert.Equal(t, 1, state.ownerQueries, + "preloaded hits and misses must use one owner query") + + _, _, exists, err := sync.pgSessionOwner(ctx, "late-miss") + require.NoError(t, err) + assert.False(t, exists) + _, _, exists, err = sync.pgSessionOwner(ctx, "late-miss") + require.NoError(t, err) + assert.False(t, exists) + assert.Equal(t, 2, state.ownerQueries, + "an owner first discovered after preload must be memoized") +} + +func TestPushIdentityOwnerCandidateIDsCoverLegacyAndArtifactAliases(t *testing.T) { + sync := &Sync{machine: "desk"} + sessions := map[string]db.Session{ + "plain": { + ID: "plain", Machine: "local", + }, + "remote-a1b2c3~imported": { + ID: "remote-a1b2c3~imported", Machine: "remote-a1b2c3", + }, + } + imported := map[string]struct{}{"remote-a1b2c3~imported": {}} + + got := sync.pushIdentityOwnerCandidateIDs( + sessions, imported, "desk-origin", "marker", []string{"old-desk"}, + ) + assert.ElementsMatch(t, []string{ + "desk-origin~plain", + "old-desk~desk-origin~plain", + "plain", + "remote-a1b2c3~imported", + "old-desk~remote-a1b2c3~imported", + "imported", + }, got) +} + func (pushSessionProbeTx) Commit() error { return nil } func (pushSessionProbeTx) Rollback() error { return nil } @@ -1060,7 +1553,7 @@ func TestSessionPushFingerprintDiffers(t *testing.T) { CreatedAt: "2026-03-11T12:00:00Z", } - fp1 := sessionPushFingerprint(base, base.Machine, "", "", "") + fp1 := sessionPushFingerprint(base, base.ID, base.Machine, "", "", "") tests := []struct { name string @@ -1170,13 +1663,14 @@ func TestSessionPushFingerprintDiffers(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { modified := tc.modify(base) - fp2 := sessionPushFingerprint(modified, modified.Machine, "", "", "") + fp2 := sessionPushFingerprint( + modified, modified.ID, modified.Machine, "", "", "") require.NotEqual(t, fp1, fp2, "fingerprint should differ after %s", tc.name) }) } - assert.Equal(t, fp1, sessionPushFingerprint(base, base.Machine, "", "", ""), + assert.Equal(t, fp1, sessionPushFingerprint(base, base.ID, base.Machine, "", "", ""), "identical sessions should produce identical fingerprints") } @@ -1196,24 +1690,24 @@ func TestSessionPushFingerprintIgnoresVolatileStatFields(t *testing.T) { LocalModifiedAt: &localModifiedAt, CreatedAt: "2026-03-11T12:00:00Z", } - baseFP := sessionPushFingerprint(base, base.Machine, "", "", "deps") + baseFP := sessionPushFingerprint(base, base.ID, base.Machine, "", "", "deps") statOnlyMtime := int64(1700000001000000000) statOnlyModifiedAt := "2026-03-11T12:00:01.000Z" statOnly := base statOnly.FileMtime = &statOnlyMtime statOnly.LocalModifiedAt = &statOnlyModifiedAt - assert.Equal(t, baseFP, sessionPushFingerprint(statOnly, statOnly.Machine, "", "", "deps"), + assert.Equal(t, baseFP, sessionPushFingerprint(statOnly, statOnly.ID, statOnly.Machine, "", "", "deps"), "file stat churn should not change push candidacy") contentChanged := statOnly contentChanged.MessageCount++ assert.NotEqual(t, baseFP, - sessionPushFingerprint(contentChanged, contentChanged.Machine, "", "", "deps"), + sessionPushFingerprint(contentChanged, contentChanged.ID, contentChanged.Machine, "", "", "deps"), "content changes should still change push candidacy") assert.NotEqual(t, baseFP, - sessionPushFingerprint(statOnly, statOnly.Machine, "", "", "changed-deps"), + sessionPushFingerprint(statOnly, statOnly.ID, statOnly.Machine, "", "", "changed-deps"), "dependent row changes should still change push candidacy") } @@ -1252,7 +1746,7 @@ func TestLocalSessionDependencyPushFingerprintTracksMessageEditsWithoutFileHash( CreatedAt: "2026-03-11T12:00:00Z", } fpBefore := sessionPushFingerprint( - session, session.Machine, "", "", depsBefore, + session, session.ID, session.Machine, "", "", depsBefore, ) require.NoError(t, localDB.ReplaceSessionMessages(sessID, []db.Message{{ @@ -1266,7 +1760,7 @@ func TestLocalSessionDependencyPushFingerprintTracksMessageEditsWithoutFileHash( ) require.NoError(t, err) fpAfter := sessionPushFingerprint( - session, session.Machine, "", "", depsAfter, + session, session.ID, session.Machine, "", "", depsAfter, ) assert.NotEqual(t, depsBefore, depsAfter) @@ -1287,12 +1781,29 @@ func TestSessionPushFingerprintIncludesUsageEventFingerprint( CreatedAt: "2026-03-11T12:00:00Z", } - withoutUsage := sessionPushFingerprint(base, base.Machine, "", "", "") - withUsage := sessionPushFingerprint(base, base.Machine, "usage-fp", "", "") + withoutUsage := sessionPushFingerprint(base, base.ID, base.Machine, "", "", "") + withUsage := sessionPushFingerprint(base, base.ID, base.Machine, "usage-fp", "", "") assert.NotEqual(t, withoutUsage, withUsage, "usage event fingerprint should affect session fingerprint") } +func TestSessionPushFingerprintTracksResolvedID(t *testing.T) { + base := db.Session{ + ID: "sess-001", + Project: "proj", + Machine: "laptop", + Agent: "claude", + CreatedAt: "2026-03-11T12:00:00Z", + } + + native := sessionPushFingerprint(base, base.ID, base.Machine, "", "", "") + prefixed := sessionPushFingerprint( + base, prefixedSessionID(base.Machine, base.ID), base.Machine, "", "", "", + ) + assert.NotEqual(t, native, prefixed, + "resolved PG id must affect session fingerprint") +} + func TestSessionPushFingerprintTracksResolvedMachine(t *testing.T) { sentinel := db.Session{ ID: "sess-001", @@ -1302,9 +1813,11 @@ func TestSessionPushFingerprintTracksResolvedMachine(t *testing.T) { CreatedAt: "2026-03-11T12:00:00Z", } fpA := sessionPushFingerprint( - sentinel, pushedSessionMachine(sentinel, "host-a"), "", "", "") + sentinel, sentinel.ID, + pushedSessionMachine(sentinel, "host-a"), "", "", "") fpB := sessionPushFingerprint( - sentinel, pushedSessionMachine(sentinel, "host-b"), "", "", "") + sentinel, sentinel.ID, + pushedSessionMachine(sentinel, "host-b"), "", "", "") assert.NotEqual(t, fpA, fpB, "sentinel session fingerprint must change with the fallback machine") @@ -1316,9 +1829,11 @@ func TestSessionPushFingerprintTracksResolvedMachine(t *testing.T) { CreatedAt: "2026-03-11T12:00:00Z", } fp1 := sessionPushFingerprint( - real, pushedSessionMachine(real, "host-a"), "", "", "") + real, real.ID, + pushedSessionMachine(real, "host-a"), "", "", "") fp2 := sessionPushFingerprint( - real, pushedSessionMachine(real, "host-b"), "", "", "") + real, real.ID, + pushedSessionMachine(real, "host-b"), "", "", "") assert.Equal(t, fp1, fp2, "a session with a real machine ignores the fallback") } @@ -1362,6 +1877,46 @@ func TestPushedSessionMachine(t *testing.T) { } } +func TestPrefixedSessionID(t *testing.T) { + tests := []struct { + name string + machine string + id string + want string + }{ + { + name: "prefixes native id", + machine: "host-a", + id: "sess-001", + want: "host-a~sess-001", + }, + { + name: "keeps already prefixed id", + machine: "host-a", + id: "host-a~sess-001", + want: "host-a~sess-001", + }, + { + name: "keeps empty machine", + machine: "", + id: "sess-001", + want: "sess-001", + }, + { + name: "keeps empty id", + machine: "host-a", + id: "", + want: "", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, prefixedSessionID(tc.machine, tc.id)) + }) + } +} + func TestSessionPushFingerprintNoFieldCollisions( t *testing.T, ) { @@ -1376,8 +1931,8 @@ func TestSessionPushFingerprintNoFieldCollisions( CreatedAt: "2026-03-11T12:00:00Z", } assert.NotEqual(t, - sessionPushFingerprint(s1, s1.Machine, "", "", ""), - sessionPushFingerprint(s2, s2.Machine, "", "", ""), + sessionPushFingerprint(s1, s1.ID, s1.Machine, "", "", ""), + sessionPushFingerprint(s2, s2.ID, s2.Machine, "", "", ""), "length-prefixed fingerprints should not collide") } @@ -1662,7 +2217,16 @@ func TestFinalizePushStateMergesPriorFingerprints( require.NoError(t, finalizePushState( store, cutoff, cycle2Sessions, priorFingerprints, - map[string]string{"sess-002": sessionPushFingerprint(cycle2Sessions[0], cycle2Sessions[0].Machine, "", "", "")}, + map[string]string{ + "sess-002": sessionPushFingerprint( + cycle2Sessions[0], + cycle2Sessions[0].ID, + cycle2Sessions[0].Machine, + "", + "", + "", + ), + }, )) raw := store.values[lastPushBoundaryStateKey] diff --git a/internal/postgres/schema.go b/internal/postgres/schema.go index 5c684561e..ca58f72b2 100644 --- a/internal/postgres/schema.go +++ b/internal/postgres/schema.go @@ -230,10 +230,6 @@ CREATE INDEX IF NOT EXISTS idx_pinned_session CREATE INDEX IF NOT EXISTS idx_pinned_created ON pinned_messages (created_at DESC); -CREATE INDEX IF NOT EXISTS idx_pinned_source_uuid - ON pinned_messages (session_id, source_uuid) - WHERE source_uuid <> ''; - CREATE TABLE IF NOT EXISTS model_pricing ( model_pattern TEXT PRIMARY KEY, input_per_mtok DOUBLE PRECISION NOT NULL DEFAULT 0, @@ -758,6 +754,11 @@ func EnsureSchema( `thinking_text TEXT NOT NULL DEFAULT ''`, "adding messages.thinking_text", }, + { + "pinned_messages", "source_uuid", + `source_uuid TEXT NOT NULL DEFAULT ''`, + "adding pinned_messages.source_uuid", + }, { "sessions", "termination_status", `termination_status TEXT`, @@ -985,12 +986,17 @@ func createPartialIndexesPG(ctx context.Context, db *sql.DB) error { ON sessions(cwd) WHERE cwd != ''`, `CREATE INDEX IF NOT EXISTS idx_sessions_project_git_branch ON sessions(project, git_branch) WHERE git_branch != ''`, + `CREATE INDEX IF NOT EXISTS idx_sessions_source_session + ON sessions(source_session_id) WHERE source_session_id != ''`, `CREATE INDEX IF NOT EXISTS idx_messages_compact_boundary ON messages(session_id, ordinal) WHERE is_compact_boundary = TRUE`, `CREATE INDEX IF NOT EXISTS idx_messages_sidechain ON messages(session_id) WHERE is_sidechain = TRUE`, `CREATE INDEX IF NOT EXISTS idx_messages_source_uuid ON messages(source_uuid) WHERE source_uuid != ''`, + `CREATE INDEX IF NOT EXISTS idx_pinned_source_uuid + ON pinned_messages(session_id, source_uuid) + WHERE source_uuid <> ''`, `CREATE INDEX IF NOT EXISTS idx_messages_usage_covering ON messages(timestamp, session_id, ordinal, model, claude_message_id, claude_request_id) @@ -1004,6 +1010,12 @@ func createPartialIndexesPG(ctx context.Context, db *sql.DB) error { // SQLite partial index so legacy schemas migrate cleanly. `CREATE INDEX IF NOT EXISTS idx_tool_calls_file_path ON tool_calls(file_path) WHERE file_path IS NOT NULL`, + `CREATE INDEX IF NOT EXISTS idx_tool_calls_subagent_session + ON tool_calls(subagent_session_id) + WHERE subagent_session_id IS NOT NULL`, + `CREATE INDEX IF NOT EXISTS idx_tool_result_events_subagent_session + ON tool_result_events(subagent_session_id) + WHERE subagent_session_id IS NOT NULL`, // idx_messages_session_role backs the dense-flow unit-range boundary // fetch (user ordinals by session), mirroring the SQLite index. `CREATE INDEX IF NOT EXISTS idx_messages_session_role diff --git a/internal/postgres/schema_pgtest_test.go b/internal/postgres/schema_pgtest_test.go index 640dd4915..d321a9076 100644 --- a/internal/postgres/schema_pgtest_test.go +++ b/internal/postgres/schema_pgtest_test.go @@ -158,3 +158,54 @@ func TestToolCallsFilePathIndex(t *testing.T) { require.NoError(t, err, "checking idx_tool_calls_file_path") assert.True(t, exists, "idx_tool_calls_file_path index missing") } + +func TestEnsureSchemaMigratesPinnedMessageSourceUUIDBeforeIndex(t *testing.T) { + pgURL := testPGURL(t) + cleanSchemaTestPG(t, pgURL) + t.Cleanup(func() { cleanSchemaTestPG(t, pgURL) }) + + pg, err := Open(pgURL, schemaTestSchema, true) + require.NoError(t, err, "connecting to pg") + defer pg.Close() + + ctx := context.Background() + _, err = pg.ExecContext(ctx, + `CREATE SCHEMA IF NOT EXISTS `+schemaTestSchema) + require.NoError(t, err, "creating test schema") + _, err = pg.ExecContext(ctx, ` + CREATE TABLE pinned_messages ( + id BIGSERIAL PRIMARY KEY, + session_id TEXT NOT NULL, + message_id INT NOT NULL, + ordinal INT NOT NULL, + note TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + UNIQUE (session_id, message_id) + )`) + require.NoError(t, err, "creating legacy pinned_messages table") + + require.NoError(t, EnsureSchema(ctx, pg, schemaTestSchema), + "EnsureSchema should add source_uuid before creating its index") + + var columnExists bool + err = pg.QueryRowContext(ctx, ` + SELECT EXISTS ( + SELECT 1 FROM information_schema.columns + WHERE table_schema = $1 + AND table_name = 'pinned_messages' + AND column_name = 'source_uuid' + )`, schemaTestSchema).Scan(&columnExists) + require.NoError(t, err, "checking pinned_messages.source_uuid") + assert.True(t, columnExists, "pinned_messages.source_uuid missing") + + var indexExists bool + err = pg.QueryRowContext(ctx, ` + SELECT EXISTS ( + SELECT 1 FROM pg_indexes + WHERE schemaname = $1 + AND tablename = 'pinned_messages' + AND indexname = 'idx_pinned_source_uuid' + )`, schemaTestSchema).Scan(&indexExists) + require.NoError(t, err, "checking idx_pinned_source_uuid") + assert.True(t, indexExists, "idx_pinned_source_uuid index missing") +} diff --git a/internal/postgres/schema_test.go b/internal/postgres/schema_test.go index b06d193c4..e08d158f7 100644 --- a/internal/postgres/schema_test.go +++ b/internal/postgres/schema_test.go @@ -6,6 +6,7 @@ import ( "database/sql/driver" "errors" "io" + "slices" "strings" "sync" "testing" @@ -109,19 +110,20 @@ func (c *schemaProbeConn) Begin() (driver.Tx, error) { func (c *schemaProbeConn) ExecContext( _ context.Context, query string, args []driver.NamedValue, ) (driver.Result, error) { + normalized := strings.ToLower(query) + if strings.Contains(normalized, "idx_pinned_source_uuid") && + c.state.hasColumn("pinned_messages", "session_id") && + !c.state.hasColumn("pinned_messages", "source_uuid") { + return nil, errors.New(`ERROR: column "source_uuid" does not exist (SQLSTATE 42703)`) + } c.state.mu.Lock() c.state.execs = append(c.state.execs, query) c.state.execArgs = append( c.state.execArgs, append([]driver.NamedValue(nil), args...), ) c.state.mu.Unlock() - normalized := strings.ToLower(query) if strings.Contains(normalized, "alter table") { - c.state.mu.Lock() - c.state.alterTableExecs = append( - c.state.alterTableExecs, query, - ) - c.state.mu.Unlock() + c.state.recordAlterTable(query) } if strings.Contains(normalized, "insert into sync_metadata") && len(args) > 0 { @@ -136,6 +138,65 @@ func (c *schemaProbeConn) ExecContext( return driver.RowsAffected(0), nil } +func (s *schemaProbeState) hasColumn(table, column string) bool { + s.mu.Lock() + defer s.mu.Unlock() + return slices.Contains(s.existingColumnNames[table], column) +} + +func (s *schemaProbeState) recordAlterTable(query string) { + s.mu.Lock() + defer s.mu.Unlock() + s.alterTableExecs = append(s.alterTableExecs, query) + + table, ok := alterTableName(query) + if !ok { + return + } + if s.existingColumnNames == nil { + s.existingColumnNames = map[string][]string{} + } + for _, column := range alterTableColumns(query) { + exists := slices.Contains(s.existingColumnNames[table], column) + if !exists { + s.existingColumnNames[table] = append( + s.existingColumnNames[table], column, + ) + } + } +} + +func alterTableName(query string) (string, bool) { + const prefix = `ALTER TABLE "` + _, after, ok := strings.Cut(query, prefix) + if !ok { + return "", false + } + rest := after + before0, _, ok0 := strings.Cut(rest, `"`) + if !ok0 { + return "", false + } + return before0, true +} + +func alterTableColumns(query string) []string { + const marker = "ADD COLUMN IF NOT EXISTS " + parts := strings.Split(query, marker) + if len(parts) < 2 { + return nil + } + columns := make([]string, 0, len(parts)-1) + for _, part := range parts[1:] { + fields := strings.Fields(part) + if len(fields) == 0 { + continue + } + columns = append(columns, strings.Trim(fields[0], `",`)) + } + return columns +} + func (c *schemaProbeConn) QueryContext( _ context.Context, query string, args []driver.NamedValue, ) (driver.Rows, error) { @@ -860,6 +921,12 @@ func TestEnsureSchemaCreatesSessionTraversalIndex(t *testing.T) { assert.Contains(t, state.executedSQL(), "CREATE INDEX IF NOT EXISTS idx_sessions_parent") + assert.Contains(t, state.executedSQL(), + "CREATE INDEX IF NOT EXISTS idx_sessions_source_session") + assert.Contains(t, state.executedSQL(), + "CREATE INDEX IF NOT EXISTS idx_tool_calls_subagent_session") + assert.Contains(t, state.executedSQL(), + "CREATE INDEX IF NOT EXISTS idx_tool_result_events_subagent_session") } func TestEnsureSchemaGroupsMissingColumnMigrationsByTable(t *testing.T) { @@ -899,6 +966,9 @@ func TestEnsureSchemaGroupsMissingColumnMigrationsByTable(t *testing.T) { "tool_calls": { "call_index", "file_path", }, + "pinned_messages": { + "source_uuid", + }, }) require.NoError(t, EnsureSchema(context.Background(), db, "agentsview")) @@ -911,3 +981,25 @@ func TestEnsureSchemaGroupsMissingColumnMigrationsByTable(t *testing.T) { // it contributes no ALTER. assert.Equal(t, 3, state.alterTableExecCount(), "ALTER TABLE execs") } + +func TestEnsureSchemaMigratesPinnedMessageSourceUUID(t *testing.T) { + db, state := newSchemaProbeDB(t, map[string][]string{ + "sessions": { + "has_total_output_tokens", + "has_peak_context_tokens", + }, + "messages": { + "has_context_tokens", + "has_output_tokens", + }, + "pinned_messages": { + "id", "session_id", "message_id", "ordinal", + "note", "created_at", + }, + }) + + require.NoError(t, EnsureSchema(context.Background(), db, "agentsview")) + + assert.Contains(t, state.executedSQL(), + "ALTER TABLE \"pinned_messages\" ADD COLUMN IF NOT EXISTS source_uuid TEXT NOT NULL DEFAULT ''") +} diff --git a/internal/postgres/sessions.go b/internal/postgres/sessions.go index a7f94d2ce..98191c4f1 100644 --- a/internal/postgres/sessions.go +++ b/internal/postgres/sessions.go @@ -1121,6 +1121,31 @@ func (s *Store) GetAgents( } // GetMachines returns distinct machine names. +// MachineSessionCounts returns the number of non-deleted sessions per machine, +// keyed by machine name. +func (s *Store) MachineSessionCounts(ctx context.Context) (map[string]int, error) { + rows, err := s.pg.QueryContext(ctx, + `SELECT machine, COUNT(*) FROM sessions + WHERE deleted_at IS NULL + GROUP BY machine`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + counts := map[string]int{} + for rows.Next() { + var machine string + var count int + if err := rows.Scan(&machine, &count); err != nil { + return nil, err + } + counts[machine] = count + } + return counts, rows.Err() +} + func (s *Store) GetMachines( ctx context.Context, excludeOneShot, excludeAutomated bool, diff --git a/internal/postgres/store.go b/internal/postgres/store.go index 978cabfb3..0154f9eae 100644 --- a/internal/postgres/store.go +++ b/internal/postgres/store.go @@ -462,6 +462,51 @@ func (s *Store) SoftDeleteSessions(ids []string) (int, error) { return total, nil } +// SoftDeleteSessionsReturningIDs moves multiple sessions to the trash and +// returns the IDs that changed state. +func (s *Store) SoftDeleteSessionsReturningIDs(ids []string) ([]string, error) { + if len(ids) == 0 { + return nil, nil + } + deleted := make([]string, 0, len(ids)) + const batchSize = 500 + for start := 0; start < len(ids); start += batchSize { + end := min(start+batchSize, len(ids)) + pb := ¶mBuilder{} + placeholders := make([]string, 0, end-start) + for _, id := range ids[start:end] { + placeholders = append(placeholders, pb.add(id)) + } + rows, err := s.pg.Query( + `UPDATE sessions + SET deleted_at = NOW(), + updated_at = NOW() + WHERE id IN (`+strings.Join(placeholders, ",")+ + `) AND deleted_at IS NULL + RETURNING id`, + pb.args..., + ) + if err != nil { + return deleted, mapPGWriteError("soft deleting sessions", err) + } + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + _ = rows.Close() + return deleted, fmt.Errorf("scanning soft deleted session: %w", err) + } + deleted = append(deleted, id) + } + if err := rows.Close(); err != nil { + return deleted, fmt.Errorf("closing soft deleted session rows: %w", err) + } + if err := rows.Err(); err != nil { + return deleted, fmt.Errorf("iterating soft deleted sessions: %w", err) + } + } + return deleted, nil +} + // RestoreSession restores a trashed session. func (s *Store) RestoreSession(id string) (int64, error) { res, err := s.pg.Exec( diff --git a/internal/postgres/sync.go b/internal/postgres/sync.go index d70f122dc..1322ac09b 100644 --- a/internal/postgres/sync.go +++ b/internal/postgres/sync.go @@ -52,6 +52,7 @@ func (s *scopedSyncStateStore) ensureMigration() error { "last_push_at", lastPushBoundaryStateKey, lastPushTargetFingerprintKey, + artifactIdentityModeStateKey, } { scopedKey := s.scopedKey(key) scopedValue, err := s.base.GetSyncState(scopedKey) diff --git a/internal/secrets/rules_test.go b/internal/secrets/rules_test.go index 4c798ba5f..df58e4575 100644 --- a/internal/secrets/rules_test.go +++ b/internal/secrets/rules_test.go @@ -8,6 +8,13 @@ import ( "github.com/stretchr/testify/require" ) +// testAWSAccessKeyID builds a scanner fixture at runtime so repository push +// protection does not mistake the deliberately credential-shaped test value +// for a usable credential in source control. +func testAWSAccessKeyID() string { + return strings.Join([]string{"AK", "IA", "7QHWN2DKR4FYPLJM"}, "") +} + func TestDefiniteRules(t *testing.T) { cases := []struct { name string @@ -186,7 +193,8 @@ func TestCandidateRules(t *testing.T) { // (high-entropy assignments, JWTs, basic-auth URLs) entirely. func TestScanDefiniteReturnsOnlyDefinite(t *testing.T) { // One definite AWS key and one candidate high-entropy assignment. - text := "aws AKIA7QHWN2DKR4FYPLJM and SECRET=Xa9Kd03Lm5Qp7Rt2Vw8Zb4Nc6" + text := "aws " + testAWSAccessKeyID() + + " and SECRET=Xa9Kd03Lm5Qp7Rt2Vw8Zb4Nc6" full := Scan(text) require.Len(t, full, 2, "precondition: Scan should report 2 matches (1 definite, 1 candidate)") @@ -203,7 +211,8 @@ func TestScanDefiniteReturnsOnlyDefinite(t *testing.T) { // same spans (rule, offsets, redaction) that Scan reports for definite rules, // so findings stored by the inline path and the full scan stay consistent. func TestScanDefiniteMatchesScanDefiniteSubset(t *testing.T) { - text := "key AKIA7QHWN2DKR4FYPLJM tok ghp_8Hk3Wn7Dz4Rp2Vx9Mb6Tj0Qc5Lm1Yp8Bv4Hg" + + text := "key " + testAWSAccessKeyID() + + " tok ghp_8Hk3Wn7Dz4Rp2Vx9Mb6Tj0Qc5Lm1Yp8Bv4Hg" + " SECRET=Xa9Kd03Lm5Qp7Rt2Vw8Zb4Nc6" var wantDef []Match for _, m := range Scan(text) { @@ -254,9 +263,10 @@ func TestRulesVersionStableAndHex(t *testing.T) { func TestVerify(t *testing.T) { // Non-grouped rule: the stored span is the full regex match. - awsSrc := "export KEY=AKIA7QHWN2DKR4FYPLJM done" - s := strings.Index(awsSrc, "AKIA") - e := s + len("AKIA7QHWN2DKR4FYPLJM") + awsKey := testAWSAccessKeyID() + awsSrc := "export KEY=" + awsKey + " done" + s := strings.Index(awsSrc, awsKey) + e := s + len(awsKey) assert.True(t, Verify("aws-access-key", awsSrc, s, e), "Verify should accept a valid AWS key at its coordinates") assert.False(t, Verify("aws-access-key", awsSrc, 0, 6), @@ -278,7 +288,7 @@ func TestVerify(t *testing.T) { // produces coordinates, Verify accepts them on the unchanged source, and // rejects them once the bytes at those coordinates are no longer the secret. func TestVerifyDetectsChangedSource(t *testing.T) { - source := "export AWS=AKIA7QHWN2DKR4FYPLJM" + source := "export AWS=" + testAWSAccessKeyID() // Seed from canonical Scan (what produces findings and what Verify uses). matches := Scan(source) require.NotEmpty(t, matches, "expected at least one match in source") @@ -379,9 +389,7 @@ func TestHighEntropyPaddingCapture(t *testing.T) { break } } - if m == nil { - t.Fatalf("no high-entropy match in %q; got %+v", c.text, got) - } + require.NotNil(t, m, "no high-entropy match in %q; got %+v", c.text, got) span := c.text[m.Start:m.End] if !strings.HasSuffix(span, c.suffix) { t.Errorf("captured span %q does not end with %q", diff --git a/internal/server/artifact_cursors.go b/internal/server/artifact_cursors.go new file mode 100644 index 000000000..cc5478003 --- /dev/null +++ b/internal/server/artifact_cursors.go @@ -0,0 +1,499 @@ +package server + +import ( + "context" + "crypto/rand" + "database/sql" + "encoding/base64" + "errors" + "fmt" + "io/fs" + "os" + "sync" + "time" + + "go.kenn.io/agentsview/internal/artifact" +) + +const ( + artifactCursorTTL = 5 * time.Minute + maxArtifactCursors = 64 +) + +type artifactCursorRegistry struct { + mu sync.Mutex + tokens map[string]*artifactCursorLease + active map[*artifactCursorLease]struct{} + closed bool +} + +type artifactCursorLease struct { + registry *artifactCursorRegistry + cursor *artifactSnapshotCursor + scope string + token string + timer *time.Timer + constructCancel context.CancelFunc + released bool +} + +type artifactSnapshotCursor struct { + mu sync.Mutex + db *sql.DB + path string + scope string + lastKind int + lastName string + closing bool + inFlight int + cleaned bool +} + +type artifactSnapshotItem struct { + kind int + name string +} + +func newArtifactCursorRegistry() *artifactCursorRegistry { + return &artifactCursorRegistry{ + tokens: make(map[string]*artifactCursorLease), + active: make(map[*artifactCursorLease]struct{}), + } +} + +func (r *artifactCursorRegistry) buildSnapshot( + ctx context.Context, + root string, + scope string, + fill func(context.Context, func(int, string) error) error, +) (*artifactCursorLease, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + constructionCtx, cancel := context.WithCancel(ctx) + lease, err := r.reserve(scope, cancel) + if err != nil { + cancel() + return nil, err + } + cursor, err := newArtifactSnapshot(constructionCtx, root, scope, + func(insert func(int, string) error) error { + return fill(constructionCtx, insert) + }) + if err != nil { + lease.release() + return nil, err + } + if err := r.attach(lease, cursor); err != nil { + cursor.close() + lease.release() + return nil, err + } + return lease, nil +} + +func (r *artifactCursorRegistry) reserve( + scope string, cancel context.CancelFunc, +) (*artifactCursorLease, error) { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed { + return nil, fs.ErrClosed + } + if len(r.active) >= maxArtifactCursors { + return nil, fmt.Errorf("%w: too many active artifact cursors", artifact.ErrArtifactConflict) + } + lease := &artifactCursorLease{ + registry: r, + scope: scope, + constructCancel: cancel, + } + r.active[lease] = struct{}{} + return lease, nil +} + +func (r *artifactCursorRegistry) attach( + lease *artifactCursorLease, cursor *artifactSnapshotCursor, +) error { + r.mu.Lock() + if r.closed || lease == nil || lease.released { + r.mu.Unlock() + return fs.ErrClosed + } + if _, ok := r.active[lease]; !ok { + r.mu.Unlock() + return fs.ErrClosed + } + lease.cursor = cursor + cancel := lease.constructCancel + lease.constructCancel = nil + r.mu.Unlock() + if cancel != nil { + cancel() + } + return nil +} + +func newArtifactSnapshot( + ctx context.Context, + root string, + scope string, + fill func(func(int, string) error) error, +) (*artifactSnapshotCursor, error) { + if err := os.MkdirAll(root, 0o755); err != nil { + return nil, err + } + file, err := os.CreateTemp(root, ".peer-cursor-*.sqlite") + if err != nil { + return nil, err + } + path := file.Name() + if err := file.Close(); err != nil { + _ = os.Remove(path) + return nil, err + } + database, err := sql.Open("sqlite3", path+"?_journal_mode=OFF&_synchronous=OFF&_temp_store=FILE") + if err != nil { + _ = os.Remove(path) + return nil, err + } + cursor := &artifactSnapshotCursor{db: database, path: path, scope: scope, lastKind: -1} + cleanup := true + defer func() { + if cleanup { + cursor.close() + } + }() + if _, err := database.ExecContext(ctx, + `CREATE TABLE items (kind INTEGER NOT NULL, name TEXT NOT NULL, PRIMARY KEY (kind, name)) WITHOUT ROWID`); err != nil { + return nil, err + } + tx, err := database.BeginTx(ctx, nil) + if err != nil { + return nil, err + } + committed := false + defer func() { + if !committed { + _ = tx.Rollback() + } + }() + statement, err := tx.PrepareContext(ctx, `INSERT OR IGNORE INTO items (kind, name) VALUES (?, ?)`) + if err != nil { + return nil, err + } + err = fill(func(kind int, name string) error { + if err := ctx.Err(); err != nil { + return err + } + _, err := statement.ExecContext(ctx, kind, name) + return err + }) + closeErr := statement.Close() + if err != nil { + return nil, err + } + if closeErr != nil { + return nil, closeErr + } + if err := tx.Commit(); err != nil { + return nil, err + } + committed = true + cleanup = false + return cursor, nil +} + +func (c *artifactSnapshotCursor) page( + ctx context.Context, limit int, +) (items []artifactSnapshotItem, more bool, retErr error) { + database, lastKind, lastName, err := c.beginPage() + if err != nil { + return nil, false, err + } + defer func() { + closing, cleanupDB, cleanupPath := c.finishPage() + cleanupArtifactSnapshot(cleanupDB, cleanupPath) + if closing { + items = nil + more = false + retErr = errors.Join(retErr, fs.ErrClosed) + } + }() + + rows, err := database.QueryContext(ctx, ` + SELECT kind, name + FROM items + WHERE kind > ? OR (kind = ? AND name > ?) + ORDER BY kind, name + LIMIT ?`, lastKind, lastKind, lastName, limit+1) + if err != nil { + return nil, false, err + } + defer rows.Close() + items = make([]artifactSnapshotItem, 0, limit+1) + for rows.Next() { + var item artifactSnapshotItem + if err := rows.Scan(&item.kind, &item.name); err != nil { + return nil, false, err + } + items = append(items, item) + } + if err := rows.Err(); err != nil { + return nil, false, err + } + more = len(items) > limit + if more { + items = items[:limit] + } + if len(items) > 0 { + last := items[len(items)-1] + c.mu.Lock() + if c.closing { + c.mu.Unlock() + return nil, false, fs.ErrClosed + } + c.lastKind = last.kind + c.lastName = last.name + c.mu.Unlock() + } + return items, more, nil +} + +func (c *artifactSnapshotCursor) beginPage() (*sql.DB, int, string, error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closing || c.db == nil { + return nil, 0, "", fs.ErrClosed + } + if c.inFlight != 0 { + return nil, 0, "", fmt.Errorf("%w: artifact cursor page already active", artifact.ErrArtifactConflict) + } + c.inFlight++ + return c.db, c.lastKind, c.lastName, nil +} + +func (c *artifactSnapshotCursor) finishPage() (bool, *sql.DB, string) { + c.mu.Lock() + defer c.mu.Unlock() + c.inFlight-- + if !c.closing || c.inFlight != 0 { + return c.closing, nil, "" + } + database, path := c.takeCleanupLocked() + return true, database, path +} + +func (c *artifactSnapshotCursor) close() { + if c == nil { + return + } + c.mu.Lock() + c.closing = true + var database *sql.DB + var path string + if c.inFlight == 0 { + database, path = c.takeCleanupLocked() + } + c.mu.Unlock() + cleanupArtifactSnapshot(database, path) +} + +func (c *artifactSnapshotCursor) takeCleanupLocked() (*sql.DB, string) { + if c.cleaned { + return nil, "" + } + c.cleaned = true + database := c.db + c.db = nil + return database, c.path +} + +func cleanupArtifactSnapshot(database *sql.DB, path string) { + if database != nil { + _ = database.Close() + } + if path == "" { + return + } + for _, suffix := range []string{"", "-journal", "-wal", "-shm"} { + _ = os.Remove(path + suffix) + } +} + +func artifactCursorToken() (string, error) { + var token [24]byte + if _, err := rand.Read(token[:]); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(token[:]), nil +} + +func (r *artifactCursorRegistry) claim(token, scope string) (*artifactCursorLease, error) { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed { + return nil, fs.ErrClosed + } + lease, ok := r.tokens[token] + if !ok || lease.scope != scope || lease.released || lease.cursor == nil { + return nil, fmt.Errorf("%w: invalid or expired artifact cursor", artifact.ErrArtifactInvalid) + } + delete(r.tokens, token) + lease.token = "" + if lease.timer != nil { + lease.timer.Stop() + lease.timer = nil + } + return lease, nil +} + +func (r *artifactCursorRegistry) retain(lease *artifactCursorLease) (string, error) { + token, err := artifactCursorToken() + if err != nil { + return "", err + } + r.mu.Lock() + if r.closed { + r.mu.Unlock() + return "", fs.ErrClosed + } + if lease == nil || lease.released || lease.cursor == nil || lease.token != "" { + r.mu.Unlock() + return "", fs.ErrClosed + } + if _, ok := r.active[lease]; !ok { + r.mu.Unlock() + return "", fs.ErrClosed + } + r.tokens[token] = lease + lease.token = token + lease.timer = time.AfterFunc(artifactCursorTTL, func() { + r.release(token) + }) + r.mu.Unlock() + return token, nil +} + +func (r *artifactCursorRegistry) release(token string) bool { + r.mu.Lock() + lease, ok := r.tokens[token] + if !ok || lease.token != token || lease.released { + r.mu.Unlock() + return false + } + cursor, cancel := r.releaseLeaseLocked(lease) + r.mu.Unlock() + if cancel != nil { + cancel() + } + if cursor != nil { + cursor.close() + } + return true +} + +func (l *artifactCursorLease) release() { + if l == nil || l.registry == nil { + return + } + l.registry.mu.Lock() + if l.released { + l.registry.mu.Unlock() + return + } + cursor, cancel := l.registry.releaseLeaseLocked(l) + l.registry.mu.Unlock() + if cancel != nil { + cancel() + } + if cursor != nil { + cursor.close() + } +} + +func (r *artifactCursorRegistry) releaseLeaseLocked( + lease *artifactCursorLease, +) (*artifactSnapshotCursor, context.CancelFunc) { + lease.released = true + delete(r.active, lease) + if lease.token != "" { + delete(r.tokens, lease.token) + lease.token = "" + } + if lease.timer != nil { + lease.timer.Stop() + lease.timer = nil + } + cursor := lease.cursor + lease.cursor = nil + cancel := lease.constructCancel + lease.constructCancel = nil + return cursor, cancel +} + +func (l *artifactCursorLease) page( + ctx context.Context, limit int, +) ([]artifactSnapshotItem, bool, error) { + if l == nil || l.registry == nil { + return nil, false, fs.ErrClosed + } + l.registry.mu.Lock() + if l.released || l.cursor == nil { + l.registry.mu.Unlock() + return nil, false, fs.ErrClosed + } + cursor := l.cursor + l.registry.mu.Unlock() + return cursor.page(ctx, limit) +} + +func (r *artifactCursorRegistry) close() { + r.mu.Lock() + if r.closed { + r.mu.Unlock() + return + } + r.closed = true + cursors := make([]*artifactSnapshotCursor, 0, len(r.active)) + cancels := make([]context.CancelFunc, 0, len(r.active)) + for lease := range r.active { + cursor, cancel := r.releaseLeaseLocked(lease) + if cursor != nil { + cursors = append(cursors, cursor) + } + if cancel != nil { + cancels = append(cancels, cancel) + } + } + r.mu.Unlock() + for _, cancel := range cancels { + cancel() + } + for _, cursor := range cursors { + cursor.close() + } +} + +func finishArtifactSnapshotPage( + ctx context.Context, + registry *artifactCursorRegistry, + lease *artifactCursorLease, + limit int, +) ([]artifactSnapshotItem, string, error) { + items, more, err := lease.page(ctx, limit) + if err != nil { + lease.release() + return nil, "", err + } + if !more { + lease.release() + return items, "", nil + } + next, err := registry.retain(lease) + if err != nil { + lease.release() + return nil, "", err + } + return items, next, nil +} diff --git a/internal/server/artifact_cursors_internal_test.go b/internal/server/artifact_cursors_internal_test.go new file mode 100644 index 000000000..ab5253950 --- /dev/null +++ b/internal/server/artifact_cursors_internal_test.go @@ -0,0 +1,294 @@ +package server + +import ( + "context" + "errors" + "fmt" + "io/fs" + "path/filepath" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/mattn/go-sqlite3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" +) + +func TestArtifactCursorRegistryCapsConcurrentSnapshotConstruction(t *testing.T) { + root := t.TempDir() + registry := newArtifactCursorRegistry() + t.Cleanup(registry.close) + release := make(chan struct{}) + started := make(chan struct{}, maxArtifactCursors) + results := make(chan error, maxArtifactCursors) + var concurrent atomic.Int32 + var maximum atomic.Int32 + + for index := range maxArtifactCursors { + go func() { + lease, err := registry.buildSnapshot(t.Context(), root, + fmt.Sprintf("scope:%d", index), + func(context.Context, func(int, string) error) error { + current := concurrent.Add(1) + defer concurrent.Add(-1) + for { + observed := maximum.Load() + if current <= observed || maximum.CompareAndSwap(observed, current) { + break + } + } + started <- struct{}{} + <-release + return nil + }) + if lease != nil { + lease.release() + } + results <- err + }() + } + for range maxArtifactCursors { + <-started + } + + var overflowWork atomic.Bool + overflow, err := registry.buildSnapshot(t.Context(), root, "overflow", + func(context.Context, func(int, string) error) error { + overflowWork.Store(true) + return nil + }) + require.Nil(t, overflow) + require.Error(t, err) + assert.ErrorIs(t, err, artifact.ErrArtifactConflict) + assert.False(t, overflowWork.Load(), "capacity must be reserved before snapshot work") + assert.Equal(t, int32(maxArtifactCursors), maximum.Load()) + cursorFiles, globErr := filepath.Glob(filepath.Join(root, ".peer-cursor-*.sqlite")) + require.NoError(t, globErr) + assert.LessOrEqual(t, len(cursorFiles), maxArtifactCursors) + + close(release) + for range maxArtifactCursors { + require.NoError(t, <-results) + } + cursorFiles, globErr = filepath.Glob(filepath.Join(root, ".peer-cursor-*.sqlite")) + require.NoError(t, globErr) + assert.Empty(t, cursorFiles) +} + +func TestArtifactCursorRegistryClaimedSnapshotsRetainCapacity(t *testing.T) { + root := t.TempDir() + registry := newArtifactCursorRegistry() + t.Cleanup(registry.close) + claimed := make([]*artifactCursorLease, 0, maxArtifactCursors) + for index := range maxArtifactCursors { + scope := fmt.Sprintf("scope:%d", index) + lease, err := registry.buildSnapshot(t.Context(), root, scope, + func(_ context.Context, insert func(int, string) error) error { + if err := insert(0, "a"); err != nil { + return err + } + return insert(0, "b") + }) + require.NoError(t, err) + _, token, err := finishArtifactSnapshotPage(t.Context(), registry, lease, 1) + require.NoError(t, err) + require.NotEmpty(t, token) + lease, err = registry.claim(token, scope) + require.NoError(t, err) + claimed = append(claimed, lease) + } + + var overflowWork atomic.Bool + overflow, err := registry.buildSnapshot(t.Context(), root, "overflow", + func(context.Context, func(int, string) error) error { + overflowWork.Store(true) + return nil + }) + require.Nil(t, overflow) + require.Error(t, err) + assert.ErrorIs(t, err, artifact.ErrArtifactConflict) + assert.False(t, overflowWork.Load()) + + for _, lease := range claimed { + lease.release() + } + assertCursorReservationAvailable(t, registry, root) +} + +func TestArtifactCursorRegistryReleasesFailedAndCanceledConstruction(t *testing.T) { + for _, tt := range []struct { + name string + ctx func() context.Context + fill func(context.Context, func(int, string) error) error + want error + }{ + { + name: "fill failure", + ctx: t.Context, + fill: func(context.Context, func(int, string) error) error { + return errors.New("fill failed") + }, + want: errors.New("fill failed"), + }, + { + name: "caller cancellation", + ctx: func() context.Context { + ctx, cancel := context.WithCancel(t.Context()) + cancel() + return ctx + }, + fill: func(ctx context.Context, _ func(int, string) error) error { + return ctx.Err() + }, + want: context.Canceled, + }, + } { + t.Run(tt.name, func(t *testing.T) { + root := t.TempDir() + registry := newArtifactCursorRegistry() + t.Cleanup(registry.close) + ctx := tt.ctx() + lease, err := registry.buildSnapshot(ctx, root, "scope", tt.fill) + require.Nil(t, lease) + require.Error(t, err) + assert.ErrorContains(t, err, tt.want.Error()) + cursorFiles, globErr := filepath.Glob(filepath.Join(root, ".peer-cursor-*.sqlite")) + require.NoError(t, globErr) + assert.Empty(t, cursorFiles) + assertCursorReservationAvailable(t, registry, root) + }) + } +} + +func TestArtifactCursorRegistryShutdownCancelsConstructionAndRemovesSnapshot(t *testing.T) { + root := t.TempDir() + registry := newArtifactCursorRegistry() + started := make(chan struct{}) + result := make(chan error, 1) + go func() { + lease, err := registry.buildSnapshot(t.Context(), root, "scope", + func(ctx context.Context, _ func(int, string) error) error { + close(started) + <-ctx.Done() + return ctx.Err() + }) + if lease != nil { + lease.release() + } + result <- err + }() + <-started + + registry.close() + err := <-result + require.Error(t, err) + assert.True(t, errors.Is(err, context.Canceled) || errors.Is(err, fs.ErrClosed)) + cursorFiles, globErr := filepath.Glob(filepath.Join(root, ".peer-cursor-*.sqlite")) + require.NoError(t, globErr) + assert.Empty(t, cursorFiles) +} + +func TestArtifactCursorRegistryForcedShutdownDoesNotWaitForClaimedPage(t *testing.T) { + root := t.TempDir() + registry := newArtifactCursorRegistry() + lease, err := registry.buildSnapshot(t.Context(), root, "scope", + func(_ context.Context, insert func(int, string) error) error { + if err := insert(0, "a"); err != nil { + return err + } + return insert(0, "b") + }) + require.NoError(t, err) + _, token, err := finishArtifactSnapshotPage(t.Context(), registry, lease, 1) + require.NoError(t, err) + require.NotEmpty(t, token) + lease, err = registry.claim(token, "scope") + require.NoError(t, err) + require.NotNil(t, lease.cursor) + cursor := lease.cursor + cursorPath := cursor.path + cursor.db.SetMaxOpenConns(1) + queryEntered := make(chan struct{}) + releaseQuery := make(chan struct{}) + var enterOnce sync.Once + connection, err := cursor.db.Conn(t.Context()) + require.NoError(t, err) + err = connection.Raw(func(driverConnection any) error { + sqliteConnection, ok := driverConnection.(*sqlite3.SQLiteConn) + if !ok { + return fmt.Errorf("snapshot connection is %T, want *sqlite3.SQLiteConn", driverConnection) + } + return sqliteConnection.RegisterFunc("hold_cursor_page", func(name string) string { + enterOnce.Do(func() { close(queryEntered) }) + <-releaseQuery + return name + }, true) + }) + require.NoError(t, err) + require.NoError(t, connection.Close()) + _, err = cursor.db.ExecContext(t.Context(), `ALTER TABLE items RENAME TO base_items`) + require.NoError(t, err) + _, err = cursor.db.ExecContext(t.Context(), ` + CREATE VIEW items AS + SELECT kind, hold_cursor_page(name) AS name + FROM base_items`) + require.NoError(t, err) + + type pageResult struct { + err error + } + pageDone := make(chan pageResult, 1) + pageCtx, cancelPage := context.WithCancel(t.Context()) + defer cancelPage() + go func() { + _, _, pageErr := lease.page(pageCtx, 1) + pageDone <- pageResult{err: pageErr} + }() + select { + case <-queryEntered: + case <-time.After(time.Second): + require.FailNow(t, "claimed page never entered the SQLite query") + } + require.Equal(t, 1, cursor.db.Stats().InUse) + + closeDone := make(chan struct{}) + go func() { + registry.close() + close(closeDone) + }() + closedPromptly := false + select { + case <-closeDone: + closedPromptly = true + case <-time.After(200 * time.Millisecond): + } + assert.FileExists(t, cursorPath, + "snapshot cleanup must wait until the active page releases its pin") + cancelPage() + close(releaseQuery) + result := <-pageDone + <-closeDone + + assert.True(t, closedPromptly, "forced shutdown must not wait for an active page query") + assert.ErrorIs(t, result.err, context.Canceled) + assert.ErrorIs(t, result.err, fs.ErrClosed) + assert.NoFileExists(t, cursorPath) + assert.True(t, cursor.cleaned) + assert.Nil(t, cursor.db) + registry.close() +} + +func assertCursorReservationAvailable( + t *testing.T, registry *artifactCursorRegistry, root string, +) { + t.Helper() + lease, err := registry.buildSnapshot(t.Context(), root, "available", + func(context.Context, func(int, string) error) error { return nil }) + require.NoError(t, err) + require.NotNil(t, lease) + lease.release() +} diff --git a/internal/server/artifact_http_transport_test.go b/internal/server/artifact_http_transport_test.go new file mode 100644 index 000000000..33a2a5f3b --- /dev/null +++ b/internal/server/artifact_http_transport_test.go @@ -0,0 +1,267 @@ +package server_test + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/dbtest" +) + +func TestArtifactHTTPTransportCanceledExchangeReleasesAllServerCursorsPromptly(t *testing.T) { + const ( + peerOrigin = "peer-0000-a1b2c3" + token = "secret" + deleteWait = 3 * time.Second + ) + peer := httptest.NewUnstartedServer(nil) + peerURL := "http://" + peer.Listener.Addr().String() + te := setupArtifact(t, + withAuth(token), + withArtifactOrigin("zzserver-a1b2c3"), + withPublicURL(peerURL), + ) + for index := range 513 { + origin := fmt.Sprintf("peer-%04d-a1b2c3", index) + seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, + fmt.Appendf(nil, "origin-%04d", index)) + } + for index := range 513 { + seedArtifactStore(t, te.artifactStore, peerOrigin, artifact.KindRaw, + fmt.Appendf(nil, "raw-%04d", index)) + } + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + var deletes, canceledDeletes atomic.Int32 + peer.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet && strings.Contains(r.URL.Path, "/raw/") { + cancel() + <-r.Context().Done() + return + } + if r.Method == http.MethodDelete && strings.Contains(r.URL.Path, "/cursors/") { + response := httptest.NewRecorder() + te.srv.Handler().ServeHTTP(response, r) + deletes.Add(1) + select { + case <-r.Context().Done(): + canceledDeletes.Add(1) + case <-time.After(deleteWait): + } + for key, values := range response.Header() { + w.Header()[key] = append([]string(nil), values...) + } + w.WriteHeader(response.Code) + _, _ = io.Copy(w, response.Body) + return + } + te.srv.Handler().ServeHTTP(w, r) + }) + peer.Start() + defer peer.Close() + + clientDB, clientDir := newClientNode(t, "sess-1", "alpha") + _, err := artifact.Sync(ctx, clientDB, artifact.SyncOptions{ + DataDir: clientDir, + Target: peerURL, + Origin: "zzclient-a1b2c3", + Token: token, + }) + require.Error(t, err) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, int32(2), deletes.Load(), "origin and index cursors must both be released") + assert.Equal(t, int32(2), canceledDeletes.Load(), + "each cursor cleanup request must use its bounded context instead of the 120-second peer timeout") +} + +func TestArtifactHTTPTransportPreExchangeFailuresReleasePreparedCursor(t *testing.T) { + const ( + attempts = 3 + token = "secret" + ) + peer := httptest.NewUnstartedServer(nil) + peerURL := "http://" + peer.Listener.Addr().String() + te := setupArtifact(t, withAuth(token), withArtifactOrigin("server-a1b2c3"), withPublicURL(peerURL)) + for index := range 513 { + origin := fmt.Sprintf("peer-%04d-a1b2c3", index) + seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, + fmt.Appendf(nil, "origin-%04d", index)) + } + + var deletes atomic.Int32 + peer.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodDelete && strings.Contains(r.URL.Path, "/cursors/") { + deletes.Add(1) + } + te.srv.Handler().ServeHTTP(w, r) + }) + peer.Start() + defer peer.Close() + clientDB, clientDir := newClientNode(t, "sess-1", "alpha") + + for range attempts { + _, err := artifact.Sync(t.Context(), clientDB, artifact.SyncOptions{ + DataDir: clientDir, + Target: peerURL, + Origin: "invalid/origin", + Token: token, + }) + require.Error(t, err) + assert.ErrorContains(t, err, "invalid artifact origin") + } + + assert.Equal(t, int32(attempts), deletes.Load(), + "each failed sync must release its prepared origin cursor") +} + +// newClientNode opens a client-only artifact store (a db plus data dir) that +// drives artifact.Sync against an HTTP peer. +func newClientNode(t *testing.T, sessionID, project string) (*db.DB, string) { + t.Helper() + dir := t.TempDir() + database, err := db.Open(filepath.Join(dir, "client.db")) + require.NoError(t, err) + t.Cleanup(func() { database.Close() }) + dbtest.SeedSession(t, database, sessionID, project, func(s *db.Session) { + s.MessageCount = 2 + s.UserMessageCount = 1 + }) + require.NoError(t, database.ReplaceSessionMessages(sessionID, []db.Message{ + {SessionID: sessionID, Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + {SessionID: sessionID, Ordinal: 1, Role: "assistant", Content: "world", ContentLength: 5}, + })) + return database, dir +} + +func TestArtifactHTTPTransportSyncsSessionsAndMetadata(t *testing.T) { + ctx := context.Background() + const token = "secret" + const aOrigin = "laptop-a1b2c3" + + // Node B: a real server exposing the artifact peer API behind auth. + te := setupArtifact(t, withAuth(token), withArtifactOrigin("desktop-d4e5f6")) + peer := httptest.NewServer(te.srv.Handler()) + defer peer.Close() + + aDB, aDir := newClientNode(t, "sess-1", "alpha") + + syncToPeer := func() { + _, err := artifact.Sync(ctx, aDB, artifact.SyncOptions{ + DataDir: aDir, + Target: peer.URL, + Origin: aOrigin, + Token: token, + }) + require.NoError(t, err) + } + + // A pushes its session over HTTP; B imports it on receipt. + syncToPeer() + importedID := aOrigin + "~sess-1" + gotB, err := te.db.GetSession(ctx, importedID) + require.NoError(t, err) + require.NotNil(t, gotB, "peer should import the pushed session") + assert.Equal(t, "alpha", gotB.Project) + + // A renames the session and syncs again; the metadata event is enumerated + // via the index route, posted, and replayed on B. + display := "Renamed on A" + require.NoError(t, aDB.RenameSession("sess-1", &display)) + repository, err := artifact.OpenRepository(ctx, aDir) + require.NoError(t, err) + recorder := artifact.NewMetadataRecorder(aDB, artifact.MetadataRecorderOptions{ + Store: repository.Content(), + Origin: aOrigin, + }) + value, err := json.Marshal(struct { + DisplayName *string `json:"display_name"` + }{DisplayName: &display}) + require.NoError(t, err) + _, err = recorder.Append(ctx, artifact.MetadataEventInput{ + SessionID: "sess-1", + Op: artifact.MetadataOpRename, + Value: value, + }) + require.NoError(t, err) + require.NoError(t, repository.Close()) + + syncToPeer() + gotB, err = te.db.GetSession(ctx, importedID) + require.NoError(t, err) + require.NotNil(t, gotB) + require.NotNil(t, gotB.DisplayName) + assert.Equal(t, display, *gotB.DisplayName) +} + +func TestArtifactHTTPTransportPullsRemoteSessions(t *testing.T) { + ctx := context.Background() + const token = "secret" + const bOrigin = "desktop-d4e5f6" + + // Node B owns a session but has not run a separate artifact publisher. + te := setupArtifact(t, withAuth(token), withArtifactOrigin(bOrigin)) + peer := httptest.NewServer(te.srv.Handler()) + defer peer.Close() + dbtest.SeedSession(t, te.db, "remote-1", "bravo", func(s *db.Session) { + s.MessageCount = 2 + s.UserMessageCount = 1 + }) + require.NoError(t, te.db.ReplaceSessionMessages("remote-1", []db.Message{ + {SessionID: "remote-1", Ordinal: 0, Role: "user", Content: "ping", ContentLength: 4}, + {SessionID: "remote-1", Ordinal: 1, Role: "assistant", Content: "pong", ContentLength: 4}, + })) + displayName := "Renamed before HTTP publishing" + require.NoError(t, te.db.RenameSession("remote-1", &displayName)) + + aDB, aDir := newClientNode(t, "sess-1", "alpha") + syncResult, err := artifact.Sync(ctx, aDB, artifact.SyncOptions{ + DataDir: aDir, + Target: peer.URL, + Origin: "laptop-a1b2c3", + Token: token, + }) + require.NoError(t, err) + assert.Equal(t, 1, syncResult.ImportedSessions) + assert.Equal(t, 1, syncResult.ImportedMetadata) + pendingImports, err := aDB.PendingArtifactImports(ctx, 1, 10) + require.NoError(t, err) + assert.Empty(t, pendingImports) + + // A pulled B's session. + gotA, err := aDB.GetSession(ctx, bOrigin+"~remote-1") + require.NoError(t, err) + require.NotNil(t, gotA, "client should pull the remote session") + assert.Equal(t, "bravo", gotA.Project) + require.NotNil(t, gotA.DisplayName) + assert.Equal(t, displayName, *gotA.DisplayName) +} + +func TestArtifactHTTPTransportRejectsBadToken(t *testing.T) { + te := setupArtifact(t, withAuth("secret"), withArtifactOrigin("desktop-d4e5f6")) + peer := httptest.NewServer(te.srv.Handler()) + defer peer.Close() + aDB, aDir := newClientNode(t, "sess-1", "alpha") + + _, err := artifact.Sync(context.Background(), aDB, artifact.SyncOptions{ + DataDir: aDir, + Target: peer.URL, + Origin: "laptop-a1b2c3", + Token: "wrong", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "peer") +} diff --git a/internal/server/artifact_iterator_external_test.go b/internal/server/artifact_iterator_external_test.go new file mode 100644 index 000000000..06558bdfc --- /dev/null +++ b/internal/server/artifact_iterator_external_test.go @@ -0,0 +1,45 @@ +package server_test + +import ( + "errors" + "io" + "testing" + + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" +) + +func collectArtifactOrigins(t *testing.T, store artifact.ArtifactStore, limit int) []string { + t.Helper() + iterator, err := store.Origins(t.Context()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, iterator.Close()) }) + var result []string + for { + items, nextErr := iterator.Next(t.Context(), limit) + require.True(t, nextErr == nil || errors.Is(nextErr, io.EOF)) + result = append(result, items...) + if errors.Is(nextErr, io.EOF) { + return result + } + } +} + +func collectArtifactEntries( + t *testing.T, store artifact.ArtifactStore, origin string, kind artifact.Kind, limit int, +) []artifact.Entry { + t.Helper() + iterator, err := store.Entries(t.Context(), origin, kind) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, iterator.Close()) }) + var result []artifact.Entry + for { + items, nextErr := iterator.Next(t.Context(), limit) + require.True(t, nextErr == nil || errors.Is(nextErr, io.EOF)) + result = append(result, items...) + if errors.Is(nextErr, io.EOF) { + return result + } + } +} diff --git a/internal/server/artifact_iterator_test.go b/internal/server/artifact_iterator_test.go new file mode 100644 index 000000000..cf8fa6055 --- /dev/null +++ b/internal/server/artifact_iterator_test.go @@ -0,0 +1,29 @@ +package server + +import ( + "errors" + "io" + "testing" + + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" +) + +func collectArtifactEntries( + t *testing.T, store artifact.ArtifactStore, origin string, kind artifact.Kind, limit int, +) []artifact.Entry { + t.Helper() + iterator, err := store.Entries(t.Context(), origin, kind) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, iterator.Close()) }) + var result []artifact.Entry + for { + items, nextErr := iterator.Next(t.Context(), limit) + require.True(t, nextErr == nil || errors.Is(nextErr, io.EOF)) + result = append(result, items...) + if errors.Is(nextErr, io.EOF) { + return result + } + } +} diff --git a/internal/server/artifact_lifecycle_internal_test.go b/internal/server/artifact_lifecycle_internal_test.go new file mode 100644 index 000000000..792429ca3 --- /dev/null +++ b/internal/server/artifact_lifecycle_internal_test.go @@ -0,0 +1,1756 @@ +package server + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/dbtest" +) + +const artifactLifecycleOrigin = "lifecycle-a1b2c3" + +type lifecycleArtifactStore struct { + mu sync.Mutex + entries map[artifact.Ref][]byte + corrupt map[artifact.Ref]bool + openStarted chan struct{} + openRelease chan struct{} + createStarted chan struct{} + createRelease chan struct{} + closeCalls atomic.Int32 + closed atomic.Bool + closeErr error +} + +type repairGateArtifactStore struct { + artifact.ArtifactStore + started chan struct{} + err error +} + +func (s *lifecycleArtifactStore) RepairContent( + context.Context, artifact.Identity, io.Reader, +) error { + return artifact.ErrArtifactUnsupported +} + +func (s *repairGateArtifactStore) RepairContent( + ctx context.Context, _ artifact.Identity, _ io.Reader, +) error { + if s.started != nil { + select { + case s.started <- struct{}{}: + default: + } + } + if s.err != nil { + return s.err + } + <-ctx.Done() + return ctx.Err() +} + +func newLifecycleArtifactStore() *lifecycleArtifactStore { + return &lifecycleArtifactStore{ + entries: make(map[artifact.Ref][]byte), corrupt: make(map[artifact.Ref]bool), + } +} + +func (s *lifecycleArtifactStore) Create( + ctx context.Context, + ref artifact.Ref, + identity artifact.Identity, + _ string, + body io.Reader, +) (artifact.CreateResult, error) { + if s.createStarted != nil { + select { + case s.createStarted <- struct{}{}: + default: + } + select { + case <-s.createRelease: + case <-ctx.Done(): + return artifact.CreateResult{}, ctx.Err() + } + } + data, err := io.ReadAll(body) + if err != nil { + return artifact.CreateResult{}, err + } + if err := ctx.Err(); err != nil { + return artifact.CreateResult{}, err + } + hash := sha256.Sum256(data) + if hex.EncodeToString(hash[:]) != identity.SHA256 || int64(len(data)) != identity.Size { + return artifact.CreateResult{}, artifact.ErrArtifactInvalid + } + s.mu.Lock() + defer s.mu.Unlock() + if existing, ok := s.entries[ref]; ok { + if !bytes.Equal(existing, data) { + return artifact.CreateResult{}, artifact.ErrArtifactConflict + } + return artifact.CreateResult{Entry: lifecycleEntry(ref, data)}, nil + } + s.entries[ref] = append([]byte(nil), data...) + return artifact.CreateResult{Entry: lifecycleEntry(ref, data), Created: true}, nil +} + +func (s *lifecycleArtifactStore) Stat( + ctx context.Context, ref artifact.Ref, +) (artifact.Entry, error) { + if err := ctx.Err(); err != nil { + return artifact.Entry{}, err + } + s.mu.Lock() + defer s.mu.Unlock() + data, ok := s.entries[ref] + if !ok { + return artifact.Entry{}, artifact.ErrArtifactNotFound + } + return lifecycleEntry(ref, data), nil +} + +func (s *lifecycleArtifactStore) Open( + ctx context.Context, ref artifact.Ref, +) (artifact.Entry, artifact.VerifiedReader, error) { + if s.openStarted != nil { + select { + case s.openStarted <- struct{}{}: + default: + } + select { + case <-s.openRelease: + case <-ctx.Done(): + return artifact.Entry{}, nil, ctx.Err() + } + } + entry, err := s.Stat(ctx, ref) + if err != nil { + return artifact.Entry{}, nil, err + } + s.mu.Lock() + data := append([]byte(nil), s.entries[ref]...) + corrupt := s.corrupt[ref] + s.mu.Unlock() + return entry, &lifecycleVerifiedReader{ + Reader: bytes.NewReader(data), corrupt: corrupt, + }, nil +} + +func (s *lifecycleArtifactStore) ListOrigins( + ctx context.Context, cursor artifact.Cursor, limit int, +) ([]string, artifact.Cursor, error) { + if err := ctx.Err(); err != nil { + return nil, "", err + } + if cursor != "" || limit < 1 { + return nil, "", artifact.ErrArtifactInvalid + } + s.mu.Lock() + defer s.mu.Unlock() + for ref := range s.entries { + return []string{ref.Origin}, "", nil + } + return []string{}, "", nil +} + +func (s *lifecycleArtifactStore) List( + ctx context.Context, + origin string, + kind artifact.Kind, + cursor artifact.Cursor, + limit int, +) (artifact.Page, error) { + if err := ctx.Err(); err != nil { + return artifact.Page{}, err + } + if cursor != "" || limit < 1 { + return artifact.Page{}, artifact.ErrArtifactInvalid + } + s.mu.Lock() + defer s.mu.Unlock() + page := artifact.Page{} + for ref, data := range s.entries { + if ref.Origin == origin && ref.Kind == kind { + page.Items = append(page.Items, lifecycleEntry(ref, data)) + } + } + return page, nil +} + +func (s *lifecycleArtifactStore) Origins(context.Context) (artifact.OriginIterator, error) { + s.mu.Lock() + defer s.mu.Unlock() + seen := make(map[string]struct{}) + for ref := range s.entries { + seen[ref.Origin] = struct{}{} + } + items := make([]string, 0, len(seen)) + for origin := range seen { + items = append(items, origin) + } + sort.Strings(items) + return &lifecycleOriginIterator{items: items}, nil +} + +func (s *lifecycleArtifactStore) Entries( + _ context.Context, origin string, kind artifact.Kind, +) (artifact.EntryIterator, error) { + s.mu.Lock() + defer s.mu.Unlock() + items := make([]artifact.Entry, 0) + for ref, data := range s.entries { + if ref.Origin == origin && ref.Kind == kind { + items = append(items, lifecycleEntry(ref, data)) + } + } + sort.Slice(items, func(left, right int) bool { + return items[left].Ref.Name < items[right].Ref.Name + }) + return &lifecycleEntryIterator{items: items}, nil +} + +type lifecycleOriginIterator struct { + items []string + offset int + closed bool +} + +func (i *lifecycleOriginIterator) Next(ctx context.Context, limit int) ([]string, error) { + if i.closed { + return nil, os.ErrClosed + } + if err := ctx.Err(); err != nil { + return nil, err + } + if limit < 1 { + return nil, artifact.ErrArtifactInvalid + } + end := min(i.offset+limit, len(i.items)) + page := append([]string(nil), i.items[i.offset:end]...) + i.offset = end + if end == len(i.items) { + return page, io.EOF + } + return page, nil +} + +func (i *lifecycleOriginIterator) Close() error { + i.closed = true + return nil +} + +type lifecycleEntryIterator struct { + items []artifact.Entry + offset int + closed bool +} + +func (i *lifecycleEntryIterator) Next( + ctx context.Context, limit int, +) ([]artifact.Entry, error) { + if i.closed { + return nil, os.ErrClosed + } + if err := ctx.Err(); err != nil { + return nil, err + } + if limit < 1 { + return nil, artifact.ErrArtifactInvalid + } + end := min(i.offset+limit, len(i.items)) + page := append([]artifact.Entry(nil), i.items[i.offset:end]...) + i.offset = end + if end == len(i.items) { + return page, io.EOF + } + return page, nil +} + +func (i *lifecycleEntryIterator) Close() error { + i.closed = true + return nil +} + +func (*lifecycleArtifactStore) Quarantine(context.Context, artifact.Ref, string) error { + return nil +} + +func (*lifecycleArtifactStore) Trash(context.Context, artifact.Ref) error { return nil } + +func (*lifecycleArtifactStore) Pack(context.Context, int64) (artifact.PackResult, error) { + return artifact.PackResult{}, nil +} + +func (*lifecycleArtifactStore) LooseBacklog(context.Context) (artifact.LooseBacklog, error) { + return artifact.LooseBacklog{}, nil +} + +func (s *lifecycleArtifactStore) Close() error { + s.closeCalls.Add(1) + s.closed.Store(true) + return s.closeErr +} + +func lifecycleEntry(ref artifact.Ref, data []byte) artifact.Entry { + hash := sha256.Sum256(data) + return artifact.Entry{Ref: ref, Identity: artifact.Identity{ + SHA256: hex.EncodeToString(hash[:]), + Size: int64(len(data)), + }} +} + +type lifecycleVerifiedReader struct { + *bytes.Reader + closed bool + corrupt bool +} + +func (r *lifecycleVerifiedReader) Verify() error { + if r.closed { + return errors.New("reader closed") + } + _, err := r.Seek(0, io.SeekEnd) + if err != nil { + return err + } + if r.corrupt { + return artifact.ErrArtifactCorrupt + } + return nil +} + +func (r *lifecycleVerifiedReader) Close() error { + r.closed = true + return nil +} + +func TestArtifactResetDrainsActiveOperationsBeforeStoreSwap(t *testing.T) { + oldStore := newLifecycleArtifactStore() + newStore := newLifecycleArtifactStore() + lifetime := artifactOperationLifetime{store: oldStore} + store, release, err := lifetime.acquire() + require.NoError(t, err) + assert.Same(t, oldStore, store) + type resetDrainResult struct { + store artifact.ArtifactStore + err error + } + drained := make(chan resetDrainResult, 1) + go func() { + owned, _, drainErr := lifetime.beginReset(t.Context()) + drained <- resetDrainResult{store: owned, err: drainErr} + }() + + require.Eventually(t, func() bool { + _, probeRelease, probeErr := lifetime.acquire() + if probeErr != nil { + return true + } + probeRelease() + return false + }, time.Second, time.Millisecond, "reset never entered its drain state") + select { + case <-drained: + t.Fatal("reset passed an active artifact operation") + default: + } + release() + drainResult := <-drained + require.NoError(t, drainResult.err) + assert.Same(t, oldStore, drainResult.store) + require.NoError(t, lifetime.finishReset(newStore)) + + store, release, err = lifetime.acquire() + require.NoError(t, err) + assert.Same(t, newStore, store) + release() +} + +func TestArtifactResetCanceledWhileDrainingUnwindsResetState(t *testing.T) { + store := newLifecycleArtifactStore() + lifetime := artifactOperationLifetime{store: store} + _, release, err := lifetime.acquire() + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + drained := make(chan error, 1) + go func() { + _, _, drainErr := lifetime.beginReset(ctx) + drained <- drainErr + }() + + require.Eventually(t, func() bool { + _, probeRelease, probeErr := lifetime.acquire() + if probeErr != nil { + return true + } + probeRelease() + return false + }, time.Second, time.Millisecond, "reset never entered its drain state") + cancel() + require.ErrorIs(t, <-drained, context.Canceled) + release() + + acquired, releaseAgain, err := lifetime.acquire() + require.NoError(t, err) + assert.Same(t, store, acquired) + releaseAgain() +} + +func TestArtifactResetShutdownWhileDrainingPreventsLateReset(t *testing.T) { + store := newLifecycleArtifactStore() + lifetime := artifactOperationLifetime{store: store} + _, release, err := lifetime.acquire() + require.NoError(t, err) + type resetDrainResult struct { + store artifact.ArtifactStore + err error + } + drained := make(chan resetDrainResult, 1) + go func() { + owned, _, drainErr := lifetime.beginReset(t.Context()) + drained <- resetDrainResult{store: owned, err: drainErr} + }() + + require.Eventually(t, func() bool { + _, probeRelease, probeErr := lifetime.acquire() + if probeErr != nil { + return true + } + probeRelease() + return false + }, time.Second, time.Millisecond, "reset never entered its drain state") + require.NoError(t, lifetime.closeWhenIdle()) + release() + drainResult := <-drained + require.ErrorContains(t, drainResult.err, "closing") + assert.Nil(t, drainResult.store) + assert.Equal(t, int32(1), store.closeCalls.Load()) + _, _, err = lifetime.acquire() + require.Error(t, err) +} + +func TestArtifactResetShutdownAfterAdmissionBeforeMutationDoesNotMoveVault(t *testing.T) { + dataDir := t.TempDir() + database := dbtest.OpenTestDBAt(t, filepath.Join(dataDir, "sessions.db")) + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, WriteTimeout: time.Second, + AuthToken: "daemon-secret", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactRepository(repository)) + locked := true + srv.sessionLifecycleMu.Lock() + t.Cleanup(func() { + if locked { + srv.sessionLifecycleMu.Unlock() + } + require.NoError(t, srv.Shutdown(t.Context())) + require.NoError(t, repository.Close()) + }) + admitted := make(chan struct{}) + var admittedOnce sync.Once + srv.beforeSessionLifecycleLock = func() { + admittedOnce.Do(func() { close(admitted) }) + } + + response := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/reset", http.NoBody) + request.Host = "127.0.0.1:0" + request.RemoteAddr = "127.0.0.1:1234" + request.Header.Set("Authorization", "Bearer daemon-secret") + done := make(chan struct{}) + go func() { + srv.Handler().ServeHTTP(response, request) + close(done) + }() + + select { + case <-admitted: + case <-time.After(time.Second): + require.FailNow(t, "reset did not reach the admitted session-lifecycle boundary") + } + require.NoError(t, srv.closeArtifactResources()) + srv.sessionLifecycleMu.Unlock() + locked = false + select { + case <-done: + case <-time.After(time.Second): + require.FailNow(t, "reset did not unwind after artifact shutdown") + } + + assert.Equal(t, http.StatusServiceUnavailable, response.Code, response.Body.String()) + assert.DirExists(t, filepath.Join(dataDir, "artifacts")) + moved, err := filepath.Glob(filepath.Join(dataDir, "artifacts.reset-*")) + require.NoError(t, err) + assert.Empty(t, moved, "shutdown must prevent a reset admitted before session drain from moving the vault") +} + +func TestArtifactResetShutdownDeadlineIsBoundedDuringPostMoveRepublish(t *testing.T) { + dataDir := t.TempDir() + database := dbtest.OpenTestDBAt(t, filepath.Join(dataDir, "sessions.db")) + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, WriteTimeout: time.Second, + AuthToken: "daemon-secret", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactRepository(repository)) + republishStarted := make(chan struct{}) + releaseRepublish := make(chan struct{}) + srv.republishArtifactRepositoryReset = func( + ctx context.Context, + dataDir string, + database *db.DB, + _ string, + _ *artifact.Repository, + result artifact.RepositoryResetResult, + ) (artifact.RepositoryResetResult, error) { + close(republishStarted) + <-releaseRepublish + return result, nil + } + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + srv.SetPort(listener.Addr().(*net.TCPAddr).Port) + serveDone := make(chan error, 1) + go func() { serveDone <- srv.Serve(listener) }() + t.Cleanup(func() { + require.NoError(t, srv.Shutdown(t.Context())) + require.NoError(t, repository.Close()) + }) + + requestDone := make(chan error, 1) + go func() { + request, requestErr := http.NewRequest( + http.MethodPost, + "http://"+listener.Addr().String()+"/api/v1/artifacts/reset", + http.NoBody, + ) + if requestErr != nil { + requestDone <- requestErr + return + } + request.Header.Set("Authorization", "Bearer daemon-secret") + response, requestErr := http.DefaultClient.Do(request) + if response != nil { + _ = response.Body.Close() + } + requestDone <- requestErr + }() + select { + case <-republishStarted: + case <-time.After(time.Second): + require.FailNow(t, "reset did not reach post-move republish") + } + + shutdownCtx, cancel := context.WithTimeout(t.Context(), 10*time.Millisecond) + defer cancel() + shutdownDone := make(chan error, 1) + go func() { shutdownDone <- srv.Shutdown(shutdownCtx) }() + var shutdownErr error + bounded := false + select { + case shutdownErr = <-shutdownDone: + bounded = true + case <-time.After(250 * time.Millisecond): + assert.Fail(t, "shutdown exceeded its deadline while reset republish was blocked") + } + close(releaseRepublish) + if !bounded { + shutdownErr = <-shutdownDone + } + assert.ErrorIs(t, shutdownErr, context.DeadlineExceeded) + require.NoError(t, <-requestDone) + assert.ErrorIs(t, <-serveDone, http.ErrServerClosed) + assert.DirExists(t, filepath.Join(dataDir, "artifacts")) + moved, err := filepath.Glob(filepath.Join(dataDir, "artifacts.reset-*")) + require.NoError(t, err) + assert.Len(t, moved, 1, "post-admission reset must preserve its moved-aside diagnostic vault") +} + +func TestArtifactResetRequestCancellationAfterMoveKeepsFreshRepositoryUsable(t *testing.T) { + dataDir := t.TempDir() + database := dbtest.OpenTestDBAt(t, filepath.Join(dataDir, "sessions.db")) + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, WriteTimeout: time.Second, + AuthToken: "daemon-secret", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactRepository(repository)) + republishStarted := make(chan struct{}) + releaseRepublish := make(chan struct{}) + var freshRepository *artifact.Repository + srv.republishArtifactRepositoryReset = func( + ctx context.Context, + dataDir string, + database *db.DB, + _ string, + fresh *artifact.Repository, + result artifact.RepositoryResetResult, + ) (artifact.RepositoryResetResult, error) { + freshRepository = fresh + close(republishStarted) + <-releaseRepublish + return result, ctx.Err() + } + t.Cleanup(func() { + require.NoError(t, srv.Shutdown(t.Context())) + require.NoError(t, repository.Close()) + if freshRepository != nil { + require.NoError(t, freshRepository.Close()) + } + }) + + ctx, cancel := context.WithCancel(t.Context()) + response := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/reset", http.NoBody).WithContext(ctx) + request.Host = "127.0.0.1:0" + request.RemoteAddr = "127.0.0.1:1234" + request.Header.Set("Authorization", "Bearer daemon-secret") + done := make(chan struct{}) + go func() { + srv.Handler().ServeHTTP(response, request) + close(done) + }() + select { + case <-republishStarted: + case <-time.After(time.Second): + require.FailNow(t, "reset did not reach post-move republish") + } + cancel() + close(releaseRepublish) + select { + case <-done: + case <-time.After(time.Second): + require.FailNow(t, "canceled reset request did not return") + } + assert.Equal(t, http.StatusInternalServerError, response.Code, response.Body.String()) + + store, release, err := srv.acquireArtifactStore() + require.NoError(t, err, "the fresh post-move repository must remain owned by the server") + defer release() + assert.Same(t, freshRepository, srv.artifactRepository) + assert.False(t, srv.artifactRepository.Closed()) + _, err = store.Origins(t.Context()) + require.NoError(t, err) + _, pending, err := database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.True(t, pending, "canceled republish must retain durable recovery authority") + require.NoError(t, srv.publishLocalArtifacts(t.Context(), store), + "the retained fresh repository must support a later full publication") + assert.True(t, srv.artifactBaselineDone) + _, pending, err = database.ArtifactResetRepublishPending(t.Context()) + require.NoError(t, err) + assert.False(t, pending, "ordinary publication must finish reset recovery") +} + +func TestArtifactResetDrainBlocksCursorReleaseRegistryAccess(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := newLifecycleArtifactStore() + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + owned, _, err := srv.artifactOps.beginReset(t.Context()) + require.NoError(t, err) + + response := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodDelete, "/api/v1/artifacts/cursors/stale", nil) + request.Host = "127.0.0.1:0" + request.RemoteAddr = "127.0.0.1:1234" + request.Header.Set("Origin", "http://127.0.0.1:0") + srv.Handler().ServeHTTP(response, request) + assert.Equal(t, http.StatusServiceUnavailable, response.Code) + require.NoError(t, srv.artifactOps.finishReset(owned)) +} + +func TestArtifactLifecycleRoutesUseInjectedStoreAndCloseItOnce(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := newLifecycleArtifactStore() + body := []byte("injected artifact bytes") + hash := sha256.Sum256(body) + name := hex.EncodeToString(hash[:]) + ref, err := artifact.NewRef(artifactLifecycleOrigin, artifact.KindRaw, name) + require.NoError(t, err) + store.entries[ref] = body + + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: t.TempDir(), WriteTimeout: time.Second, + }, database, nil, WithArtifactStore(store)) + + for _, path := range []string{ + "/api/v1/artifacts/origins", + "/api/v1/artifacts/" + artifactLifecycleOrigin + "/index", + "/api/v1/artifacts/" + artifactLifecycleOrigin + "/raw/" + name, + } { + response := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, path, nil) + request.Host = "127.0.0.1:0" + srv.Handler().ServeHTTP(response, request) + assert.Equal(t, http.StatusOK, response.Code, path) + } + + require.NoError(t, srv.Shutdown(t.Context())) + require.NoError(t, srv.Shutdown(t.Context())) + assert.Equal(t, int32(1), store.closeCalls.Load()) +} + +func TestArtifactLifecycleReadOnlyAndRemoteServersOmitMutationRoutes(t *testing.T) { + dir := t.TempDir() + writable, err := db.Open(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + require.NoError(t, writable.Close()) + readonly, err := db.OpenReadOnly(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, readonly.Close()) }) + + for name, database := range map[string]db.Store{ + "read-only SQLite": readonly, + "remote": lifecycleRemoteStore{}, + } { + t.Run(name, func(t *testing.T) { + store := newLifecycleArtifactStore() + srv := New(config.Config{Host: "127.0.0.1", DataDir: dir}, database, nil, + WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + for _, path := range []string{ + "/api/v1/artifacts/finalize", + "/api/v1/artifacts/exchange", + "/api/v1/artifacts/maintenance", + "/api/v1/artifacts/reset", + "/api/v1/artifacts/" + artifactLifecycleOrigin + "/raw/" + strings.Repeat("0", 64), + } { + response := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(nil)) + request.Host = "127.0.0.1:0" + request.Header.Set("Origin", "http://127.0.0.1:0") + srv.Handler().ServeHTTP(response, request) + assert.Contains(t, []int{http.StatusNotFound, http.StatusMethodNotAllowed}, response.Code, path) + } + }) + } +} + +func TestArtifactResetRequiresAuthenticatedDirectLoopback(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dataDir := t.TempDir() + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, WriteTimeout: time.Second, + AuthToken: "daemon-secret", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactRepository(repository)) + t.Cleanup(func() { + require.NoError(t, srv.Shutdown(t.Context())) + require.NoError(t, repository.Close()) + }) + request := func(remoteAddr string, authenticated, forwarded bool) *http.Request { + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/reset", http.NoBody) + req.Host = "127.0.0.1:0" + req.RemoteAddr = remoteAddr + if authenticated { + req.Header.Set("Authorization", "Bearer daemon-secret") + } + if forwarded { + req.Header.Set("X-Forwarded-For", "203.0.113.10") + } + return req + } + + unauthorized := httptest.NewRecorder() + srv.Handler().ServeHTTP(unauthorized, request("127.0.0.1:1234", false, false)) + assert.Equal(t, http.StatusUnauthorized, unauthorized.Code) + proxied := httptest.NewRecorder() + srv.Handler().ServeHTTP(proxied, request("127.0.0.1:1234", true, true)) + assert.Equal(t, http.StatusForbidden, proxied.Code) + remote := httptest.NewRecorder() + srv.Handler().ServeHTTP(remote, request("203.0.113.10:1234", true, false)) + assert.Equal(t, http.StatusForbidden, remote.Code) + local := httptest.NewRecorder() + srv.Handler().ServeHTTP(local, request("127.0.0.1:1234", true, false)) + assert.Equal(t, http.StatusOK, local.Code, local.Body.String()) +} + +func TestArtifactLifecycleExchangeRequiresAuthenticatedDirectLoopback(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dataDir := t.TempDir() + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, WriteTimeout: time.Second, + AuthToken: "daemon-secret", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactRepository(repository)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + target := t.TempDir() + body, err := json.Marshal(artifactExchangeRequest{Target: target}) + require.NoError(t, err) + request := func(remoteAddr string, authenticated, forwarded bool) *http.Request { + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/exchange", bytes.NewReader(body)) + req.Host = "127.0.0.1:0" + req.RemoteAddr = remoteAddr + if authenticated { + req.Header.Set("Authorization", "Bearer daemon-secret") + } + if forwarded { + req.Header.Set("X-Forwarded-For", "203.0.113.10") + } + return req + } + + unauthorized := httptest.NewRecorder() + srv.Handler().ServeHTTP(unauthorized, request("127.0.0.1:1234", false, false)) + assert.Equal(t, http.StatusUnauthorized, unauthorized.Code) + + proxied := httptest.NewRecorder() + srv.Handler().ServeHTTP(proxied, request("127.0.0.1:1234", true, true)) + assert.Equal(t, http.StatusForbidden, proxied.Code) + + remote := httptest.NewRecorder() + srv.Handler().ServeHTTP(remote, request("203.0.113.10:1234", true, false)) + assert.Equal(t, http.StatusForbidden, remote.Code) + + local := httptest.NewRecorder() + srv.Handler().ServeHTTP(local, request("127.0.0.1:1234", true, false)) + require.Equal(t, http.StatusOK, local.Code, local.Body.String()) + var response artifactExchangeResponse + require.NoError(t, json.Unmarshal(local.Body.Bytes(), &response)) + assert.Equal(t, artifactLifecycleOrigin, response.Origin) +} + +func TestArtifactLifecycleMaintenanceRequiresAuthenticatedDirectLoopback(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dataDir := t.TempDir() + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, WriteTimeout: time.Second, + AuthToken: "daemon-secret", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactRepository(repository)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + request := func(remoteAddr string, authenticated, forwarded bool) *http.Request { + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/maintenance", + strings.NewReader(`{"dry_run":true}`)) + req.Host = "127.0.0.1:0" + req.RemoteAddr = remoteAddr + if authenticated { + req.Header.Set("Authorization", "Bearer daemon-secret") + } + if forwarded { + req.Header.Set("X-Forwarded-For", "203.0.113.10") + } + return req + } + + unauthorized := httptest.NewRecorder() + srv.Handler().ServeHTTP(unauthorized, request("127.0.0.1:1234", false, false)) + assert.Equal(t, http.StatusUnauthorized, unauthorized.Code) + + proxied := httptest.NewRecorder() + srv.Handler().ServeHTTP(proxied, request("127.0.0.1:1234", true, true)) + assert.Equal(t, http.StatusForbidden, proxied.Code) + + remote := httptest.NewRecorder() + srv.Handler().ServeHTTP(remote, request("203.0.113.10:1234", true, false)) + assert.Equal(t, http.StatusForbidden, remote.Code) + + local := httptest.NewRecorder() + srv.Handler().ServeHTTP(local, request("127.0.0.1:1234", true, false)) + assert.Equal(t, http.StatusOK, local.Code, local.Body.String()) +} + +func TestArtifactLifecycleShutdownTimeoutDefersStoreCloseUntilRouteCompletes(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := newLifecycleArtifactStore() + store.openStarted = make(chan struct{}, 1) + store.openRelease = make(chan struct{}) + body := []byte("held artifact") + entry := lifecycleEntry(artifact.Ref{}, body) + ref, err := artifact.NewRef(artifactLifecycleOrigin, artifact.KindRaw, entry.Identity.SHA256) + require.NoError(t, err) + store.entries[ref] = body + + srv := New(config.Config{Host: "127.0.0.1", WriteTimeout: time.Second}, database, nil, + WithArtifactStore(store)) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + srv.SetPort(listener.Addr().(*net.TCPAddr).Port) + serveDone := make(chan error, 1) + go func() { serveDone <- srv.Serve(listener) }() + + requestDone := make(chan error, 1) + go func() { + response, err := http.Get("http://" + listener.Addr().String() + + "/api/v1/artifacts/" + artifactLifecycleOrigin + "/raw/" + ref.Name) + if response != nil { + _ = response.Body.Close() + } + requestDone <- err + }() + select { + case <-store.openStarted: + case <-time.After(time.Second): + t.Fatal("artifact read did not start") + } + + shutdownCtx, cancel := context.WithTimeout(t.Context(), 10*time.Millisecond) + defer cancel() + assert.ErrorIs(t, srv.Shutdown(shutdownCtx), context.DeadlineExceeded) + assert.False(t, store.closed.Load(), "timed-out shutdown must not close a store in active use") + + close(store.openRelease) + require.NoError(t, <-requestDone) + assert.Eventually(t, store.closed.Load, time.Second, time.Millisecond) + assert.Equal(t, int32(1), store.closeCalls.Load()) + assert.ErrorIs(t, <-serveDone, http.ErrServerClosed) +} + +func TestArtifactLifecycleVerifiedReadFailsClosedAfterProducingBytes(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := newLifecycleArtifactStore() + body := []byte("content that must not be served") + entry := lifecycleEntry(artifact.Ref{}, body) + ref, err := artifact.NewRef(artifactLifecycleOrigin, artifact.KindRaw, entry.Identity.SHA256) + require.NoError(t, err) + store.entries[ref] = body + store.corrupt[ref] = true + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + request := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/"+artifactLifecycleOrigin+"/raw/"+ref.Name, nil) + request.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + + assert.Equal(t, http.StatusInternalServerError, response.Code) + assert.NotContains(t, response.Body.String(), string(body)) +} + +func TestArtifactPostRepairsExactQueuedDocbankContent(t *testing.T) { + database, repository, ref, _, body := seedQueuedCorruptArtifact(t) + + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, + WithArtifactRepository(repository)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + request := httptest.NewRequest(http.MethodPost, + "/api/v1/artifacts/"+ref.Origin+"/raw/"+ref.Name, bytes.NewReader(body)) + request.Host = "127.0.0.1:0" + request.RemoteAddr = "127.0.0.1:1234" + request.Header.Set("Origin", "http://127.0.0.1:0") + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, "body: %s", response.Body.String()) + + pending, err := database.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, err) + assert.Empty(t, pending) + _, reader, err := repository.Content().Open(t.Context(), ref) + require.NoError(t, err) + got, err := io.ReadAll(reader) + require.NoError(t, err) + require.NoError(t, reader.Close()) + assert.Equal(t, body, got) +} + +func TestArtifactPostPreservesExactRepairClaimWhenRepairFails(t *testing.T) { + database, repository, ref, identity, body := seedQueuedCorruptArtifact(t) + store := &repairGateArtifactStore{ + ArtifactStore: repository.Content(), err: errors.New("forced repair failure"), + } + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + _, err := srv.humaPostArtifact(t.Context(), &artifactPostInput{ + Origin: ref.Origin, Kind: string(ref.Kind), Name: ref.Name, + ImportMode: "deferred", Body: bytes.NewReader(body), + }) + require.Error(t, err) + + pending, pendingErr := database.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, pendingErr) + require.Len(t, pending, 1) + assert.Equal(t, identity.SHA256, pending[0].SHA256) + assert.Equal(t, identity.Size, pending[0].Size) + require.NoError(t, repository.Close()) +} + +func TestArtifactPostPreservesExactRepairClaimWhenCanceled(t *testing.T) { + database, repository, ref, identity, body := seedQueuedCorruptArtifact(t) + store := &repairGateArtifactStore{ + ArtifactStore: repository.Content(), started: make(chan struct{}, 1), + } + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + ctx, cancel := context.WithCancel(t.Context()) + done := make(chan error, 1) + go func() { + _, err := srv.humaPostArtifact(ctx, &artifactPostInput{ + Origin: ref.Origin, Kind: string(ref.Kind), Name: ref.Name, + ImportMode: "deferred", Body: bytes.NewReader(body), + }) + done <- err + }() + select { + case <-store.started: + case <-time.After(time.Second): + t.Fatal("repair did not start") + } + cancel() + require.ErrorIs(t, <-done, context.Canceled) + + pending, err := database.PendingArtifactRepairs(t.Context(), 10) + require.NoError(t, err) + require.Len(t, pending, 1) + assert.Equal(t, identity.SHA256, pending[0].SHA256) + assert.Equal(t, identity.Size, pending[0].Size) + require.NoError(t, repository.Close()) +} + +func TestArtifactExchangeRejectsSecretURLWithoutResponseOrLogDisclosure(t *testing.T) { + dataDir := t.TempDir() + database := dbtest.OpenTestDBAt(t, filepath.Join(dataDir, "sessions.db")) + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + srv := New(config.Config{ + Host: "127.0.0.1", DataDir: dataDir, AuthToken: "server-auth", + }, database, nil, + WithArtifactRepository(repository)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + const secret = "exchange-secret-value" + request := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/exchange", + strings.NewReader(`{"target":"https://user:`+secret+ + `@example.invalid/archive?token=`+secret+`#`+secret+`","token":"peer-`+secret+`"}`)) + request.Host = "127.0.0.1:0" + request.RemoteAddr = "127.0.0.1:1234" + request.Header.Set("Origin", "http://127.0.0.1:0") + request.Header.Set("Authorization", "Bearer server-auth") + var logs bytes.Buffer + previousLogOutput := log.Writer() + log.SetOutput(&logs) + defer log.SetOutput(previousLogOutput) + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + + assert.Equal(t, http.StatusBadRequest, response.Code, response.Body.String()) + assert.NotContains(t, response.Body.String(), secret) + assert.NotContains(t, logs.String(), secret) +} + +func seedQueuedCorruptArtifact( + t *testing.T, +) (*db.DB, *artifact.Repository, artifact.Ref, artifact.Identity, []byte) { + t.Helper() + dataDir := t.TempDir() + database := dbtest.OpenTestDBAt(t, filepath.Join(dataDir, "sessions.db")) + repository, err := artifact.OpenRepository(t.Context(), dataDir) + require.NoError(t, err) + body := []byte("trusted peer artifact content") + identity := lifecycleEntry(artifact.Ref{}, body).Identity + ref, err := artifact.NewRef(artifactLifecycleOrigin, artifact.KindRaw, identity.SHA256) + require.NoError(t, err) + _, err = repository.Content().Create(t.Context(), ref, identity, + "application/octet-stream", bytes.NewReader(body)) + require.NoError(t, err) + require.NoError(t, database.EnqueueArtifactRepair(t.Context(), db.ArtifactRepair{ + Origin: ref.Origin, Kind: string(ref.Kind), Name: ref.Name, + SHA256: identity.SHA256, Size: identity.Size, + })) + blobPath := filepath.Join(dataDir, "artifacts", "blobs", identity.SHA256[:2], identity.SHA256) + require.NoError(t, os.WriteFile(blobPath, []byte("corrupt"), 0o600)) + return database, repository, ref, identity, body +} + +func TestArtifactLifecycleMetadataCreateHoldsStoreThroughTimedOutShutdown(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + require.NoError(t, artifact.AdoptOrigin(database, artifactLifecycleOrigin)) + store := newLifecycleArtifactStore() + store.createStarted = make(chan struct{}, 1) + store.createRelease = make(chan struct{}) + srv := New(config.Config{ + Host: "127.0.0.1", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactStore(store)) + + appendDone := make(chan error, 1) + go func() { + appendDone <- srv.appendMetadataEvent(t.Context(), artifact.MetadataEventInput{ + SessionID: "session-1", Op: artifact.MetadataOpStar, + }) + }() + select { + case <-store.createStarted: + case <-time.After(time.Second): + t.Fatal("metadata create did not start") + } + + shutdownCtx, cancel := context.WithTimeout(t.Context(), 10*time.Millisecond) + defer cancel() + require.NoError(t, srv.Shutdown(shutdownCtx)) + assert.False(t, store.closed.Load(), + "timed-out shutdown must not close a store during metadata create") + + close(store.createRelease) + require.NoError(t, <-appendDone) + assert.Eventually(t, store.closed.Load, time.Second, time.Millisecond) + assert.Equal(t, int32(1), store.closeCalls.Load()) +} + +func TestArtifactLifecycleBulkMetadataCreateHoldsStoreThroughTimedOutShutdown(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dbtest.SeedSession(t, database, "session-1", "bulk metadata lifecycle") + dbtest.SeedSession(t, database, "session-2", "bulk metadata lifecycle") + require.NoError(t, artifact.AdoptOrigin(database, artifactLifecycleOrigin)) + store := newLifecycleArtifactStore() + store.createStarted = make(chan struct{}, 1) + store.createRelease = make(chan struct{}) + srv := New(config.Config{ + Host: "127.0.0.1", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactStore(store)) + + requestDone := make(chan error, 1) + go func() { + input := &bulkStarInput{} + input.Body.SessionIDs = []string{"session-1", "session-2"} + _, err := srv.humaBulkStar(t.Context(), input) + requestDone <- err + }() + select { + case <-store.createStarted: + case <-time.After(time.Second): + t.Fatal("bulk metadata create did not start") + } + + shutdownCtx, cancel := context.WithTimeout(t.Context(), 10*time.Millisecond) + defer cancel() + require.NoError(t, srv.Shutdown(shutdownCtx)) + assert.False(t, store.closed.Load(), + "timed-out shutdown must not close a store during bulk metadata create") + + close(store.createRelease) + require.NoError(t, <-requestDone) + assert.Eventually(t, store.closed.Load, time.Second, time.Millisecond) + assert.Equal(t, int32(1), store.closeCalls.Load()) +} + +func TestArtifactLifecycleMetadataRepairAndAppendShareShutdownLease(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dbtest.SeedSession(t, database, "session-1", "metadata repair lifecycle") + require.NoError(t, artifact.AdoptOrigin(database, artifactLifecycleOrigin)) + store := newLifecycleArtifactStore() + recorder := artifact.NewMetadataRecorder(database, artifact.MetadataRecorderOptions{ + Origin: artifactLifecycleOrigin, Store: store, + }) + _, err := recorder.Append(t.Context(), artifact.MetadataEventInput{ + SessionID: "session-1", Op: artifact.MetadataOpStar, + }) + require.NoError(t, err) + _, err = recorder.Append(t.Context(), artifact.MetadataEventInput{ + SessionID: "session-1", Op: artifact.MetadataOpUnstar, + }) + require.NoError(t, err) + store.openStarted = make(chan struct{}, 1) + store.openRelease = make(chan struct{}) + srv := New(config.Config{ + Host: "127.0.0.1", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactStore(store)) + + ensureDone := make(chan error, 1) + go func() { + ensureDone <- srv.ensureLocalMetadataEvent(t.Context(), artifact.MetadataEventInput{ + SessionID: "session-1", Op: artifact.MetadataOpStar, + }, "starred", artifact.MetadataOpStar) + }() + select { + case <-store.openStarted: + case <-time.After(time.Second): + t.Fatal("metadata repair did not open its provenance event") + } + + require.NoError(t, srv.Shutdown(t.Context())) + assert.False(t, store.closed.Load(), + "shutdown must retain the store through the repair and append transaction") + + close(store.openRelease) + require.NoError(t, <-ensureDone) + assert.Eventually(t, store.closed.Load, time.Second, time.Millisecond) + assert.Equal(t, int32(1), store.closeCalls.Load()) + assert.Len(t, store.entries, 3, + "a new star event must publish after repair observes the newer unstar state") +} + +func TestArtifactLifecycleSpontaneousServeExitClosesOwnedStoreOnce(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := newLifecycleArtifactStore() + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, + WithArtifactStore(store)) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + serveDone := make(chan error, 1) + go func() { serveDone <- srv.Serve(listener) }() + + require.NoError(t, listener.Close()) + require.Error(t, <-serveDone) + assert.Eventually(t, store.closed.Load, time.Second, time.Millisecond) + assert.Equal(t, int32(1), store.closeCalls.Load()) + require.NoError(t, srv.Shutdown(t.Context())) + assert.Equal(t, int32(1), store.closeCalls.Load()) +} + +func TestArtifactLifecycleServePreservesServerClosedIdentity(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := newLifecycleArtifactStore() + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, + WithArtifactStore(store)) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + serveDone := make(chan error, 1) + go func() { serveDone <- srv.Serve(listener) }() + require.Eventually(t, func() bool { + srv.mu.RLock() + defer srv.mu.RUnlock() + return srv.httpSrv != nil + }, time.Second, time.Millisecond) + + require.NoError(t, srv.Shutdown(t.Context())) + serveErr := <-serveDone + assert.True(t, serveErr == http.ErrServerClosed, + "clean artifact cleanup must preserve the exact HTTP sentinel") + assert.Equal(t, int32(1), store.closeCalls.Load()) +} + +func TestArtifactLifecycleServeJoinsStoreCloseFailure(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + closeErr := errors.New("closing artifact store") + store := newLifecycleArtifactStore() + store.closeErr = closeErr + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, + WithArtifactStore(store)) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + serveDone := make(chan error, 1) + go func() { serveDone <- srv.Serve(listener) }() + require.Eventually(t, func() bool { + srv.mu.RLock() + defer srv.mu.RUnlock() + return srv.httpSrv != nil + }, time.Second, time.Millisecond) + + assert.ErrorIs(t, srv.Shutdown(t.Context()), closeErr) + serveErr := <-serveDone + assert.ErrorIs(t, serveErr, http.ErrServerClosed) + assert.ErrorIs(t, serveErr, closeErr) + assert.Equal(t, int32(1), store.closeCalls.Load()) +} + +func TestArtifactLifecyclePeerStatusReportsCorruptExactHeadWithoutFallback(t *testing.T) { + dir := t.TempDir() + writable, err := db.Open(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + store := newLifecycleArtifactStore() + checkpoint := func(sequence int, body []byte, corrupt bool) { + t.Helper() + ref, err := artifact.NewRef(artifactLifecycleOrigin, artifact.KindCheckpoints, + fmt.Sprintf("cp-%010d.json", sequence)) + require.NoError(t, err) + store.entries[ref] = body + store.corrupt[ref] = corrupt + } + checkpoint(1, []byte(`{"origin":"lifecycle-a1b2c3","seq":1,"sessions":{},"v":1}`+"\n"), false) + corruptHead := []byte(`{"origin":"lifecycle-a1b2c3","seq":2,"sessions":{},"v":1}` + "\n") + checkpoint(2, corruptHead, true) + headIdentity := lifecycleEntry(artifact.Ref{}, corruptHead).Identity + require.NoError(t, writable.RecordArtifactCheckpointHead(t.Context(), db.ArtifactCheckpointHead{ + Origin: artifactLifecycleOrigin, Sequence: 2, + SessionMapSHA256: strings.Repeat("0", 64), + CheckpointSHA256: headIdentity.SHA256, CheckpointSize: headIdentity.Size, + }, nil)) + require.NoError(t, writable.Close()) + readonly, err := db.OpenReadOnly(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, readonly.Close()) }) + srv := New(config.Config{ + Host: "127.0.0.1", ArtifactOriginID: artifactLifecycleOrigin, + }, readonly, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + request := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/peers", nil) + request.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + var peers artifactPeersResponse + require.NoError(t, json.Unmarshal(response.Body.Bytes(), &peers)) + require.Len(t, peers.Peers, 1) + assert.Equal(t, 2, peers.Peers[0].CheckpointSeq) + assert.Equal(t, "error", peers.Peers[0].Status) +} + +type lifecycleRemoteStore struct{ db.Store } + +func (lifecycleRemoteStore) ReadOnly() bool { return true } + +type pagedLifecycleStore struct { + *lifecycleArtifactStore + originCount int + entryCount int + originCalls atomic.Int32 + listCalls atomic.Int32 + statCalls atomic.Int32 + openCalls atomic.Int32 + releaseCalls atomic.Int32 + entryCloseCalls atomic.Int32 + originErr error +} + +func (s *pagedLifecycleStore) Stat( + ctx context.Context, ref artifact.Ref, +) (artifact.Entry, error) { + s.statCalls.Add(1) + return s.lifecycleArtifactStore.Stat(ctx, ref) +} + +func (s *pagedLifecycleStore) Open( + ctx context.Context, ref artifact.Ref, +) (artifact.Entry, artifact.VerifiedReader, error) { + s.openCalls.Add(1) + return s.lifecycleArtifactStore.Open(ctx, ref) +} + +func (s *pagedLifecycleStore) Origins(context.Context) (artifact.OriginIterator, error) { + return &pagedOriginIterator{store: s}, nil +} + +func (s *pagedLifecycleStore) Entries( + _ context.Context, origin string, kind artifact.Kind, +) (artifact.EntryIterator, error) { + return &pagedEntryIterator{store: s, origin: origin, kind: kind}, nil +} + +type pagedOriginIterator struct { + store *pagedLifecycleStore + offset int + closed atomic.Bool +} + +func (i *pagedOriginIterator) Next(ctx context.Context, limit int) ([]string, error) { + if i.closed.Load() { + return nil, os.ErrClosed + } + if err := ctx.Err(); err != nil { + return nil, err + } + if i.store.originErr != nil { + return nil, i.store.originErr + } + i.store.originCalls.Add(1) + end := min(i.offset+limit, i.store.originCount) + origins := make([]string, 0, end-i.offset) + for index := i.offset; index < end; index++ { + origins = append(origins, fmt.Sprintf("origin-%05d-a1b2c3", index)) + } + i.offset = end + if end == i.store.originCount { + return origins, io.EOF + } + return origins, nil +} + +func (i *pagedOriginIterator) Close() error { + if i.closed.CompareAndSwap(false, true) { + i.store.releaseCalls.Add(1) + } + return nil +} + +type pagedEntryIterator struct { + store *pagedLifecycleStore + origin string + kind artifact.Kind + offset int + closed atomic.Bool +} + +func (i *pagedEntryIterator) Next( + ctx context.Context, limit int, +) ([]artifact.Entry, error) { + if i.closed.Load() { + return nil, os.ErrClosed + } + if err := ctx.Err(); err != nil { + return nil, err + } + if i.kind != artifact.KindSegments { + return []artifact.Entry{}, io.EOF + } + i.store.listCalls.Add(1) + end := min(i.offset+limit, i.store.entryCount) + items := make([]artifact.Entry, 0, end-i.offset) + for index := i.offset; index < end; index++ { + name := fmt.Sprintf("%064x.ndjson", index+1) + ref, err := artifact.NewRef(i.origin, artifact.KindSegments, name) + if err != nil { + return nil, err + } + items = append(items, artifact.Entry{Ref: ref}) + } + i.offset = end + if end == i.store.entryCount { + return items, io.EOF + } + return items, nil +} + +func (i *pagedEntryIterator) Close() error { + if i.closed.CompareAndSwap(false, true) { + i.store.entryCloseCalls.Add(1) + } + return nil +} + +func (s *pagedLifecycleStore) ListOrigins( + ctx context.Context, cursor artifact.Cursor, limit int, +) ([]string, artifact.Cursor, error) { + if err := ctx.Err(); err != nil { + return nil, "", err + } + if s.originErr != nil { + return nil, "", s.originErr + } + s.originCalls.Add(1) + start := 0 + if cursor != "" { + var err error + start, err = strconv.Atoi(string(cursor)) + if err != nil { + return nil, "", artifact.ErrArtifactInvalid + } + } + end := min(start+limit, s.originCount) + origins := make([]string, 0, end-start) + for index := start; index < end; index++ { + origins = append(origins, fmt.Sprintf("origin-%05d-a1b2c3", index)) + } + if end == s.originCount { + return origins, "", nil + } + return origins, artifact.Cursor(strconv.Itoa(end)), nil +} + +func TestArtifactPeersUsesBoundedScopedOriginCursor(t *testing.T) { + dir := t.TempDir() + writable, err := db.Open(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + require.NoError(t, writable.Close()) + readonly, err := db.OpenReadOnly(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, readonly.Close()) }) + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 10_000, + } + srv := New(config.Config{Host: "127.0.0.1"}, readonly, nil, + WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + request := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/peers?limit=2", nil) + request.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + var first artifactPeersResponse + require.NoError(t, json.Unmarshal(response.Body.Bytes(), &first)) + require.Len(t, first.Peers, 2) + assert.NotEmpty(t, first.NextCursor) + assert.Equal(t, int32(1), store.originCalls.Load(), + "peer work must be bounded by the requested page") + assert.Zero(t, store.listCalls.Load(), + "peer status must not scan checkpoint history") + + wrongScope := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/origins?limit=2&cursor="+first.NextCursor, nil) + wrongScope.Host = "127.0.0.1:0" + wrongScopeResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(wrongScopeResponse, wrongScope) + assert.Equal(t, http.StatusBadRequest, wrongScopeResponse.Code, + "peer cursors must not be accepted by origin enumeration") + + request = httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/peers?limit=2&cursor="+first.NextCursor, nil) + request.Host = "127.0.0.1:0" + response = httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + var second artifactPeersResponse + require.NoError(t, json.Unmarshal(response.Body.Bytes(), &second)) + require.Len(t, second.Peers, 2) + assert.NotEqual(t, first.Peers[0].Origin, second.Peers[0].Origin) + assert.Equal(t, int32(2), store.originCalls.Load()) +} + +func TestArtifactPeersPointOpensOnlyTheProvenanceSelectedCheckpoint(t *testing.T) { + dir := t.TempDir() + database, err := db.Open(filepath.Join(dir, "sessions.db")) + require.NoError(t, err) + origin := "origin-00000-a1b2c3" + body := []byte(`{"origin":"origin-00000-a1b2c3","seq":73,"sessions":{},"v":1}` + "\n") + identity := lifecycleEntry(artifact.Ref{}, body).Identity + require.NoError(t, database.RecordArtifactPeerCheckpointHead(t.Context(), + db.ArtifactPeerCheckpointHead{ + Origin: origin, Sequence: 73, + CheckpointSHA256: identity.SHA256, CheckpointSize: identity.Size, + })) + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 10_000, + entryCount: 10_000, + } + ref, err := artifact.NewRef(origin, artifact.KindCheckpoints, "cp-0000000073.json") + require.NoError(t, err) + store.entries[ref] = body + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, + WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + request := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/peers?limit=1", nil) + request.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + var peers artifactPeersResponse + require.NoError(t, json.Unmarshal(response.Body.Bytes(), &peers)) + require.Len(t, peers.Peers, 1) + assert.Equal(t, 73, peers.Peers[0].CheckpointSeq) + assert.Equal(t, "in_sync", peers.Peers[0].Status) + assert.Zero(t, store.listCalls.Load(), + "checkpoint status must not enumerate the 10k unrelated entries") + assert.Equal(t, int32(1), store.statCalls.Load()) + assert.Equal(t, int32(1), store.openCalls.Load()) +} + +func (s *pagedLifecycleStore) List( + ctx context.Context, + origin string, + kind artifact.Kind, + cursor artifact.Cursor, + limit int, +) (artifact.Page, error) { + if kind != artifact.KindSegments { + return artifact.Page{}, nil + } + if err := ctx.Err(); err != nil { + return artifact.Page{}, err + } + s.listCalls.Add(1) + start := 0 + if cursor != "" { + var err error + start, err = strconv.Atoi(string(cursor)) + if err != nil { + return artifact.Page{}, artifact.ErrArtifactInvalid + } + } + end := min(start+limit, s.entryCount) + items := make([]artifact.Entry, 0, end-start) + for index := start; index < end; index++ { + name := fmt.Sprintf("%064x.ndjson", index+1) + ref, err := artifact.NewRef(origin, artifact.KindSegments, name) + if err != nil { + return artifact.Page{}, err + } + items = append(items, artifact.Entry{Ref: ref}) + } + page := artifact.Page{Items: items} + if end < s.entryCount { + page.Next = artifact.Cursor(strconv.Itoa(end)) + } + return page, nil +} + +func TestArtifactLifecycleEnumerationOnlyConsumesRequestedStorePage(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), + originCount: 10_000, + entryCount: 10_000, + } + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + origins := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/origins?limit=10", nil) + origins.Host = "127.0.0.1:0" + originsResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(originsResponse, origins) + require.Equal(t, http.StatusOK, originsResponse.Code) + assert.Equal(t, int32(1), store.originCalls.Load(), + "the first HTTP page must not enumerate the remaining archive") + + index := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/origin-00000-a1b2c3/index?limit=10", nil) + index.Host = "127.0.0.1:0" + indexResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(indexResponse, index) + require.Equal(t, http.StatusOK, indexResponse.Code) + assert.Equal(t, int32(1), store.listCalls.Load(), + "the first HTTP index page must not enumerate the remaining collection") +} + +func TestArtifactLifecycleConcurrentEnumerationCoalescesDirtyPublicationAndBaseline(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dbtest.SeedSession(t, database, "local-1", "project") + displayName := "curated before publication" + require.NoError(t, database.RenameSession("local-1", &displayName)) + store := newLifecycleArtifactStore() + srv := New(config.Config{ + Host: "127.0.0.1", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + const callers = 8 + var group sync.WaitGroup + group.Add(callers) + statuses := make(chan int, callers) + for range callers { + go func() { + defer group.Done() + request := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/origins", nil) + request.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + statuses <- response.Code + }() + } + group.Wait() + close(statuses) + for status := range statuses { + assert.Equal(t, http.StatusOK, status) + } + + store.mu.Lock() + defer store.mu.Unlock() + kindCounts := make(map[artifact.Kind]int) + for ref := range store.entries { + if ref.Origin == artifactLifecycleOrigin { + kindCounts[ref.Kind]++ + } + } + assert.Equal(t, 1, kindCounts[artifact.KindCheckpoints], + "concurrent enumeration must converge on one unchanged checkpoint") + assert.Equal(t, 1, kindCounts[artifact.KindMeta], + "metadata baseline must run once for concurrent enumeration") + assert.Equal(t, 1, kindCounts[artifact.KindManifests]) + assert.Equal(t, 1, kindCounts[artifact.KindSegments]) +} + +func TestArtifactLifecycleIndividualReadDoesNotFlushDirtyPublication(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + dbtest.SeedSession(t, database, "local-1", "project") + store := newLifecycleArtifactStore() + body := []byte("already stored peer content") + entry := lifecycleEntry(artifact.Ref{}, body) + ref, err := artifact.NewRef("peer-a1b2c3", artifact.KindRaw, entry.Identity.SHA256) + require.NoError(t, err) + store.entries[ref] = body + srv := New(config.Config{ + Host: "127.0.0.1", ArtifactOriginID: artifactLifecycleOrigin, + }, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + request := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/peer-a1b2c3/raw/"+ref.Name, nil) + request.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, request) + require.Equal(t, http.StatusOK, response.Code, response.Body.String()) + assert.Equal(t, body, response.Body.Bytes()) + + store.mu.Lock() + defer store.mu.Unlock() + assert.Len(t, store.entries, 1, + "an individual artifact read must not publish the dirty local queue") +} + +func TestArtifactLifecycleCursorReleaseClosesUnderlyingStoreCursor(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), + originCount: 100, + } + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + firstRequest := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/origins?limit=10", nil) + firstRequest.Host = "127.0.0.1:0" + firstResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(firstResponse, firstRequest) + require.Equal(t, http.StatusOK, firstResponse.Code) + var first artifactOriginsResponse + require.NoError(t, json.Unmarshal(firstResponse.Body.Bytes(), &first)) + require.NotEmpty(t, first.NextCursor) + + releaseRequest := httptest.NewRequest(http.MethodDelete, + "/api/v1/artifacts/cursors/"+first.NextCursor, nil) + releaseRequest.Host = "127.0.0.1:0" + releaseRequest.Header.Set("Origin", "http://127.0.0.1:0") + releaseResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(releaseResponse, releaseRequest) + require.Equal(t, http.StatusNoContent, releaseResponse.Code) + assert.Equal(t, int32(1), store.releaseCalls.Load()) + + staleRequest := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/origins?limit=10&cursor="+first.NextCursor, nil) + staleRequest.Host = "127.0.0.1:0" + staleResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(staleResponse, staleRequest) + assert.Equal(t, http.StatusBadRequest, staleResponse.Code) +} + +func TestArtifactLifecycleCanceledContinuationClosesUnderlyingStoreCursor(t *testing.T) { + database := dbtest.OpenTestDBAt(t, filepath.Join(t.TempDir(), "sessions.db")) + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), + originCount: 100, + } + srv := New(config.Config{Host: "127.0.0.1"}, database, nil, WithArtifactStore(store)) + t.Cleanup(func() { require.NoError(t, srv.Shutdown(t.Context())) }) + + firstRequest := httptest.NewRequest(http.MethodGet, "/api/v1/artifacts/origins?limit=10", nil) + firstRequest.Host = "127.0.0.1:0" + firstResponse := httptest.NewRecorder() + srv.Handler().ServeHTTP(firstResponse, firstRequest) + require.Equal(t, http.StatusOK, firstResponse.Code) + var first artifactOriginsResponse + require.NoError(t, json.Unmarshal(firstResponse.Body.Bytes(), &first)) + require.NotEmpty(t, first.NextCursor) + + store.originErr = context.Canceled + continuation := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/origins?limit=10&cursor="+first.NextCursor, nil) + continuation.Host = "127.0.0.1:0" + response := httptest.NewRecorder() + srv.Handler().ServeHTTP(response, continuation) + assert.Equal(t, int32(1), store.releaseCalls.Load()) +} diff --git a/internal/server/artifact_maintenance_internal_test.go b/internal/server/artifact_maintenance_internal_test.go new file mode 100644 index 000000000..77013c362 --- /dev/null +++ b/internal/server/artifact_maintenance_internal_test.go @@ -0,0 +1,39 @@ +package server + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestParseArtifactMaintenanceDuration(t *testing.T) { + seconds := func(value int64) *int64 { return &value } + cases := []struct { + name string + exact string + legacy *int64 + want time.Duration + wantErr bool + }{ + {name: "exact nanoseconds", exact: "1.5us", want: 1500 * time.Nanosecond}, + {name: "legacy seconds", legacy: seconds(7), want: 7 * time.Second}, + {name: "absent", want: 0}, + {name: "conflicting exact and legacy", exact: "1s", legacy: seconds(1), wantErr: true}, + {name: "negative exact", exact: "-1ns", wantErr: true}, + {name: "negative legacy", legacy: seconds(-1), wantErr: true}, + {name: "overflowing legacy", legacy: seconds(maxArtifactMaintenanceGraceSeconds + 1), wantErr: true}, + {name: "invalid exact", exact: "tomorrow", wantErr: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := parseArtifactMaintenanceDuration(tc.exact, tc.legacy) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} diff --git a/internal/server/artifact_peer_test.go b/internal/server/artifact_peer_test.go new file mode 100644 index 000000000..e0da77379 --- /dev/null +++ b/internal/server/artifact_peer_test.go @@ -0,0 +1,1173 @@ +package server_test + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "runtime" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/dbtest" + "go.kenn.io/agentsview/internal/server" + "go.kenn.io/docbank" +) + +type artifactOriginsBody struct { + Origins []string `json:"origins"` + NextCursor string `json:"next_cursor,omitempty"` +} + +type artifactPostBody struct { + Origin string `json:"origin"` + Kind string `json:"kind"` + Name string `json:"name"` + Hash string `json:"hash,omitempty"` + Size int64 `json:"size"` + Duplicate bool `json:"duplicate"` +} + +type artifactIndexPageBody struct { + artifact.OriginArtifactIndex + NextCursor string `json:"next_cursor,omitempty"` +} + +type repeatedByteReader struct{ value byte } + +func (r repeatedByteReader) Read(p []byte) (int, error) { + for index := range p { + p[index] = r.value + } + return len(p), nil +} + +type cancelingRequestReader struct { + cancel context.CancelFunc + value byte + sent bool +} + +func (r *cancelingRequestReader) Read(p []byte) (int, error) { + if r.sent { + return 0, context.Canceled + } + if len(p) > 64<<10 { + p = p[:64<<10] + } + for index := range p { + p[index] = r.value + } + r.sent = true + r.cancel() + return len(p), nil +} + +type cancelAfterChecksContext struct { + context.Context + checks atomic.Int64 + cancelAt int64 + cancel context.CancelFunc +} + +func (c *cancelAfterChecksContext) Err() error { + if c.checks.Add(1) >= c.cancelAt { + if c.cancel != nil { + c.cancel() + return c.Context.Err() + } + return context.Canceled + } + return nil +} + +type countingHTTPResponse struct { + header http.Header + status int + bytes int64 +} + +func seedArtifactStore( + t *testing.T, store artifact.ArtifactStore, origin string, kind artifact.Kind, body []byte, +) artifact.Ref { + t.Helper() + hash := sha256.Sum256(body) + name := hex.EncodeToString(hash[:]) + if kind == artifact.KindSegments { + name += ".ndjson" + } + ref, err := artifact.NewRef(origin, kind, name) + require.NoError(t, err) + result, err := store.Create(t.Context(), ref, artifact.Identity{ + SHA256: hex.EncodeToString(hash[:]), Size: int64(len(body)), + }, "application/octet-stream", bytes.NewReader(body)) + require.NoError(t, err) + assert.Equal(t, ref, result.Entry.Ref) + return ref +} + +func exportArtifactFixture( + t *testing.T, ctx context.Context, database *db.DB, origin string, +) (artifact.ArtifactStore, artifact.ExportResult) { + t.Helper() + repository, err := artifact.OpenRepository(ctx, t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + result, err := artifact.ExportToStore(ctx, database, repository.Content(), artifact.ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + return repository.Content(), result +} + +func oneArtifactRef( + t *testing.T, store artifact.ArtifactStore, origin string, kind artifact.Kind, +) artifact.Ref { + t.Helper() + entries := collectArtifactEntries(t, store, origin, kind, 10) + require.Len(t, entries, 1) + return entries[0].Ref +} + +func wireArtifact( + t *testing.T, store artifact.ArtifactStore, ref artifact.Ref, +) (artifact.WireRef, []byte) { + t.Helper() + _, reader, err := store.Open(t.Context(), ref) + require.NoError(t, err) + defer func() { require.NoError(t, reader.Close()) }() + wire, err := artifact.ToWireRef(ref) + require.NoError(t, err) + var body bytes.Buffer + require.NoError(t, artifact.EncodeWire(t.Context(), ref, reader, &body)) + require.NoError(t, reader.Verify()) + return wire, body.Bytes() +} + +func TestArtifactMaintenanceRouteUsesDaemonOwnedStore(t *testing.T) { + te := setupArtifact(t, withAuth("secret")) + body := strings.NewReader(`{"grace_seconds":3600,"quarantine_grace_seconds":7200,"max_objects":8,"max_bytes":1048576}`) + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/maintenance", body) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer secret") + recorder := httptest.NewRecorder() + te.handler.ServeHTTP(recorder, req) + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + var response struct { + Logical struct { + Origins int `json:"origins"` + } `json:"logical"` + Physical struct { + Supported bool `json:"supported"` + } `json:"physical"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) + assert.Zero(t, response.Logical.Origins) + assert.True(t, response.Physical.Supported, + "the daemon-owned Docbank store must service physical maintenance") +} + +func TestArtifactMaintenanceRouteRejectsOverflowingGrace(t *testing.T) { + te := setupArtifact(t, withAuth("secret")) + body := strings.NewReader(`{"grace_seconds":9223372036854775807}`) + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/maintenance", body) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer secret") + recorder := httptest.NewRecorder() + te.handler.ServeHTTP(recorder, req) + assert.Equal(t, http.StatusBadRequest, recorder.Code, recorder.Body.String()) +} + +func TestArtifactMaintenanceRouteRejectsCursorFromAnotherStage(t *testing.T) { + te := setupArtifact(t, withAuth("secret")) + body := strings.NewReader(`{"max_objects":1,"gc_cursor":"agentsview-artifact-maintenance:v1:empty-trash"}`) + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/maintenance", body) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer secret") + recorder := httptest.NewRecorder() + + te.handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusBadRequest, recorder.Code, recorder.Body.String()) +} + +func TestArtifactMaintenanceRouteRejectsOversizedBudgetBeforeLogicalRetention(t *testing.T) { + store := &maintenanceBudgetProbeStore{} + te := setupWithServerOpts(t, []server.Option{server.WithArtifactStore(store)}, withAuth("secret")) + body := strings.NewReader(fmt.Sprintf(`{"max_objects":%d}`, docbank.MaxMaintenanceObjects+1)) + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/maintenance", body) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer secret") + recorder := httptest.NewRecorder() + + te.handler.ServeHTTP(recorder, req) + + assert.Equal(t, http.StatusBadRequest, recorder.Code, recorder.Body.String()) + assert.Zero(t, store.listOrigins, + "invalid physical limits must be rejected before logical retention reads the store") +} + +func TestArtifactMaintenanceRoutePreservesExplicitZeroBudgets(t *testing.T) { + store := &maintenanceBudgetProbeStore{} + te := setupWithServerOpts(t, []server.Option{server.WithArtifactStore(store)}, withAuth("secret")) + body := strings.NewReader(`{"max_objects":0,"max_bytes":0}`) + req := httptest.NewRequest(http.MethodPost, "/api/v1/artifacts/maintenance", body) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer secret") + recorder := httptest.NewRecorder() + + te.handler.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + assert.Equal(t, artifact.WorkBudget{}, store.emptyTrash) + assert.Equal(t, artifact.WorkBudget{}, store.gc) + assert.Equal(t, artifact.WorkBudget{}, store.repack) +} + +type maintenanceBudgetProbeStore struct { + artifact.ArtifactStore + listOrigins int + emptyTrash artifact.WorkBudget + gc artifact.WorkBudget + repack artifact.WorkBudget +} + +func (s *maintenanceBudgetProbeStore) Origins(context.Context) (artifact.OriginIterator, error) { + s.listOrigins++ + return emptyArtifactOriginIterator{}, nil +} + +type emptyArtifactOriginIterator struct{} + +func (emptyArtifactOriginIterator) Next(context.Context, int) ([]string, error) { + return nil, io.EOF +} + +func (emptyArtifactOriginIterator) Close() error { return nil } + +func (s *maintenanceBudgetProbeStore) Verify( + context.Context, artifact.WorkBudget, +) (artifact.MaintenanceResult, error) { + return artifact.MaintenanceResult{}, nil +} + +func (s *maintenanceBudgetProbeStore) EmptyTrash( + _ context.Context, _ time.Duration, budget artifact.WorkBudget, +) (artifact.MaintenanceResult, error) { + s.emptyTrash = budget + return artifact.MaintenanceResult{}, nil +} + +func (s *maintenanceBudgetProbeStore) GarbageCollect( + _ context.Context, budget artifact.WorkBudget, +) (artifact.MaintenanceResult, error) { + s.gc = budget + return artifact.MaintenanceResult{}, nil +} + +func (s *maintenanceBudgetProbeStore) Repack( + _ context.Context, budget artifact.WorkBudget, +) (artifact.MaintenanceResult, error) { + s.repack = budget + return artifact.MaintenanceResult{}, nil +} + +func (w *countingHTTPResponse) Header() http.Header { + if w.header == nil { + w.header = make(http.Header) + } + return w.header +} + +func (w *countingHTTPResponse) WriteHeader(status int) { w.status = status } + +func (w *countingHTTPResponse) Write(p []byte) (int, error) { + if w.status == 0 { + w.status = http.StatusOK + } + w.bytes += int64(len(p)) + return len(p), nil +} + +func repeatedByteHash(t *testing.T, value byte, size int64) string { + t.Helper() + hasher := sha256.New() + _, err := io.CopyN(hasher, repeatedByteReader{value: value}, size) + require.NoError(t, err) + return hex.EncodeToString(hasher.Sum(nil)) +} + +func measureArtifactRouteAlloc(t *testing.T, run func()) uint64 { + t.Helper() + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + run() + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc +} + +func TestArtifactPeerPostRouteMemoryRemainsBounded(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + const origin = "peer-a1b2c3" + measure := func(size int64) uint64 { + name := repeatedByteHash(t, 'p', size) + request := httptest.NewRequest(http.MethodPost, + "/api/v1/artifacts/"+origin+"/raw/"+name, + io.LimitReader(repeatedByteReader{value: 'p'}, size), + ) + request.ContentLength = size + request.Header.Set("Content-Type", "application/octet-stream") + request.Header.Set("X-Agentsview-Artifact-Import", "deferred") + response := &countingHTTPResponse{} + allocated := measureArtifactRouteAlloc(t, func() { + te.handler.ServeHTTP(response, request) + }) + assert.Equal(t, http.StatusOK, response.status) + return allocated + } + + small := measure(1 << 20) + large := measure(24 << 20) + assert.Less(t, large, small+(4<<20), + "production POST route allocation must not scale with request bytes") +} + +func TestArtifactPeerGetRouteMemoryRemainsBounded(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + const origin = "peer-a1b2c3" + measure := func(size int64) uint64 { + name := repeatedByteHash(t, 'g', size) + post := httptest.NewRequest(http.MethodPost, + "/api/v1/artifacts/"+origin+"/raw/"+name, + io.LimitReader(repeatedByteReader{value: 'g'}, size), + ) + post.ContentLength = size + post.Header.Set("Content-Type", "application/octet-stream") + post.Header.Set("X-Agentsview-Artifact-Import", "deferred") + postResponse := &countingHTTPResponse{} + te.handler.ServeHTTP(postResponse, post) + require.Equal(t, http.StatusOK, postResponse.status) + + request := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/"+origin+"/raw/"+name, nil, + ) + response := &countingHTTPResponse{} + allocated := measureArtifactRouteAlloc(t, func() { + te.handler.ServeHTTP(response, request) + }) + assert.Equal(t, http.StatusOK, response.status) + assert.Equal(t, size, response.bytes) + return allocated + } + + small := measure(1 << 20) + large := measure(24 << 20) + assert.Less(t, large, small+(4<<20), + "production GET route allocation must not scale with response bytes") +} + +func TestArtifactPeerRoutesHonorCancellationWithoutSuccessOrPublication(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + const origin = "peer-a1b2c3" + + postCtx, cancelPost := context.WithCancel(t.Context()) + postName := repeatedByteHash(t, 'c', 2<<20) + post := httptest.NewRequest(http.MethodPost, + "/api/v1/artifacts/"+origin+"/raw/"+postName, + &cancelingRequestReader{cancel: cancelPost, value: 'c'}, + ).WithContext(postCtx) + post.Header.Set("Content-Type", "application/octet-stream") + postResponse := &countingHTTPResponse{} + te.handler.ServeHTTP(postResponse, post) + assert.NotEqual(t, http.StatusOK, postResponse.status) + assert.NoFileExists(t, filepath.Join(te.dataDir, "artifacts", origin, "raw", postName)) + + getSize := int64(2 << 20) + getName := repeatedByteHash(t, 'd', getSize) + seed := httptest.NewRequest(http.MethodPost, + "/api/v1/artifacts/"+origin+"/raw/"+getName, + io.LimitReader(repeatedByteReader{value: 'd'}, getSize), + ) + seed.Header.Set("Content-Type", "application/octet-stream") + seed.Header.Set("X-Agentsview-Artifact-Import", "deferred") + seedResponse := &countingHTTPResponse{} + te.handler.ServeHTTP(seedResponse, seed) + require.Equal(t, http.StatusOK, seedResponse.status) + + getContext := &cancelAfterChecksContext{Context: t.Context(), cancelAt: 8} + get := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/"+origin+"/raw/"+getName, nil, + ).WithContext(getContext) + getResponse := &countingHTTPResponse{} + te.handler.ServeHTTP(getResponse, get) + assert.NotEqual(t, http.StatusOK, getResponse.status) + assert.Zero(t, getResponse.bytes) + assert.GreaterOrEqual(t, getContext.checks.Load(), getContext.cancelAt) +} + +func TestArtifactPeerIndexPagesWireNamesWithOpaqueCursor(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + const origin = "peer-a1b2c3" + refs := make(map[string]artifact.Ref, 513) + for index := range 513 { + body := fmt.Appendf(nil, "raw-%04d", index) + ref := seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, body) + refs[ref.Name] = ref + } + + firstResponse := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/"+origin+"/index", nil, "") + assertStatus(t, firstResponse, http.StatusOK) + var first artifactIndexPageBody + require.NoError(t, json.Unmarshal(firstResponse.Body.Bytes(), &first)) + assert.Len(t, first.Raw, 512) + require.NotEmpty(t, first.NextCursor) + assert.NotContains(t, first.NextCursor, first.Raw[len(first.Raw)-1], + "cursor must be opaque rather than a raw filename") + firstNames := make(map[string]struct{}, len(first.Raw)) + for _, name := range first.Raw { + firstNames[name] = struct{}{} + } + remaining := "" + for name := range refs { + if _, found := firstNames[name]; !found { + remaining = name + break + } + } + require.NotEmpty(t, remaining) + require.NoError(t, te.artifactStore.Trash(t.Context(), refs[remaining])) + + secondResponse := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/"+origin+"/index?cursor="+url.QueryEscape(first.NextCursor), nil, "") + assertStatus(t, secondResponse, http.StatusOK) + var second artifactIndexPageBody + require.NoError(t, json.Unmarshal(secondResponse.Body.Bytes(), &second)) + assert.Len(t, second.Raw, 1) + assert.Empty(t, second.NextCursor) + assert.Equal(t, remaining, second.Raw[0], + "later pages must come from the initial snapshot") +} + +func TestArtifactPeerOriginsUseBoundedSnapshotCursor(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + refs := make(map[string]artifact.Ref, 512) + for index := range 512 { + origin := fmt.Sprintf("peer-%04d-a1b2c3", index) + refs[origin] = seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, + fmt.Appendf(nil, "origin-%04d", index)) + } + + firstResponse := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/origins?limit=512", nil, "") + assertStatus(t, firstResponse, http.StatusOK) + var first artifactOriginsBody + require.NoError(t, json.Unmarshal(firstResponse.Body.Bytes(), &first)) + assert.Len(t, first.Origins, 512) + require.NotEmpty(t, first.NextCursor) + + removed := "peer-0511-a1b2c3" + require.NoError(t, te.artifactStore.Trash(t.Context(), refs[removed])) + secondResponse := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/origins?cursor="+url.QueryEscape(first.NextCursor), nil, "") + assertStatus(t, secondResponse, http.StatusOK) + var second artifactOriginsBody + require.NoError(t, json.Unmarshal(secondResponse.Body.Bytes(), &second)) + assert.Equal(t, []string{removed}, second.Origins, + "later pages must advance a stable snapshot instead of rescanning") + assert.Empty(t, second.NextCursor) +} + +func TestArtifactPeerCursorReleaseInvalidatesContinuation(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + for index := range 513 { + origin := fmt.Sprintf("peer-%04d-a1b2c3", index) + seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, + fmt.Appendf(nil, "origin-%04d", index)) + } + + firstResponse := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/origins?limit=512", nil, "") + assertStatus(t, firstResponse, http.StatusOK) + var first artifactOriginsBody + require.NoError(t, json.Unmarshal(firstResponse.Body.Bytes(), &first)) + require.NotEmpty(t, first.NextCursor) + + releaseResponse := artifactPeerRequest(t, te, http.MethodDelete, + "/api/v1/artifacts/cursors/"+url.PathEscape(first.NextCursor), nil, "") + assertStatus(t, releaseResponse, http.StatusNoContent) + + staleResponse := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/origins?cursor="+url.QueryEscape(first.NextCursor), nil, "") + assertStatus(t, staleResponse, http.StatusBadRequest) +} + +func TestArtifactPeerHighCardinalityOriginsRetainOneBoundedCursor(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + for index := range 513 { + origin := fmt.Sprintf("peer-%04d-a1b2c3", index) + seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, + fmt.Appendf(nil, "origin-%04d", index)) + } + + response := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/origins?limit=512", nil, "") + assertStatus(t, response, http.StatusOK) + var page artifactOriginsBody + require.NoError(t, json.Unmarshal(response.Body.Bytes(), &page)) + assert.Len(t, page.Origins, 512) + require.NotEmpty(t, page.NextCursor) + release := artifactPeerRequest(t, te, http.MethodDelete, + "/api/v1/artifacts/cursors/"+url.PathEscape(page.NextCursor), nil, "") + assertStatus(t, release, http.StatusNoContent) + stale := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/origins?cursor="+url.QueryEscape(page.NextCursor), nil, "") + assertStatus(t, stale, http.StatusBadRequest) +} + +func TestArtifactPeerOriginSnapshotCancellationReleasesSpool(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + for index := range 513 { + origin := fmt.Sprintf("peer-%04d-a1b2c3", index) + seedArtifactStore(t, te.artifactStore, origin, artifact.KindRaw, + fmt.Appendf(nil, "origin-%04d", index)) + } + base, cancel := context.WithCancel(t.Context()) + t.Cleanup(cancel) + ctx := &cancelAfterChecksContext{Context: base, cancelAt: 2, cancel: cancel} + request := httptest.NewRequest(http.MethodGet, + "/api/v1/artifacts/origins?limit=512", nil).WithContext(ctx) + response := httptest.NewRecorder() + + te.handler.ServeHTTP(response, request) + + assert.NotEqual(t, http.StatusOK, response.Code) + assert.GreaterOrEqual(t, ctx.checks.Load(), ctx.cancelAt) +} + +type artifactRemoteStore struct { + db.Store +} + +func (artifactRemoteStore) ReadOnly() bool { return true } + +func (artifactRemoteStore) MachineSessionCounts( + context.Context, +) (map[string]int, error) { + return map[string]int{}, nil +} + +func (artifactRemoteStore) CountMetadataConflicts(context.Context) (int, error) { + return 0, nil +} + +func TestArtifactPeerRoutesRequireBearerAuthWhenConfigured(t *testing.T) { + te := setupArtifact(t, withAuth("secret")) + + w := artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/origins", nil, "") + assertStatus(t, w, http.StatusUnauthorized) + + w = artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/origins", nil, "secret") + assertStatus(t, w, http.StatusOK) +} + +func TestArtifactPeerRoutesPostDuplicateAndFetch(t *testing.T) { + te := setupArtifact(t, withAuth("secret")) + origin := "peer-a1b2c3" + metadataBody, metadataName := peerMetadataArtifact( + origin, + "2026-06-14T010203.000000001Z-00000000000000000000", + ) + + w := artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+origin+"/meta/"+url.PathEscape(metadataName), + metadataBody, "secret", + ) + assertStatus(t, w, http.StatusOK) + posted := decode[artifactPostBody](t, w) + assert.False(t, posted.Duplicate) + assert.Equal(t, "meta", posted.Kind) + assert.Equal(t, metadataName, posted.Name) + + w = artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+origin+"/meta/"+url.PathEscape(metadataName), + metadataBody, "secret", + ) + assertStatus(t, w, http.StatusOK) + posted = decode[artifactPostBody](t, w) + assert.True(t, posted.Duplicate) + + w = artifactPeerRequest( + t, te, http.MethodGet, + "/api/v1/artifacts/"+origin+"/meta/"+url.PathEscape(metadataName), + nil, "secret", + ) + assertStatus(t, w, http.StatusOK) + assert.Equal(t, "application/octet-stream", w.Header().Get("Content-Type")) + assert.Equal(t, metadataBody, w.Body.Bytes()) + + checkpoint := []byte(`{"origin":"peer-a1b2c3","seq":1,"sessions":{},"v":1}` + "\n") + w = artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+origin+"/checkpoints/cp-0000000001.json", + checkpoint, "secret", + ) + assertStatus(t, w, http.StatusOK) + + w = artifactPeerRequest( + t, te, http.MethodGet, + "/api/v1/artifacts/"+origin+"/checkpoint", + nil, "secret", + ) + assertStatus(t, w, http.StatusOK) + assert.Equal(t, "application/octet-stream", w.Header().Get("Content-Type")) + assert.Equal(t, checkpoint, w.Body.Bytes()) + + w = artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/origins", nil, "secret") + assertStatus(t, w, http.StatusOK) + origins := decode[artifactOriginsBody](t, w) + assert.Contains(t, origins.Origins, origin) +} + +func TestArtifactPeerRawAndZstdFetchesUseBinaryMediaTypeAndExactBytes(t *testing.T) { + te := setupArtifact(t) + origin := "peer-a1b2c3" + rawBody := []byte{0x00, 0xff, 0x80, 0x7f, 0x01} + rawHash := sha256.Sum256(rawBody) + rawName := hex.EncodeToString(rawHash[:]) + + request := httptest.NewRequest(http.MethodPost, + "/api/v1/artifacts/"+origin+"/raw/"+rawName, bytes.NewReader(rawBody)) + request.Header.Set("Content-Type", "application/octet-stream") + request.Header.Set("X-Agentsview-Artifact-Import", "deferred") + response := httptest.NewRecorder() + te.handler.ServeHTTP(response, request) + assertStatus(t, response, http.StatusOK) + + w := artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/"+origin+"/raw/"+rawName, nil, "") + assertStatus(t, w, http.StatusOK) + assert.Equal(t, "application/octet-stream", w.Header().Get("Content-Type")) + assert.Equal(t, rawBody, w.Body.Bytes()) + + peerDB, err := db.Open(filepath.Join(t.TempDir(), "peer.db")) + require.NoError(t, err) + t.Cleanup(func() { peerDB.Close() }) + first := "binary contract" + dbtest.SeedSession(t, peerDB, "sess-1", "alpha", func(session *db.Session) { + session.FirstMessage = &first + }) + require.NoError(t, peerDB.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + })) + peerStore, _ := exportArtifactFixture(t, t.Context(), peerDB, origin) + segmentRef := oneArtifactRef(t, peerStore, origin, artifact.KindSegments) + segmentWire, segmentBody := wireArtifact(t, peerStore, segmentRef) + postArtifactBodyDeferred(t, te, segmentWire, segmentBody) + + w = artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/"+origin+"/segments/"+url.PathEscape(segmentWire.Name), + nil, "") + assertStatus(t, w, http.StatusOK) + assert.Equal(t, "application/octet-stream", w.Header().Get("Content-Type")) + assert.Equal(t, segmentBody, w.Body.Bytes()) +} + +func TestArtifactPeerMutationRoutesAreAbsentInRemoteMode(t *testing.T) { + dir := tempDirWithRetryCleanup(t) + cfg := config.Config{ + Host: "127.0.0.1", + Port: 0, + DataDir: dir, + WriteTimeout: 30 * time.Second, + } + srv := server.New(cfg, artifactRemoteStore{}, nil) + te := &testEnv{ + srv: srv, + handler: wrapTestHandler(cfg, srv.Handler()), + dataDir: dir, + } + origin := "peer-a1b2c3" + metadataBody, metadataName := peerMetadataArtifact( + origin, + "2026-06-14T010203.000000001Z-00000000000000000000", + ) + + w := artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+origin+"/meta/"+url.PathEscape(metadataName), + metadataBody, "", + ) + + assertStatus(t, w, http.StatusNotFound) +} + +func TestArtifactPeerReadRoutesServeInjectedStoreInRemoteMode(t *testing.T) { + dir := tempDirWithRetryCleanup(t) + cfg := config.Config{ + Host: "127.0.0.1", + Port: 0, + DataDir: dir, + WriteTimeout: 30 * time.Second, + } + repository, err := artifact.OpenRepository(t.Context(), dir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + origin := "peer-a1b2c3" + raw := []byte("remote artifact") + ref := seedArtifactStore(t, repository.Content(), origin, artifact.KindRaw, raw) + + srv := server.New(cfg, artifactRemoteStore{}, nil, server.WithArtifactRepository(repository)) + te := &testEnv{ + srv: srv, + handler: wrapTestHandler(cfg, srv.Handler()), + dataDir: dir, + } + + for _, path := range []string{ + "/api/v1/artifacts/origins", + "/api/v1/artifacts/" + origin + "/index", + "/api/v1/artifacts/" + origin + "/raw/" + ref.Name, + } { + w := artifactPeerRequest(t, te, http.MethodGet, path, nil, "") + assertStatus(t, w, http.StatusOK) + } +} + +func TestArtifactPeerMutationRoutesAreAbsentForReadOnlySQLite(t *testing.T) { + dir := tempDirWithRetryCleanup(t) + dbPath := filepath.Join(dir, "test.db") + writable, err := db.Open(dbPath) + require.NoError(t, err) + require.NoError(t, writable.Close()) + readonly, err := db.OpenReadOnly(dbPath) + require.NoError(t, err) + t.Cleanup(func() { readonly.Close() }) + + cfg := config.Config{ + Host: "127.0.0.1", + Port: 0, + DataDir: dir, + DBPath: dbPath, + ArtifactOriginID: "desktop-d4e5f6", + WriteTimeout: 30 * time.Second, + } + srv := server.New(cfg, readonly, nil) + te := &testEnv{ + srv: srv, + handler: wrapTestHandler(cfg, srv.Handler()), + db: readonly, + dataDir: dir, + } + origin := "peer-a1b2c3" + metadataBody, metadataName := peerMetadataArtifact( + origin, + "2026-06-14T010203.000000001Z-00000000000000000000", + ) + + w := artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+origin+"/meta/"+url.PathEscape(metadataName), + metadataBody, "", + ) + + assertStatus(t, w, http.StatusNotFound) +} + +func TestArtifactResetDaemonRouteMovesAsideAndReportsManualRecovery(t *testing.T) { + dir := tempDirWithRetryCleanup(t) + database := dbtest.OpenTestDBAt(t, filepath.Join(dir, "test.db")) + origin := "desktop-d4e5f6" + startedAt := "2026-06-14T01:02:03Z" + require.NoError(t, database.UpsertSession(db.Session{ + ID: "local-session", Machine: "local", Agent: "codex", Project: "project-a", + StartedAt: &startedAt, CreatedAt: startedAt, + })) + repository, err := artifact.OpenRepository(t.Context(), dir) + require.NoError(t, err) + foreign := seedArtifactStore(t, repository.Content(), "peer-a1b2c3", artifact.KindRaw, + []byte("foreign relay bytes")) + cfg := config.Config{ + Host: "127.0.0.1", DataDir: dir, DBPath: filepath.Join(dir, "test.db"), + ArtifactOriginID: origin, WriteTimeout: 30 * time.Second, + RequireAuth: true, AuthToken: "daemon-secret", + } + srv := server.New(cfg, database, nil, server.WithArtifactRepository(repository)) + t.Cleanup(func() { + require.NoError(t, srv.Shutdown(context.Background())) + require.NoError(t, repository.Close()) + }) + te := &testEnv{srv: srv, handler: wrapTestHandler(cfg, srv.Handler()), db: database, dataDir: dir} + + w := artifactPeerRequest(t, te, http.MethodPost, "/api/v1/artifacts/reset", nil, "daemon-secret") + assertStatus(t, w, http.StatusOK) + var response struct { + VaultRoot string `json:"vault_root"` + DiagnosticRoot string `json:"diagnostic_root"` + Export artifact.ExportResult `json:"export"` + ManualCleanup string `json:"manual_cleanup"` + ForeignArtifacts string `json:"foreign_artifacts"` + } + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &response)) + assert.NotEmpty(t, response.VaultRoot) + assert.DirExists(t, response.VaultRoot) + assert.DirExists(t, response.DiagnosticRoot) + assert.Equal(t, artifact.ArtifactResetManualCleanupWarning, response.ManualCleanup) + assert.Equal(t, artifact.ArtifactResetForeignRelayWarning, response.ForeignArtifacts) + assert.NotZero(t, response.Export.CheckpointSequence) + + w = artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/origins", nil, "daemon-secret") + assertStatus(t, w, http.StatusOK) + origins := decode[artifactOriginsBody](t, w) + assert.Equal(t, []string{origin}, origins.Origins) + w = artifactPeerRequest(t, te, http.MethodGet, + "/api/v1/artifacts/peer-a1b2c3/raw/"+foreign.Name, nil, "daemon-secret") + assertStatus(t, w, http.StatusNotFound) + session, err := database.GetSession(t.Context(), "local-session") + require.NoError(t, err) + assert.NotNil(t, session) +} + +type artifactPeerBody struct { + Origin string `json:"origin"` + IsLocal bool `json:"is_local"` + CheckpointSeq int `json:"checkpoint_seq"` + PublishedSessions int `json:"published_sessions"` + LocalSessions int `json:"local_sessions"` + LastPublished string `json:"last_published"` +} + +type artifactPeersBody struct { + LocalOrigin string `json:"local_origin"` + Peers []artifactPeerBody `json:"peers"` + ConflictCount int `json:"conflict_count"` + PendingImports int `json:"pending_imports"` + OldestPending string `json:"oldest_pending_at"` +} + +func TestArtifactPeersStatus(t *testing.T) { + local := "desktop-d4e5f6" + te := setupArtifact(t, withArtifactOrigin(local)) + ctx := context.Background() + first := "hi" + + // Two owned sessions, exported so the local origin gets a checkpoint. + dbtest.SeedSession(t, te.db, "local-1", "proj", func(s *db.Session) { s.FirstMessage = &first }) + dbtest.SeedSession(t, te.db, "local-2", "proj", func(s *db.Session) { s.FirstMessage = &first }) + exported, err := artifact.ExportToStore(ctx, te.db, te.artifactStore, artifact.ExportOptions{ + Origin: local, + Full: true, + }) + require.NoError(t, err) + require.Equal(t, 2, exported.ExportedSessions) + + // A foreign peer publishes one session that the server imports. + origin := "peer-a1b2c3" + peerDB, err := db.Open(filepath.Join(t.TempDir(), "peer.db")) + require.NoError(t, err) + t.Cleanup(func() { peerDB.Close() }) + dbtest.SeedSession(t, peerDB, "sess-1", "alpha", func(s *db.Session) { s.FirstMessage = &first }) + require.NoError(t, peerDB.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + })) + peerStore, _ := exportArtifactFixture(t, ctx, peerDB, origin) + postArtifactRef(t, te, peerStore, oneArtifactRef(t, peerStore, origin, artifact.KindSegments)) + postArtifactRef(t, te, peerStore, oneArtifactRef(t, peerStore, origin, artifact.KindManifests)) + postArtifactRef(t, te, peerStore, oneArtifactRef(t, peerStore, origin, artifact.KindCheckpoints)) + + w := artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/peers", nil, "") + assertStatus(t, w, http.StatusOK) + body := decode[artifactPeersBody](t, w) + + assert.Equal(t, local, body.LocalOrigin) + assert.Equal(t, 0, body.ConflictCount) + require.Len(t, body.Peers, 2) + + byOrigin := map[string]artifactPeerBody{} + for _, p := range body.Peers { + byOrigin[p.Origin] = p + } + + localPeer, ok := byOrigin[local] + require.True(t, ok, "local origin present in peers") + assert.True(t, localPeer.IsLocal) + assert.Equal(t, 2, localPeer.PublishedSessions) + assert.Equal(t, 2, localPeer.LocalSessions) + assert.NotEmpty(t, localPeer.LastPublished) + + peer, ok := byOrigin[origin] + require.True(t, ok, "foreign origin present in peers") + assert.False(t, peer.IsLocal) + assert.Equal(t, 1, peer.PublishedSessions) + assert.Equal(t, 1, peer.LocalSessions) + assert.Equal(t, 1, peer.CheckpointSeq) + + // Landing is checkpoint provenance, not visibility in the session list. + // A locally trashed replica is still fully imported from this checkpoint. + require.NoError(t, te.db.SoftDeleteSession(origin+"~sess-1")) + w = artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/peers", nil, "") + assertStatus(t, w, http.StatusOK) + body = decode[artifactPeersBody](t, w) + for _, p := range body.Peers { + if p.Origin == origin { + assert.Equal(t, 1, p.LocalSessions) + } + } +} + +func TestArtifactPeersStatusDoesNotCountUnrelatedMachineRows(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + origin := "peer-a1b2c3" + peerDB, err := db.Open(filepath.Join(t.TempDir(), "peer.db")) + require.NoError(t, err) + t.Cleanup(func() { peerDB.Close() }) + first := "published" + dbtest.SeedSession(t, peerDB, "sess-1", "alpha", func(s *db.Session) { + s.FirstMessage = &first + }) + peerStore, _ := exportArtifactFixture(t, context.Background(), peerDB, origin) + + // Publish only the checkpoint, leaving its manifest unresolved, then add a + // stale row for the same machine that is not a member of that checkpoint. + postArtifactRef(t, te, peerStore, + oneArtifactRef(t, peerStore, origin, artifact.KindCheckpoints)) + dbtest.SeedSession(t, te.db, origin+"~stale", "old", func(s *db.Session) { + s.Machine = origin + }) + + w := artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/peers", nil, "") + assertStatus(t, w, http.StatusOK) + body := decode[artifactPeersBody](t, w) + assert.Equal(t, 1, body.PendingImports) + assert.NotEmpty(t, body.OldestPending) + for _, p := range body.Peers { + if p.Origin == origin { + assert.Equal(t, 1, p.PublishedSessions) + assert.Zero(t, p.LocalSessions, + "rows outside the latest checkpoint must not satisfy peer status") + return + } + } + t.Fatal("foreign peer missing from status") +} + +func TestArtifactPeersStatusPublishesEmptyLocalOrigin(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + // Discovery publishes an explicit empty checkpoint for a configured origin. + w := artifactPeerRequest(t, te, http.MethodGet, "/api/v1/artifacts/peers", nil, "") + assertStatus(t, w, http.StatusOK) + body := decode[artifactPeersBody](t, w) + assert.Equal(t, "desktop-d4e5f6", body.LocalOrigin) + require.Len(t, body.Peers, 1) + assert.True(t, body.Peers[0].IsLocal) + assert.Equal(t, 0, body.Peers[0].PublishedSessions) + assert.Equal(t, 1, body.Peers[0].CheckpointSeq) + assert.NotEmpty(t, body.Peers[0].LastPublished) +} + +func TestArtifactPeerPostRejectsHashMismatch(t *testing.T) { + te := setupArtifact(t, withAuth("secret")) + origin := "peer-a1b2c3" + metadataBody, _ := peerMetadataArtifact( + origin, + "2026-06-14T010203.000000001Z-00000000000000000000", + ) + badName := "2026-06-14T010203.000000001Z-peer-a1b2c3-" + strings.Repeat("0", 64) + + w := artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+origin+"/meta/"+url.PathEscape(badName), + metadataBody, "secret", + ) + assertStatus(t, w, http.StatusBadRequest) +} + +func TestArtifactPeerPostImportsAndEmitsDataChanged(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + origin := "peer-a1b2c3" + peerDB, err := db.Open(filepath.Join(t.TempDir(), "peer.db")) + require.NoError(t, err) + t.Cleanup(func() { peerDB.Close() }) + + first := "hello" + started := "2026-06-14T01:02:03Z" + ended := "2026-06-14T01:03:03Z" + dbtest.SeedSession(t, peerDB, "sess-1", "alpha", func(s *db.Session) { + s.MessageCount = 2 + s.UserMessageCount = 1 + s.FirstMessage = &first + s.StartedAt = &started + s.EndedAt = &ended + }) + require.NoError(t, peerDB.ReplaceSessionMessages("sess-1", []db.Message{ + {SessionID: "sess-1", Ordinal: 0, Role: "user", Content: "hello", ContentLength: 5}, + {SessionID: "sess-1", Ordinal: 1, Role: "assistant", Content: "world", ContentLength: 5}, + })) + peerStore, _ := exportArtifactFixture(t, context.Background(), peerDB, origin) + + postArtifactRef(t, te, peerStore, oneArtifactRef(t, peerStore, origin, artifact.KindSegments)) + postArtifactRef(t, te, peerStore, oneArtifactRef(t, peerStore, origin, artifact.KindManifests)) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + req := httptest.NewRequest(http.MethodGet, "/api/v1/events", nil).WithContext(ctx) + stream := &flushRecorder{ResponseRecorder: httptest.NewRecorder()} + done := make(chan struct{}) + go func() { + te.handler.ServeHTTP(stream, req) + close(done) + }() + time.Sleep(100 * time.Millisecond) + + postArtifactRef(t, te, peerStore, oneArtifactRef(t, peerStore, origin, artifact.KindCheckpoints)) + te.waitForSSEEvent(t, stream, "data_changed", 3*time.Second) + + // Live clients only refresh the session index on the "sessions" + // scope and only invalidate hydrated session details on the + // "messages" scope; an import needs both. + assert.Eventually(t, func() bool { + scopes := dataChangedScopes(stream) + return scopes["messages"] && scopes["sessions"] + }, 3*time.Second, 10*time.Millisecond, + "import must emit data_changed with both messages and sessions scopes") + + got, err := te.db.GetSession(context.Background(), origin+"~sess-1") + require.NoError(t, err) + require.NotNil(t, got) + assert.Equal(t, origin, got.Machine) + assert.Equal(t, "alpha", got.Project) + + cancel() + <-done +} + +func TestArtifactPeerDeferredBatchImportsOnlyAtFinalize(t *testing.T) { + te := setupArtifact(t, withArtifactOrigin("desktop-d4e5f6")) + const origin = "peer-a1b2c3" + peerDB, err := db.Open(filepath.Join(t.TempDir(), "peer.db")) + require.NoError(t, err) + t.Cleanup(func() { peerDB.Close() }) + first := "batched" + dbtest.SeedSession(t, peerDB, "sess-1", "alpha", func(s *db.Session) { + s.FirstMessage = &first + }) + peerStore, _ := exportArtifactFixture(t, context.Background(), peerDB, origin) + + for _, kind := range []artifact.Kind{ + artifact.KindSegments, artifact.KindManifests, artifact.KindCheckpoints, + } { + postArtifactRefDeferred(t, te, peerStore, oneArtifactRef(t, peerStore, origin, kind)) + } + got, err := te.db.GetSession(context.Background(), origin+"~sess-1") + require.NoError(t, err) + assert.Nil(t, got, "deferred uploads must not repeatedly import partial batches") + + w := artifactPeerRequest( + t, te, http.MethodPost, "/api/v1/artifacts/finalize", nil, "", + ) + assertStatus(t, w, http.StatusOK) + got, err = te.db.GetSession(context.Background(), origin+"~sess-1") + require.NoError(t, err) + require.NotNil(t, got, "finalize must import the completed artifact batch") + assert.Equal(t, "alpha", got.Project) +} + +// dataChangedScopes collects the scope payloads of every data_changed +// event written to the SSE stream so far. +func dataChangedScopes(w *flushRecorder) map[string]bool { + scopes := make(map[string]bool) + for _, e := range parseSSE(w.BodyString()) { + if e.Event != "data_changed" { + continue + } + var payload struct { + Scope string `json:"scope"` + } + if json.Unmarshal([]byte(e.Data), &payload) == nil { + scopes[payload.Scope] = true + } + } + return scopes +} + +func artifactPeerRequest( + t *testing.T, + te *testEnv, + method string, + path string, + body []byte, + token string, +) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(method, path, bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/octet-stream") + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + w := httptest.NewRecorder() + te.handler.ServeHTTP(w, req) + return w +} + +func postArtifactRef( + t *testing.T, te *testEnv, store artifact.ArtifactStore, ref artifact.Ref, +) { + t.Helper() + wire, body := wireArtifact(t, store, ref) + w := artifactPeerRequest( + t, te, http.MethodPost, + "/api/v1/artifacts/"+ref.Origin+"/"+string(ref.Kind)+"/"+url.PathEscape(wire.Name), + body, "", + ) + assertStatus(t, w, http.StatusOK) +} + +func postArtifactBodyDeferred( + t *testing.T, te *testEnv, wire artifact.WireRef, body []byte, +) { + t.Helper() + req := httptest.NewRequest( + http.MethodPost, + "/api/v1/artifacts/"+wire.Origin+"/"+string(wire.Kind)+"/"+url.PathEscape(wire.Name), + bytes.NewReader(body), + ) + req.Header.Set("Content-Type", "application/octet-stream") + req.Header.Set("X-Agentsview-Artifact-Import", "deferred") + w := httptest.NewRecorder() + te.handler.ServeHTTP(w, req) + assertStatus(t, w, http.StatusOK) +} + +func postArtifactRefDeferred( + t *testing.T, te *testEnv, store artifact.ArtifactStore, ref artifact.Ref, +) { + t.Helper() + wire, body := wireArtifact(t, store, ref) + postArtifactBodyDeferred(t, te, wire, body) +} + +func peerMetadataArtifact(origin, hlc string) ([]byte, string) { + body := []byte(`{"hlc":"` + hlc + `","op":"rename","origin":"` + origin + `","session_gid":"` + origin + `~sess-1","v":1,"value":{"display_name":"Remote"}}` + "\n") + sum := sha256.Sum256(body) + hash := hex.EncodeToString(sum[:]) + return body, hlc + "-" + hash + ".json" +} diff --git a/internal/server/artifact_store_cursors.go b/internal/server/artifact_store_cursors.go new file mode 100644 index 000000000..c53c5f6b6 --- /dev/null +++ b/internal/server/artifact_store_cursors.go @@ -0,0 +1,368 @@ +package server + +import ( + "context" + "errors" + "fmt" + "io" + "io/fs" + "sync" + "time" + + "go.kenn.io/agentsview/internal/artifact" +) + +type artifactStoreCursorRegistry struct { + mu sync.Mutex + tokens map[string]*artifactStoreCursorLease + active int + closed bool +} + +type artifactStoreCursorLease struct { + registry *artifactStoreCursorRegistry + store artifact.ArtifactStore + originIterator artifact.OriginIterator + entryIterator artifact.EntryIterator + scope string + kindIndex int + token string + timer *time.Timer + released bool + registered bool +} + +func newArtifactStoreCursorRegistry() *artifactStoreCursorRegistry { + return &artifactStoreCursorRegistry{tokens: make(map[string]*artifactStoreCursorLease)} +} + +func (s *Server) currentArtifactStoreCursorRegistry() *artifactStoreCursorRegistry { + s.artifactStoreCursorsMu.Lock() + defer s.artifactStoreCursorsMu.Unlock() + return s.artifactStoreCursors +} + +func (s *Server) replaceArtifactStoreCursorRegistry() { + s.artifactStoreCursorsMu.Lock() + previous := s.artifactStoreCursors + replacement := newArtifactStoreCursorRegistry() + s.artifactStoreCursors = replacement + closed := s.artifactStoreCursorsClosed + s.artifactStoreCursorsMu.Unlock() + if previous != nil { + previous.close() + } + if closed { + replacement.close() + } +} + +func (s *Server) closeArtifactStoreCursorRegistry() { + s.artifactStoreCursorsMu.Lock() + registry := s.artifactStoreCursors + s.artifactStoreCursorsClosed = true + s.artifactStoreCursorsMu.Unlock() + if registry != nil { + registry.close() + } +} + +func (r *artifactStoreCursorRegistry) claim( + token, scope string, +) (*artifactStoreCursorLease, error) { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed { + return nil, fs.ErrClosed + } + lease, ok := r.tokens[token] + if !ok || lease.scope != scope || lease.released { + return nil, fmt.Errorf("%w: invalid or expired artifact cursor", artifact.ErrArtifactInvalid) + } + delete(r.tokens, token) + lease.token = "" + if lease.timer != nil { + lease.timer.Stop() + lease.timer = nil + } + return lease, nil +} + +func (r *artifactStoreCursorRegistry) retain( + lease *artifactStoreCursorLease, +) (string, error) { + token, err := artifactCursorToken() + if err != nil { + return "", err + } + r.mu.Lock() + defer r.mu.Unlock() + if r.closed || lease == nil || lease.released || lease.token != "" { + return "", fs.ErrClosed + } + if !lease.registered && r.active >= maxArtifactCursors { + return "", fmt.Errorf("%w: too many active artifact cursors", artifact.ErrArtifactConflict) + } + if !lease.registered { + lease.registered = true + r.active++ + } + lease.registry = r + lease.token = token + r.tokens[token] = lease + lease.timer = time.AfterFunc(artifactCursorTTL, func() { r.release(token) }) + return token, nil +} + +func (r *artifactStoreCursorRegistry) release(token string) bool { + r.mu.Lock() + lease, ok := r.tokens[token] + if !ok || lease.released || lease.token != token { + r.mu.Unlock() + return false + } + delete(r.tokens, token) + lease.token = "" + lease.released = true + if lease.registered { + lease.registered = false + r.active-- + } + if lease.timer != nil { + lease.timer.Stop() + lease.timer = nil + } + r.mu.Unlock() + _ = lease.closeIterators() + return true +} + +func (r *artifactStoreCursorRegistry) abandon(lease *artifactStoreCursorLease) error { + if lease == nil { + return nil + } + r.mu.Lock() + if lease.released { + r.mu.Unlock() + return nil + } + lease.released = true + if lease.registered { + lease.registered = false + r.active-- + } + if lease.token != "" { + delete(r.tokens, lease.token) + lease.token = "" + } + if lease.timer != nil { + lease.timer.Stop() + lease.timer = nil + } + r.mu.Unlock() + return lease.closeIterators() +} + +func (r *artifactStoreCursorRegistry) close() { + r.mu.Lock() + if r.closed { + r.mu.Unlock() + return + } + r.closed = true + leases := make([]*artifactStoreCursorLease, 0, len(r.tokens)) + for token, lease := range r.tokens { + delete(r.tokens, token) + lease.token = "" + lease.released = true + if lease.registered { + lease.registered = false + r.active-- + } + if lease.timer != nil { + lease.timer.Stop() + lease.timer = nil + } + leases = append(leases, lease) + } + r.mu.Unlock() + for _, lease := range leases { + _ = lease.closeIterators() + } +} + +func (l *artifactStoreCursorLease) closeIterators() error { + if l == nil { + return nil + } + var err error + if l.originIterator != nil { + err = errors.Join(err, l.originIterator.Close()) + l.originIterator = nil + } + if l.entryIterator != nil { + err = errors.Join(err, l.entryIterator.Close()) + l.entryIterator = nil + } + return err +} + +func (r *artifactStoreCursorRegistry) originPage( + ctx context.Context, + store artifact.ArtifactStore, + token string, + limit int, +) (_ []string, next string, retErr error) { + return r.originPageForScope(ctx, store, token, limit, "origins") +} + +func (r *artifactStoreCursorRegistry) peerOriginPage( + ctx context.Context, + store artifact.ArtifactStore, + token string, + limit int, +) (_ []string, next string, retErr error) { + return r.originPageForScope(ctx, store, token, limit, "peers") +} + +func (r *artifactStoreCursorRegistry) originPageForScope( + ctx context.Context, + store artifact.ArtifactStore, + token string, + limit int, + scope string, +) (_ []string, next string, retErr error) { + var lease *artifactStoreCursorLease + if token != "" { + var err error + lease, err = r.claim(token, scope) + if err != nil { + return nil, "", err + } + } else { + iterator, err := store.Origins(ctx) + if err != nil { + return nil, "", err + } + lease = &artifactStoreCursorLease{ + store: store, scope: scope, originIterator: iterator, + } + } + defer func() { + if retErr != nil { + retErr = errors.Join(retErr, r.abandon(lease)) + } + }() + origins, err := lease.originIterator.Next(ctx, limit) + done := errors.Is(err, io.EOF) + if err != nil && !done { + return nil, "", err + } + if len(origins) > limit { + return nil, "", fmt.Errorf("%w: artifact origin page exceeds requested limit", artifact.ErrArtifactInvalid) + } + if done { + return origins, "", r.abandon(lease) + } + next, err = r.retain(lease) + if err != nil { + return nil, "", errors.Join(err, r.abandon(lease)) + } + return origins, next, nil +} + +func (r *artifactStoreCursorRegistry) indexPage( + ctx context.Context, + store artifact.ArtifactStore, + origin, token string, + limit int, +) (_ artifact.OriginArtifactIndex, next string, retErr error) { + scope := "index:" + origin + var lease *artifactStoreCursorLease + if token != "" { + var err error + lease, err = r.claim(token, scope) + if err != nil { + return artifact.OriginArtifactIndex{}, "", err + } + } else { + lease = &artifactStoreCursorLease{store: store, scope: scope} + } + defer func() { + if retErr != nil { + retErr = errors.Join(retErr, r.abandon(lease)) + } + }() + index := artifact.OriginArtifactIndex{Origin: origin} + remaining := limit + for lease.kindIndex < len(artifactPeerKinds) && remaining > 0 { + kind := artifactPeerKinds[lease.kindIndex] + if lease.entryIterator == nil { + var err error + lease.entryIterator, err = lease.store.Entries(ctx, origin, kind) + if err != nil { + return artifact.OriginArtifactIndex{}, "", err + } + } + items, err := lease.entryIterator.Next(ctx, remaining) + done := errors.Is(err, io.EOF) + if err != nil && !done { + return artifact.OriginArtifactIndex{}, "", err + } + if len(items) > remaining { + return artifact.OriginArtifactIndex{}, "", fmt.Errorf( + "%w: artifact index page exceeds requested limit", artifact.ErrArtifactInvalid, + ) + } + if len(items) == 0 && !done { + return artifact.OriginArtifactIndex{}, "", fmt.Errorf( + "%w: artifact iterator made no progress", artifact.ErrArtifactInvalid, + ) + } + for _, entry := range items { + wire, err := artifact.ToWireRef(entry.Ref) + if err != nil { + return artifact.OriginArtifactIndex{}, "", err + } + if err := appendArtifactWireName(&index, kind, wire.Name); err != nil { + return artifact.OriginArtifactIndex{}, "", err + } + } + remaining -= len(items) + if done { + if err := lease.entryIterator.Close(); err != nil { + return artifact.OriginArtifactIndex{}, "", err + } + lease.entryIterator = nil + lease.kindIndex++ + } + } + if lease.kindIndex == len(artifactPeerKinds) { + return index, "", r.abandon(lease) + } + next, err := r.retain(lease) + if err != nil { + return artifact.OriginArtifactIndex{}, "", errors.Join(err, r.abandon(lease)) + } + return index, next, nil +} + +func appendArtifactWireName( + index *artifact.OriginArtifactIndex, kind artifact.Kind, name string, +) error { + switch kind { + case artifact.KindSegments: + index.Segments = append(index.Segments, name) + case artifact.KindRaw: + index.Raw = append(index.Raw, name) + case artifact.KindManifests: + index.Manifests = append(index.Manifests, name) + case artifact.KindMeta: + index.Meta = append(index.Meta, name) + case artifact.KindCheckpoints: + index.Checkpoints = append(index.Checkpoints, name) + default: + return fmt.Errorf("%w: unsupported artifact kind %q", artifact.ErrArtifactInvalid, kind) + } + return nil +} diff --git a/internal/server/artifact_store_cursors_internal_test.go b/internal/server/artifact_store_cursors_internal_test.go new file mode 100644 index 000000000..c471fe55b --- /dev/null +++ b/internal/server/artifact_store_cursors_internal_test.go @@ -0,0 +1,204 @@ +package server + +import ( + "context" + "errors" + "fmt" + "io/fs" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" +) + +func TestArtifactStoreCursorRegistryCapsRetainedSnapshots(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 2, + } + tokens := make([]string, 0, maxArtifactCursors) + for range maxArtifactCursors { + _, token, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.NoError(t, err) + require.NotEmpty(t, token) + tokens = append(tokens, token) + } + + _, token, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.ErrorIs(t, err, artifact.ErrArtifactConflict) + assert.Empty(t, token) + assert.Equal(t, maxArtifactCursors, registry.active) + for _, token := range tokens { + assert.True(t, registry.release(token)) + } + assert.Zero(t, registry.active) +} + +func TestArtifactStoreCursorRegistryFinalPageClosesIterator(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 1, + } + + origins, token, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.NoError(t, err) + assert.Equal(t, []string{"origin-00000-a1b2c3"}, origins) + assert.Empty(t, token) + assert.Equal(t, int32(1), store.releaseCalls.Load()) + assert.Zero(t, registry.active) +} + +func TestArtifactStoreCursorRegistryReplacementClosesRetainedIterator(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 2, + } + _, token, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.NoError(t, err) + require.NotEmpty(t, token) + server := &Server{artifactStoreCursors: registry} + + server.replaceArtifactStoreCursorRegistry() + + assert.Equal(t, int32(1), store.releaseCalls.Load()) + assert.Zero(t, registry.active) + assert.NotSame(t, registry, server.artifactStoreCursors) +} + +func TestArtifactStoreCursorRegistryExpiryReleasesAndRejectsReplay(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 2, + } + _, token, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.NoError(t, err) + require.NotEmpty(t, token) + + registry.mu.Lock() + require.NotNil(t, registry.tokens[token]) + registry.tokens[token].timer.Reset(time.Millisecond) + registry.mu.Unlock() + assert.Eventually(t, func() bool { return store.releaseCalls.Load() == 1 }, + time.Second, time.Millisecond) + _, _, err = registry.peerOriginPage(t.Context(), store, token, 1) + require.ErrorIs(t, err, artifact.ErrArtifactInvalid) + assert.Zero(t, registry.active) +} + +func TestArtifactStoreCursorRegistryTokensAreSingleUseAndScopeBound(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 3, + } + _, first, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.NoError(t, err) + _, second, err := registry.peerOriginPage(t.Context(), store, first, 1) + require.NoError(t, err) + require.NotEmpty(t, second) + _, _, err = registry.peerOriginPage(t.Context(), store, first, 1) + require.ErrorIs(t, err, artifact.ErrArtifactInvalid, + "a consumed token must not be replayable") + + for _, scope := range []string{ + "origins", "index:other-a1b2c3", "index:origin-a1b2c3:segments", + } { + _, err := registry.claim(second, scope) + require.ErrorIs(t, err, artifact.ErrArtifactInvalid, scope) + } + _, continued, err := registry.peerOriginPage(t.Context(), store, second, 1) + require.NoError(t, err, "wrong-scope attempts must not consume the peer token") + assert.Empty(t, continued) +} + +type failingArtifactStoreCursor struct { + *pagedLifecycleStore + err error +} + +type failingArtifactOriginIterator struct { + store *pagedLifecycleStore + err error + closed bool +} + +func (i *failingArtifactOriginIterator) Next( + context.Context, int, +) ([]string, error) { + return nil, i.err +} + +func (i *failingArtifactOriginIterator) Close() error { + if !i.closed { + i.closed = true + i.store.releaseCalls.Add(1) + } + return nil +} + +func (s *failingArtifactStoreCursor) Origins( + context.Context, +) (artifact.OriginIterator, error) { + return &failingArtifactOriginIterator{store: s.pagedLifecycleStore, err: s.err}, nil +} + +func (s *failingArtifactStoreCursor) ListOrigins( + context.Context, artifact.Cursor, int, +) ([]string, artifact.Cursor, error) { + return nil, artifact.Cursor("owned-on-error"), s.err +} + +func TestArtifactStoreCursorRegistryReleasesCursorsReturnedWithFailure(t *testing.T) { + for _, failure := range []error{errors.New("fill failure"), context.Canceled} { + t.Run(failure.Error(), func(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &failingArtifactStoreCursor{ + pagedLifecycleStore: &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), + }, + err: failure, + } + _, _, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.ErrorIs(t, err, failure) + assert.Equal(t, int32(1), store.releaseCalls.Load()) + assert.Zero(t, registry.active) + }) + } +} + +func TestArtifactStoreCursorRegistryShutdownDefersClaimedPageRelease(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), originCount: 2, + } + _, token, err := registry.peerOriginPage(t.Context(), store, "", 1) + require.NoError(t, err) + lease, err := registry.claim(token, "peers") + require.NoError(t, err) + + registry.close() + assert.Zero(t, store.releaseCalls.Load(), + "shutdown must not release a cursor while its page is claimed") + err = registry.abandon(lease) + require.NoError(t, err) + assert.Equal(t, int32(1), store.releaseCalls.Load()) + _, _, err = registry.peerOriginPage(t.Context(), store, "", 1) + require.ErrorIs(t, err, fs.ErrClosed) + assert.Zero(t, registry.active) +} + +func TestArtifactStoreCursorRegistryRejectsCrossOriginIndexContinuation(t *testing.T) { + registry := newArtifactStoreCursorRegistry() + store := &pagedLifecycleStore{ + lifecycleArtifactStore: newLifecycleArtifactStore(), entryCount: 2, + } + _, token, err := registry.indexPage(t.Context(), store, "origin-a1b2c3", "", 1) + require.NoError(t, err) + require.NotEmpty(t, token, "index fixture did not retain a cursor") + _, _, err = registry.indexPage(t.Context(), store, "other-a1b2c3", token, 1) + require.ErrorIs(t, err, artifact.ErrArtifactInvalid) + assert.True(t, registry.release(token), fmt.Sprintf("release token %q", token)) + assert.Equal(t, int32(1), store.entryCloseCalls.Load()) +} diff --git a/internal/server/auth.go b/internal/server/auth.go index a76e6934b..5e77b6cab 100644 --- a/internal/server/auth.go +++ b/internal/server/auth.go @@ -103,6 +103,9 @@ func (s *Server) authMiddleware(next http.Handler) http.Handler { } remoteSyncAuth := isRemoteSyncPath(r.URL.Path) + if isArtifactDaemonMutationPath(r.URL.Path) && !s.hasWritableArtifactStore() { + remoteSyncAuth = false + } // When auth is not required, skip token checks entirely // except for machine-to-machine remote sync archive APIs, @@ -164,7 +167,14 @@ func isSSEPath(path string) bool { } func isRemoteSyncPath(path string) bool { - return strings.HasPrefix(path, "/api/v1/remote-sync/") + return strings.HasPrefix(path, "/api/v1/remote-sync/") || + isArtifactDaemonMutationPath(path) +} + +func isArtifactDaemonMutationPath(path string) bool { + return path == "/api/v1/artifacts/exchange" || + path == "/api/v1/artifacts/maintenance" || + path == "/api/v1/artifacts/reset" } // setCORSOnAuthError adds CORS headers to 401 responses so diff --git a/internal/server/bulk_star_metadata_test.go b/internal/server/bulk_star_metadata_test.go new file mode 100644 index 000000000..1b3ca6d8f --- /dev/null +++ b/internal/server/bulk_star_metadata_test.go @@ -0,0 +1,66 @@ +package server_test + +import ( + "context" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" +) + +func TestBulkStarAppendsMetadataEvents(t *testing.T) { + te := setup(t, withArtifactOrigin("desktop-d4e5f6")) + te.seedSession(t, "s1", "alpha", 2) + te.seedSession(t, "s2", "beta", 2) + + w := te.requestJSON(t, http.MethodPost, "/api/v1/starred/bulk", + `{"session_ids":["s1","s2","missing"]}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + // Both existing sessions are starred; the missing one is skipped. + w = te.get(t, "/api/v1/starred") + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + list := decode[starredHandlerResponse](t, w) + assert.ElementsMatch(t, []string{"s1", "s2"}, list.SessionIDs) + + // A star metadata event artifact was written for each session actually + // starred, so the migrated stars converge through artifact sync. The missing + // session produces no event. + assert.Len(t, readMetadataEvents(t, te), 2, + "one star event per existing session") +} + +func TestBulkStarRetriesRepairPublishedMetadataState(t *testing.T) { + te := setup(t, withArtifactOrigin("desktop-d4e5f6")) + te.seedSession(t, "s1", "alpha", 2) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w := te.requestJSON(t, http.MethodPost, "/api/v1/starred/bulk", + `{"session_ids":["s1"]}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"s1"}, ids) + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_replay_state", "session_gid = 'desktop-d4e5f6~s1'")) + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpStar, events[0].Op) + + execTestDDL(t, te, `DROP TRIGGER fail_metadata_replay_state_insert`) + w = te.requestJSON(t, http.MethodPost, "/api/v1/starred/bulk", + `{"session_ids":["s1"]}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpStar, + serverMetadataReplayOp(t, te, "desktop-d4e5f6~s1", "starred")) + assert.Len(t, readMetadataEvents(t, te), 1) +} diff --git a/internal/server/compress.go b/internal/server/compress.go index bbde74f7e..d303d3d9e 100644 --- a/internal/server/compress.go +++ b/internal/server/compress.go @@ -213,6 +213,8 @@ func isStreamingPath(path string) bool { return true case strings.HasPrefix(path, "/api/v1/import/"): return true + case strings.HasPrefix(path, "/api/v1/artifacts/"): + return true default: return false } diff --git a/internal/server/huma_route_groups.go b/internal/server/huma_route_groups.go index 1f9c99657..2056776b9 100644 --- a/internal/server/huma_route_groups.go +++ b/internal/server/huma_route_groups.go @@ -27,6 +27,7 @@ func (s *Server) registerTypedAPIRoutes() { s.registerImportRoutes() s.registerAssetRoutes() s.registerEmbeddingsRoutes() + s.registerArtifactRoutes() } type routeGroup struct { diff --git a/internal/server/huma_routes_artifacts.go b/internal/server/huma_routes_artifacts.go new file mode 100644 index 000000000..e2bd9f005 --- /dev/null +++ b/internal/server/huma_routes_artifacts.go @@ -0,0 +1,1282 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "net/http" + "os" + "reflect" + "strconv" + "time" + + "github.com/danielgtaylor/huma/v2" + + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/db" +) + +func (s *Server) registerArtifactRoutes() { + if s.artifactStore == nil && !s.documentArtifactRoutes { + return + } + group := newRouteGroup(s.api, "/api/v1/artifacts", "Artifacts") + + get(s, group, "/origins", "List artifact origins", s.humaListArtifactOrigins) + get(s, group, "/peers", "List artifact peers", s.humaListArtifactPeers) + get(s, group, "/{origin}/index", "List artifact index for an origin", s.humaGetArtifactIndex) + writable := s.documentArtifactRoutes || s.hasWritableArtifactStore() + if writable { + post(s, group, "/finalize", "Finalize artifact uploads", s.humaFinalizeArtifacts) + } + s.documentArtifactStreamingRoutes(writable) + s.mux.HandleFunc("GET /api/v1/artifacts/{origin}/checkpoint", s.getArtifactCheckpointHTTP) + s.mux.HandleFunc("GET /api/v1/artifacts/{origin}/{kind}/{name}", s.getArtifactHTTP) + if writable { + s.mux.HandleFunc("POST /api/v1/artifacts/{origin}/{kind}/{name}", s.postArtifactHTTP) + s.mux.HandleFunc("POST /api/v1/artifacts/exchange", s.postArtifactExchangeHTTP) + s.mux.HandleFunc("POST /api/v1/artifacts/maintenance", s.postArtifactMaintenanceHTTP) + s.mux.HandleFunc("POST /api/v1/artifacts/reset", s.postArtifactResetHTTP) + } + s.mux.HandleFunc("DELETE /api/v1/artifacts/cursors/{cursor}", s.releaseArtifactCursorHTTP) +} + +func (s *Server) documentArtifactStreamingRoutes(documentPost bool) { + binary := &huma.Schema{Type: huma.TypeString, Format: "binary"} + stringParam := func(name, description string) *huma.Param { + return &huma.Param{ + Name: name, In: "path", Description: description, Required: true, + Schema: &huma.Schema{Type: huma.TypeString}, + } + } + originParam := func() *huma.Param { return stringParam("origin", "Artifact origin ID") } + + s.api.OpenAPI().AddOperation(&huma.Operation{ + OperationID: operationID(http.MethodGet, "/api/v1/artifacts/{origin}/checkpoint"), + Method: http.MethodGet, + Path: "/api/v1/artifacts/{origin}/checkpoint", + Tags: []string{"Artifacts"}, + Summary: "Get latest artifact checkpoint", + Parameters: []*huma.Param{originParam()}, + Responses: artifactStreamingResponses(&huma.Response{ + Description: "OK", + Content: map[string]*huma.MediaType{ + "application/octet-stream": {Schema: binary}, + }, + }), + }) + + path := "/api/v1/artifacts/{origin}/{kind}/{name}" + params := func() []*huma.Param { + return []*huma.Param{ + originParam(), + stringParam("kind", "Artifact kind"), + stringParam("name", "Artifact filename or hash"), + } + } + s.api.OpenAPI().AddOperation(&huma.Operation{ + OperationID: operationID(http.MethodGet, path), + Method: http.MethodGet, + Path: path, + Tags: []string{"Artifacts"}, + Summary: "Get artifact", + Parameters: params(), + Responses: artifactStreamingResponses(&huma.Response{ + Description: "OK", + Content: map[string]*huma.MediaType{ + "application/octet-stream": {Schema: binary}, + }, + }), + }) + if !documentPost { + return + } + postSchema := s.api.OpenAPI().Components.Schemas.Schema( + reflect.TypeFor[artifactPostResponse](), true, "ArtifactPostResponse", + ) + s.api.OpenAPI().AddOperation(&huma.Operation{ + OperationID: operationID(http.MethodPost, path), + Method: http.MethodPost, + Path: path, + Tags: []string{"Artifacts"}, + Summary: "Post artifact", + Parameters: params(), + RequestBody: &huma.RequestBody{ + Required: true, + Content: map[string]*huma.MediaType{ + "application/octet-stream": {Schema: binary}, + }, + }, + Responses: artifactStreamingResponses(&huma.Response{ + Description: "OK", + Content: map[string]*huma.MediaType{ + "application/json": {Schema: postSchema}, + }, + }), + }) +} + +func artifactStreamingResponses(success *huma.Response) map[string]*huma.Response { + responses := map[string]*huma.Response{"200": success} + for _, status := range []int{ + http.StatusBadRequest, + http.StatusUnauthorized, + http.StatusForbidden, + http.StatusNotFound, + http.StatusConflict, + http.StatusInternalServerError, + http.StatusNotImplemented, + http.StatusBadGateway, + http.StatusServiceUnavailable, + http.StatusGatewayTimeout, + } { + responses[strconv.Itoa(status)] = &huma.Response{ + Description: http.StatusText(status), + Content: map[string]*huma.MediaType{ + "text/plain": {Schema: &huma.Schema{Type: huma.TypeString}}, + }, + } + } + return responses +} + +type artifactOriginsInput struct { + Cursor string `query:"cursor" doc:"Opaque artifact origin cursor"` + Limit int `query:"limit" minimum:"1" maximum:"512" default:"512" doc:"Maximum origins to return"` +} + +type artifactIndexInput struct { + Origin string `path:"origin" required:"true" doc:"Artifact origin ID"` + Cursor string `query:"cursor" doc:"Opaque artifact index cursor"` + Limit int `query:"limit" minimum:"1" maximum:"512" default:"512" doc:"Maximum artifact names to return"` +} + +type artifactIndexResponse struct { + artifact.OriginArtifactIndex + NextCursor string `json:"next_cursor,omitempty"` +} + +type artifactPostInput struct { + Origin string + Kind string + Name string + ImportMode string + Body io.Reader +} + +type artifactFinalizeResponse struct { + ImportedSessions int `json:"imported_sessions"` + ImportedMessages int `json:"imported_messages"` + ImportedMetadata int `json:"imported_metadata"` + Deferred int `json:"deferred"` +} + +type artifactOriginsResponse struct { + Origins []string `json:"origins"` + NextCursor string `json:"next_cursor,omitempty"` +} + +// artifactPeer is one origin's status in the peers view: what it has published +// (from its latest checkpoint) and how much of it has landed locally. +type artifactPeer struct { + Origin string `json:"origin"` + IsLocal bool `json:"is_local"` + CheckpointSeq int `json:"checkpoint_seq"` + PublishedSessions int `json:"published_sessions"` + LocalSessions int `json:"local_sessions"` + LastPublished string `json:"last_published,omitempty"` + Status string `json:"status"` +} + +type artifactPeersResponse struct { + LocalOrigin string `json:"local_origin"` + Peers []artifactPeer `json:"peers"` + ConflictCount int `json:"conflict_count"` + PendingImports int `json:"pending_imports"` + OldestPendingAt string `json:"oldest_pending_at,omitempty"` + NextCursor string `json:"next_cursor,omitempty"` +} + +type artifactPeersInput struct { + Cursor string `query:"cursor" doc:"Opaque peer origin cursor"` + Limit int `query:"limit" minimum:"1" maximum:"512" default:"512" doc:"Maximum peer origins to return"` +} + +type artifactPostResponse struct { + Origin string `json:"origin"` + Kind string `json:"kind"` + Name string `json:"name"` + Hash string `json:"hash,omitempty"` + Size int64 `json:"size"` + Duplicate bool `json:"duplicate"` +} + +type artifactExchangeRequest struct { + Target string `json:"target"` + Token string `json:"token,omitempty"` + AllowInsecure bool `json:"allow_insecure,omitempty"` + BaselineMetadata bool `json:"baseline_metadata,omitempty"` +} + +type artifactExchangeResponse struct { + Origin string `json:"origin"` + ExportedSessions int `json:"exported_sessions"` + ImportedSessions int `json:"imported_sessions"` + ImportedMessages int `json:"imported_messages"` + ImportedMetadata int `json:"imported_metadata"` +} + +type artifactMaintenanceRequest struct { + Grace string `json:"grace,omitempty"` + QuarantineGrace string `json:"quarantine_grace,omitempty"` + GraceSeconds *int64 `json:"grace_seconds,omitempty"` + QuarantineGraceSeconds *int64 `json:"quarantine_grace_seconds,omitempty"` + MaxObjects int `json:"max_objects"` + MaxBytes int64 `json:"max_bytes"` + DryRun bool `json:"dry_run"` + TrashCursor string `json:"trash_cursor,omitempty"` + GCCursor string `json:"gc_cursor,omitempty"` + RepackCursor string `json:"repack_cursor,omitempty"` +} + +type artifactResetResponse struct { + artifact.RepositoryResetResult + ManualCleanup string `json:"manual_cleanup"` + ForeignArtifacts string `json:"foreign_artifacts"` +} + +const maxArtifactMaintenanceGraceSeconds = int64(1<<63-1) / int64(time.Second) + +func parseArtifactMaintenanceDuration(exact string, legacySeconds *int64) (time.Duration, error) { + if exact != "" && legacySeconds != nil { + return 0, errors.New("duration and legacy seconds cannot both be set") + } + if exact != "" { + duration, err := time.ParseDuration(exact) + if err != nil || duration < 0 { + return 0, errors.New("invalid duration") + } + return duration, nil + } + if legacySeconds == nil { + return 0, nil + } + if *legacySeconds < 0 || *legacySeconds > maxArtifactMaintenanceGraceSeconds { + return 0, errors.New("invalid legacy duration") + } + return time.Duration(*legacySeconds) * time.Second, nil +} + +type artifactMaintenancePhysicalResponse struct { + Supported bool `json:"supported"` + Result artifact.PhysicalMaintenanceResult `json:"result"` +} + +type artifactMaintenanceResponse struct { + Logical artifact.GCResult `json:"logical"` + Physical artifactMaintenancePhysicalResponse `json:"physical"` +} + +func (s *Server) acquireArtifactStore() (artifact.ArtifactStore, func(), error) { + store, release, err := s.artifactOps.acquire() + if err != nil { + return nil, nil, apiError(http.StatusServiceUnavailable, "artifact store not configured") + } + return store, release, nil +} + +func (s *Server) humaListArtifactOrigins( + ctx context.Context, + in *artifactOriginsInput, +) (*jsonOutput[artifactOriginsResponse], error) { + store, release, err := s.acquireArtifactStore() + if err != nil { + return nil, err + } + defer release() + if err := s.publishLocalArtifacts(ctx, store); err != nil { + return nil, err + } + limit := in.Limit + if limit == 0 { + limit = 512 + } + origins, next, err := s.currentArtifactStoreCursorRegistry().originPage( + ctx, store, in.Cursor, limit, + ) + if err != nil { + return nil, artifactRouteError("list artifact origins", err) + } + return &jsonOutput[artifactOriginsResponse]{ + Body: artifactOriginsResponse{Origins: origins, NextCursor: next}, + }, nil +} + +func (s *Server) humaGetArtifactIndex( + ctx context.Context, + in *artifactIndexInput, +) (*jsonOutput[artifactIndexResponse], error) { + store, release, err := s.acquireArtifactStore() + if err != nil { + return nil, err + } + defer release() + limit := in.Limit + if limit == 0 { + limit = 512 + } + index, next, err := s.currentArtifactStoreCursorRegistry().indexPage( + ctx, store, in.Origin, in.Cursor, limit, + ) + if err != nil { + return nil, artifactRouteError("list artifact index", err) + } + return &jsonOutput[artifactIndexResponse]{Body: artifactIndexResponse{ + OriginArtifactIndex: index, + NextCursor: next, + }}, nil +} + +func (s *Server) releaseArtifactCursorHTTP(w http.ResponseWriter, r *http.Request) { + if err := r.Context().Err(); err != nil { + return + } + _, release, err := s.acquireArtifactStore() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + defer release() + s.artifactCursors.release(r.PathValue("cursor")) + s.currentArtifactStoreCursorRegistry().release(r.PathValue("cursor")) + w.WriteHeader(http.StatusNoContent) +} + +// localArtifactOrigin returns this machine's artifact origin without creating +// one. It prefers the configured origin and falls back to the persisted DB +// value so read-only callers never mint a new identity. +func (s *Server) localArtifactOrigin() string { + if s.cfg.ArtifactOriginID != "" { + return s.cfg.ArtifactOriginID + } + if local, ok := s.db.(*db.DB); ok { + if origin, err := artifact.StoredOrigin(local); err == nil { + return origin + } + } + return "" +} + +func (s *Server) humaListArtifactPeers( + ctx context.Context, + in *artifactPeersInput, +) (*jsonOutput[artifactPeersResponse], error) { + store, release, err := s.acquireArtifactStore() + if err != nil { + return nil, err + } + defer release() + if err := s.publishLocalArtifacts(ctx, store); err != nil { + return nil, err + } + limit := in.Limit + if limit == 0 { + limit = 512 + } + origins, next, err := s.currentArtifactStoreCursorRegistry().peerOriginPage( + ctx, store, in.Cursor, limit, + ) + if err != nil { + return nil, artifactRouteError("list artifact origins", err) + } + conflicts, err := s.db.CountMetadataConflicts(ctx) + if err != nil { + return nil, internalError("count metadata conflicts", err) + } + + localOrigin := s.localArtifactOrigin() + peers := make([]artifactPeer, 0, len(origins)) + for _, origin := range origins { + isLocal := origin == localOrigin + sequence, expected, found, headErr := s.artifactPeerCheckpointHead(ctx, origin, isLocal) + landing := artifact.OriginCheckpointLanding{} + if found { + landing.Sequence = sequence + } + peerStatus := "pending" + if headErr != nil { + peerStatus = "error" + } else if found { + landing, err = artifact.CheckpointLandingStatusAtStoreHead( + ctx, store, origin, sequence, expected, s.db, isLocal, + ) + if err != nil { + landing.Sequence = sequence + } + switch { + case err == nil && landing.LandedSessionCount >= landing.SessionCount: + peerStatus = "in_sync" + case err == nil, errors.Is(err, artifact.ErrArtifactNotFound): + peerStatus = "pending" + default: + peerStatus = "error" + } + } + last := "" + if landing.Found { + last = landing.ModTime.UTC().Format(time.RFC3339) + } + peers = append(peers, artifactPeer{ + Origin: origin, + IsLocal: isLocal, + CheckpointSeq: landing.Sequence, + PublishedSessions: landing.SessionCount, + LocalSessions: landing.LandedSessionCount, + LastPublished: last, + Status: peerStatus, + }) + } + + pendingImports := 0 + oldestPendingAt := "" + if local, ok := s.db.(*db.DB); ok { + pendingImports, oldestPendingAt, err = local.ArtifactImportQueueStats(ctx) + if err != nil { + return nil, artifactRouteError("read artifact import queue", err) + } + } + return &jsonOutput[artifactPeersResponse]{ + Body: artifactPeersResponse{ + LocalOrigin: localOrigin, Peers: peers, ConflictCount: conflicts, + PendingImports: pendingImports, OldestPendingAt: oldestPendingAt, + NextCursor: next, + }, + }, nil +} + +func (s *Server) artifactPeerCheckpointHead( + ctx context.Context, origin string, isLocal bool, +) (int, artifact.Identity, bool, error) { + database, ok := s.db.(interface { + GetArtifactCheckpointHead(context.Context, string) (db.ArtifactCheckpointHead, bool, error) + GetArtifactPeerCheckpointHead(context.Context, string) (db.ArtifactPeerCheckpointHead, bool, error) + GetArtifactCheckpointLandingHead(context.Context, string) (db.ArtifactCheckpointLanding, bool, error) + }) + if !ok { + return 0, artifact.Identity{}, false, nil + } + if isLocal { + head, found, err := database.GetArtifactCheckpointHead(ctx, origin) + if err != nil || !found { + return 0, artifact.Identity{}, found, err + } + identity, err := artifact.NewIdentity(head.CheckpointSHA256, head.CheckpointSize) + return head.Sequence, identity, true, err + } + head, headFound, err := database.GetArtifactPeerCheckpointHead(ctx, origin) + if err != nil { + return 0, artifact.Identity{}, false, err + } + landing, landed, err := database.GetArtifactCheckpointLandingHead(ctx, origin) + if err != nil { + return 0, artifact.Identity{}, false, err + } + if landed && (!headFound || landing.Sequence > head.Sequence) { + return landing.Sequence, artifact.Identity{}, true, nil + } + if !headFound { + return 0, artifact.Identity{}, false, nil + } + identity, err := artifact.NewIdentity(head.CheckpointSHA256, head.CheckpointSize) + return head.Sequence, identity, true, err +} + +// publishLocalArtifacts refreshes the server's owned origin immediately before +// peer discovery. HTTP transports begin every exchange with origin discovery, +// so this makes the server a publisher without requiring a separate folder +// sync process while keeping individual artifact reads side-effect free. +func (s *Server) publishLocalArtifacts(ctx context.Context, store artifact.ArtifactStore) error { + local, ok := s.db.(*db.DB) + if !ok || local.ReadOnly() { + return nil + } + origin := s.localArtifactOrigin() + if origin == "" { + return nil + } + + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + if s.artifactRepository != nil { + _, recovered, err := artifact.RecoverRepositoryResetRepublish( + ctx, local, s.artifactRepository, origin, + ) + if err != nil { + return artifactRouteError("recover artifact repository reset", err) + } + if recovered { + s.artifactBaselineDone = true + } + } + if !s.artifactBaselineDone && s.metadata != nil { + if _, err := s.metadata.AppendBaseline(ctx); err != nil { + return artifactRouteError("baseline local artifact metadata", err) + } + s.artifactBaselineDone = true + } + if s.engine != nil { + s.engine.FlushSignals() + } + var err error + if s.artifactRepository != nil { + _, err = artifact.PublishRepositoryArtifacts( + ctx, local, s.artifactRepository, artifact.ExportOptions{Origin: origin}, + ) + } else { + _, err = artifact.ExportToStore( + ctx, local, store, artifact.ExportOptions{Origin: origin}, + ) + } + if err != nil { + return artifactRouteError("export local artifacts", err) + } + if s.artifactRepository != nil { + s.artifactRepository.NotifyBatch(ctx) + } + return nil +} + +func (s *Server) getArtifactHTTP(w http.ResponseWriter, r *http.Request) { + store, release, err := s.acquireArtifactStore() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + defer release() + spool, err := spoolStoreArtifactForServe( + r.Context(), store, r.PathValue("origin"), r.PathValue("kind"), r.PathValue("name"), + ) + s.serveArtifactSpool(w, r, spool, err) +} + +func (s *Server) getArtifactCheckpointHTTP(w http.ResponseWriter, r *http.Request) { + store, release, err := s.acquireArtifactStore() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + defer release() + spool, err := spoolLatestStoreCheckpointForServe( + r.Context(), store, r.PathValue("origin"), + ) + s.serveArtifactSpool(w, r, spool, err) +} + +func (s *Server) serveArtifactSpool( + w http.ResponseWriter, r *http.Request, spool *artifact.PeerArtifactSpool, err error, +) { + if err != nil { + if r.Context().Err() != nil { + return + } + writeArtifactHTTPError(w, artifactRouteError("get artifact", err)) + return + } + if spool == nil { + writeArtifactHTTPError(w, internalError( + "get artifact", errors.New("artifact response spool is nil"), + )) + return + } + defer func() { + if err := spool.Close(); err != nil { + log.Printf("closing artifact response spool: %v", err) + } + }() + w.Header().Set("Content-Type", "application/octet-stream") + w.Header().Set("X-Content-Type-Options", "nosniff") + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Content-Length", strconv.FormatInt(spool.Size, 10)) + w.WriteHeader(http.StatusOK) + if _, err := io.Copy(w, &artifactHTTPContextReader{ + ctx: r.Context(), reader: spool.File, + }); err != nil && r.Context().Err() == nil { + log.Printf("streaming artifact response: %v", err) + } +} + +type artifactHTTPContextReader struct { + ctx context.Context + reader io.Reader +} + +func (r *artifactHTTPContextReader) Read(p []byte) (int, error) { + if err := r.ctx.Err(); err != nil { + return 0, err + } + return r.reader.Read(p) +} + +var artifactPeerKinds = [...]artifact.Kind{ + artifact.KindSegments, + artifact.KindRaw, + artifact.KindManifests, + artifact.KindMeta, + artifact.KindCheckpoints, +} + +func latestStoreCheckpointEntry( + ctx context.Context, store artifact.ArtifactStore, origin string, +) (_ artifact.Entry, found bool, retErr error) { + iterator, err := store.Entries(ctx, origin, artifact.KindCheckpoints) + if err != nil { + return artifact.Entry{}, false, err + } + defer func() { retErr = errors.Join(retErr, iterator.Close()) }() + var latest artifact.Entry + for { + entries, nextErr := iterator.Next(ctx, 512) + if nextErr != nil && !errors.Is(nextErr, io.EOF) { + return artifact.Entry{}, false, nextErr + } + if len(entries) > 0 { + latest = entries[len(entries)-1] + found = true + } + if errors.Is(nextErr, io.EOF) { + return latest, found, nil + } + } +} + +func spoolStoreArtifactForServe( + ctx context.Context, + store artifact.ArtifactStore, + origin, kind, name string, +) (_ *artifact.PeerArtifactSpool, retErr error) { + ref, err := artifact.FromWireRef(origin, artifact.Kind(kind), name) + if err != nil { + return nil, err + } + entry, reader, err := store.Open(ctx, ref) + if err != nil { + return nil, err + } + defer func() { retErr = errors.Join(retErr, reader.Close()) }() + if entry.Ref != ref { + return nil, fmt.Errorf("%w: artifact store returned the wrong reference", artifact.ErrArtifactCorrupt) + } + response, err := os.CreateTemp("", "agentsview-peer-wire-response-*") + if err != nil { + return nil, err + } + cleanup := true + defer func() { + if cleanup { + retErr = errors.Join(retErr, response.Close(), os.Remove(response.Name())) + } + }() + if err := response.Chmod(0o600); err != nil { + return nil, err + } + if err := artifact.EncodeWire(ctx, ref, reader, response); err != nil { + return nil, err + } + if err := reader.Verify(); err != nil { + return nil, fmt.Errorf("%w: %v", artifact.ErrArtifactCorrupt, err) + } + info, err := response.Stat() + if err != nil { + return nil, err + } + if err := response.Sync(); err != nil { + return nil, err + } + if _, err := response.Seek(0, io.SeekStart); err != nil { + return nil, err + } + cleanup = false + return &artifact.PeerArtifactSpool{ + Origin: origin, Kind: kind, Name: name, + Hash: entry.Identity.SHA256, ContentType: "application/octet-stream", + Size: info.Size(), File: response, + }, nil +} + +func spoolLatestStoreCheckpointForServe( + ctx context.Context, store artifact.ArtifactStore, origin string, +) (_ *artifact.PeerArtifactSpool, retErr error) { + latest, found, err := latestStoreCheckpointEntry(ctx, store, origin) + if err != nil { + return nil, err + } + if !found { + return nil, artifact.ErrArtifactNotFound + } + wire, err := artifact.ToWireRef(latest.Ref) + if err != nil { + return nil, err + } + return spoolStoreArtifactForServe(ctx, store, origin, string(wire.Kind), wire.Name) +} + +type artifactCountingReader struct { + reader io.Reader + read int64 +} + +type serverArtifactImportRetryScheduler struct{ server *Server } + +func (s serverArtifactImportRetryScheduler) RecordChanged( + ctx context.Context, entry artifact.Entry, +) error { + local, ok := s.server.db.(*db.DB) + if !ok || local.ReadOnly() || s.server.artifactStore == nil { + return errors.New("artifact import store is not writable") + } + coordinator := artifact.NewStoreImportCoordinator( + local, s.server.artifactStore, s.server.localArtifactOrigin(), + ) + if err := coordinator.RecordChanged(ctx, entry); err != nil { + return err + } + s.server.lockSessionLifecycle() + s.server.artifactImportPending = true + s.server.sessionLifecycleMu.Unlock() + return nil +} + +func (r *artifactCountingReader) Read(p []byte) (int, error) { + n, err := r.reader.Read(p) + r.read += int64(n) + return n, err +} + +func (s *Server) humaPostArtifact( + ctx context.Context, + in *artifactPostInput, +) (*jsonOutput[artifactPostResponse], error) { + local, err := s.writableArtifactImportDB() + if err != nil { + return nil, err + } + store, release, err := s.acquireArtifactStore() + if err != nil { + return nil, err + } + defer release() + deferred := in.ImportMode == "deferred" + ref, err := artifact.FromWireRef(in.Origin, artifact.Kind(in.Kind), in.Name) + if err != nil { + return nil, artifactRouteError("post artifact", err) + } + wire, err := artifact.ToWireRef(ref) + if err != nil { + return nil, artifactRouteError("post artifact", err) + } + counting := &artifactCountingReader{reader: in.Body} + spool, err := artifact.DecodeWireToCanonicalSpool(ctx, wire, counting, + artifact.PeerWireLimits(wire.Kind)) + if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } + return nil, artifactRouteError("post artifact", err) + } + defer func() { _ = spool.Close() }() + canonicalRef := spool.Ref() + identity := spool.Identity() + repair, repairQueued, err := local.ArtifactRepairForRef( + ctx, canonicalRef.Origin, string(canonicalRef.Kind), canonicalRef.Name, + ) + if err != nil { + return nil, artifactRouteError("post artifact repair lookup", err) + } + duplicate := false + if repairQueued { + if repair.SHA256 != identity.SHA256 || repair.Size != identity.Size || + repair.Origin != canonicalRef.Origin || repair.Kind != string(canonicalRef.Kind) || + repair.Name != canonicalRef.Name { + return nil, artifactRouteError("post artifact repair", + fmt.Errorf("%w: queued repair identity does not match peer content", + artifact.ErrArtifactConflict)) + } + trusted, err := spool.Rewind() + if err != nil { + return nil, artifactRouteError("post artifact repair", err) + } + if err := artifact.RepairArtifactFromTrustedPeer( + ctx, local, store, repair, trusted, + serverArtifactImportRetryScheduler{server: s}, + ); err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } + return nil, artifactRouteError("post artifact repair", err) + } + duplicate = true + } else { + created, err := spool.Create(ctx, store) + if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return nil, ctxErr + } + return nil, artifactRouteError("post artifact", err) + } + duplicate = !created.Created + } + if !duplicate { + coordinator := artifact.NewStoreImportCoordinator( + local, store, s.localArtifactOrigin(), + ) + if err := coordinator.RecordChanged(ctx, artifact.Entry{ + Ref: canonicalRef, Identity: identity, + }); err != nil { + return nil, artifactRouteError("record changed artifact", err) + } + } + res := artifact.PeerArtifactWrite{ + Origin: in.Origin, Kind: in.Kind, Name: in.Name, + Hash: identity.SHA256, Size: counting.read, + Duplicate: duplicate, + } + if deferred { + s.lockSessionLifecycle() + s.artifactImportPending = true + s.sessionLifecycleMu.Unlock() + return artifactPostOutput(res), nil + } + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + drivesImport := canonicalRef.Kind == artifact.KindCheckpoints || + canonicalRef.Kind == artifact.KindMeta + if drivesImport || s.artifactImportPending { + importRes, err := s.importPeerArtifacts(ctx, local, store) + if err != nil { + return nil, err + } + s.artifactImportPending = importRes.Deferred > 0 + } + return artifactPostOutput(res), nil +} + +func (s *Server) postArtifactHTTP(w http.ResponseWriter, r *http.Request) { + output, err := s.humaPostArtifact(r.Context(), &artifactPostInput{ + Origin: r.PathValue("origin"), + Kind: r.PathValue("kind"), + Name: r.PathValue("name"), + ImportMode: r.Header.Get("X-Agentsview-Artifact-Import"), + Body: r.Body, + }) + if err != nil { + writeArtifactHTTPError(w, err) + return + } + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(output.Body); err != nil && r.Context().Err() == nil { + log.Printf("encoding artifact POST response: %v", err) + } +} + +func (s *Server) postArtifactExchangeHTTP(w http.ResponseWriter, r *http.Request) { + if !isLocalhostRequest(r) { + http.Error(w, "artifact exchange is only available from localhost", http.StatusForbidden) + return + } + local, err := s.writableArtifactImportDB() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + store, release, err := s.acquireArtifactStore() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + defer release() + + r.Body = http.MaxBytesReader(w, r.Body, 1<<20) + decoder := json.NewDecoder(r.Body) + decoder.DisallowUnknownFields() + var input artifactExchangeRequest + if err := decoder.Decode(&input); err != nil { + http.Error(w, "invalid artifact exchange request", http.StatusBadRequest) + return + } + if input.Target == "" { + http.Error(w, "artifact exchange target is required", http.StatusBadRequest) + return + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + http.Error(w, "invalid artifact exchange request", http.StatusBadRequest) + return + } + if err := artifact.ValidateSyncTarget(input.Target); err != nil { + writeArtifactHTTPError(w, apiError(http.StatusBadRequest, + "invalid artifact exchange target")) + return + } + + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + syncOpts := artifact.SyncOptions{ + DataDir: s.cfg.DataDir, + Target: input.Target, + Origin: s.localArtifactOrigin(), + Token: input.Token, + AllowInsecure: input.AllowInsecure, + BaselineMetadata: input.BaselineMetadata, + OnDataChanged: func() { + if s.broadcaster != nil { + s.broadcaster.Emit("data_changed") + } + }, + } + var result artifact.SyncResult + if s.artifactRepository != nil { + result, err = artifact.SyncWithRepository( + r.Context(), local, s.artifactRepository, syncOpts, + ) + } else { + result, err = artifact.SyncWithStore(r.Context(), local, store, syncOpts) + } + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return + } + // The target and peer token are caller-provided secrets. Do not log or + // reflect transport errors that may include either value. + writeArtifactHTTPError(w, apiError(http.StatusBadGateway, + "artifact exchange failed")) + return + } + s.artifactBaselineDone = s.artifactBaselineDone || input.BaselineMetadata + s.artifactImportPending = false + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(artifactExchangeResponse{ + Origin: result.Origin, ExportedSessions: result.ExportedSessions, + ImportedSessions: result.ImportedSessions, ImportedMessages: result.ImportedMessages, + ImportedMetadata: result.ImportedMetadata, + }); err != nil && r.Context().Err() == nil { + log.Printf("encoding artifact exchange response: %v", err) + } +} + +func (s *Server) postArtifactResetHTTP(w http.ResponseWriter, r *http.Request) { + if !isLocalhostRequest(r) { + http.Error(w, "artifact reset is only available from localhost", http.StatusForbidden) + return + } + local, err := s.writableArtifactImportDB() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + ownedStore, resetCtx, err := s.artifactOps.beginReset(r.Context()) + if err != nil { + writeArtifactHTTPError(w, apiError(http.StatusServiceUnavailable, err.Error())) + return + } + current := s.artifactRepository + if current == nil { + _ = s.artifactOps.finishReset(ownedStore) + writeArtifactHTTPError(w, apiError(http.StatusNotImplemented, + "artifact reset requires the local AgentsView repository")) + return + } + + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + origin := s.localArtifactOrigin() + var pending db.ArtifactResetRepublishPending + pendingPrepared := false + fresh, result, resetErr := s.beginArtifactRepositoryReset( + resetCtx, s.cfg.DataDir, origin, current, + func() error { + if origin != "" { + var err error + pending, err = artifact.PrepareRepositoryResetRepublish( + resetCtx, local, s.cfg.DataDir, origin, + ) + if err != nil { + return err + } + pendingPrepared = true + } + commitErr := s.artifactOps.commitResetMutation(resetCtx) + if commitErr != nil && pendingPrepared { + _, clearErr := local.ClearArtifactResetRepublishPending( + context.WithoutCancel(resetCtx), pending, + ) + return errors.Join(commitErr, clearErr) + } + return commitErr + }, + ) + if resetErr != nil { + replacement := artifact.ArtifactStore(nil) + status := http.StatusInternalServerError + if !current.Closed() { + replacement = ownedStore + status = http.StatusServiceUnavailable + } + if resetCtx.Err() != nil { + status = http.StatusServiceUnavailable + } + if replacement == nil { + s.metadata = nil + } + finishErr := s.artifactOps.finishReset(replacement) + http.Error(w, errors.Join(resetErr, finishErr).Error(), status) + return + } + + freshStore := fresh.Content() + if err := s.artifactOps.setResetStore(freshStore); err != nil { + s.metadata = nil + finishErr := s.artifactOps.finishReset(nil) + http.Error(w, errors.Join(err, fresh.Close(), finishErr).Error(), http.StatusInternalServerError) + return + } + s.artifactRepository = fresh + s.replaceArtifactStoreCursorRegistry() + s.metadata = artifact.NewMetadataRecorder(local, artifact.MetadataRecorderOptions{ + Origin: s.localArtifactOrigin(), + Store: freshStore, + }) + s.artifactBaselineDone = false + result, resetErr = s.republishArtifactRepositoryReset( + resetCtx, s.cfg.DataDir, local, s.localArtifactOrigin(), fresh, result, + ) + resetErr = errors.Join(resetErr, resetCtx.Err()) + if resetErr != nil { + finishErr := s.artifactOps.finishReset(freshStore) + http.Error(w, errors.Join(resetErr, finishErr).Error(), http.StatusInternalServerError) + return + } + s.artifactBaselineDone = true + if err := s.artifactOps.finishReset(freshStore); err != nil { + s.metadata = nil + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(artifactResetResponse{ + RepositoryResetResult: result, + ManualCleanup: artifact.ArtifactResetManualCleanupWarning, + ForeignArtifacts: artifact.ArtifactResetForeignRelayWarning, + }); err != nil && r.Context().Err() == nil { + log.Printf("encoding artifact reset response: %v", err) + } +} + +func (s *Server) postArtifactMaintenanceHTTP(w http.ResponseWriter, r *http.Request) { + if !isLocalhostRequest(r) { + http.Error(w, "artifact maintenance is only available from localhost", http.StatusForbidden) + return + } + if _, err := s.writableArtifactImportDB(); err != nil { + writeArtifactHTTPError(w, err) + return + } + store, release, err := s.acquireArtifactStore() + if err != nil { + writeArtifactHTTPError(w, err) + return + } + defer release() + + r.Body = http.MaxBytesReader(w, r.Body, 64<<10) + decoder := json.NewDecoder(r.Body) + decoder.DisallowUnknownFields() + var input artifactMaintenanceRequest + if err := decoder.Decode(&input); err != nil { + http.Error(w, "invalid artifact maintenance request", http.StatusBadRequest) + return + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + http.Error(w, "invalid artifact maintenance request", http.StatusBadRequest) + return + } + grace, graceErr := parseArtifactMaintenanceDuration(input.Grace, input.GraceSeconds) + quarantineGrace, quarantineGraceErr := parseArtifactMaintenanceDuration( + input.QuarantineGrace, input.QuarantineGraceSeconds, + ) + if graceErr != nil || quarantineGraceErr != nil { + http.Error(w, "invalid artifact maintenance limits", http.StatusBadRequest) + return + } + maintenanceOpts := artifact.ArtifactMaintenanceOptions{ + TrashGrace: grace, + EmptyTrash: artifact.WorkBudget{ + MaxObjects: input.MaxObjects, Cursor: input.TrashCursor, + }, + GC: artifact.WorkBudget{ + MaxObjects: input.MaxObjects, MaxBytes: input.MaxBytes, + Cursor: input.GCCursor, + }, + Repack: artifact.WorkBudget{ + MaxObjects: input.MaxObjects, MaxBytes: input.MaxBytes, + Cursor: input.RepackCursor, + }, + } + if err := artifact.ValidateArtifactMaintenanceOptions(maintenanceOpts); err != nil { + http.Error(w, "invalid artifact maintenance limits", http.StatusBadRequest) + return + } + + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + logical, err := artifact.GarbageCollect(r.Context(), artifact.GCOptions{ + Store: store, + Grace: grace, + QuarantineGrace: quarantineGrace, + DryRun: input.DryRun, + }) + if err != nil { + writeArtifactHTTPError(w, artifactRouteError("artifact retention", err)) + return + } + response := artifactMaintenanceResponse{Logical: logical} + if s.artifactRepository != nil && !input.DryRun { + response.Physical.Supported = true + response.Physical.Result, err = s.artifactRepository.RunMaintenance( + r.Context(), maintenanceOpts) + if err != nil { + writeArtifactHTTPError(w, artifactRouteError("artifact physical maintenance", err)) + return + } + } + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(response); err != nil && r.Context().Err() == nil { + log.Printf("encoding artifact maintenance response: %v", err) + } +} + +func writeArtifactHTTPError(w http.ResponseWriter, err error) { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return + } + status := http.StatusInternalServerError + message := "internal error" + var apiErr *apiErrorResponse + if errors.As(err, &apiErr) { + status = apiErr.Status + message = apiErr.Message + } + http.Error(w, message, status) +} + +func artifactPostOutput(res artifact.PeerArtifactWrite) *jsonOutput[artifactPostResponse] { + return &jsonOutput[artifactPostResponse]{ + Body: artifactPostResponse{ + Origin: res.Origin, + Kind: res.Kind, + Name: res.Name, + Hash: res.Hash, + Size: res.Size, + Duplicate: res.Duplicate, + }, + } +} + +func (s *Server) humaFinalizeArtifacts( + ctx context.Context, + _ *emptyInput, +) (*jsonOutput[artifactFinalizeResponse], error) { + local, err := s.writableArtifactImportDB() + if err != nil { + return nil, err + } + store, release, err := s.acquireArtifactStore() + if err != nil { + return nil, err + } + defer release() + + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + res, err := s.importPeerArtifacts(ctx, local, store) + if err != nil { + return nil, err + } + s.artifactImportPending = res.Deferred > 0 + return &jsonOutput[artifactFinalizeResponse]{ + Body: artifactFinalizeResponse{ + ImportedSessions: res.Sessions, + ImportedMessages: res.Messages, + ImportedMetadata: res.Metadata, + Deferred: res.Deferred, + }, + }, nil +} + +func (s *Server) writableArtifactImportDB() (*db.DB, error) { + if err := s.requireLocalArtifactStore(); err != nil { + return nil, err + } + return s.db.(*db.DB), nil +} + +func (s *Server) requireLocalArtifactStore() error { + local, ok := s.db.(*db.DB) + if !ok { + return apiError(http.StatusNotImplemented, + "artifact routes are not available in remote mode") + } + if local.ReadOnly() { + return apiError(http.StatusNotImplemented, + "artifact routes are not available in read-only mode") + } + return nil +} + +func (s *Server) hasWritableArtifactStore() bool { + local, ok := s.db.(*db.DB) + return ok && !local.ReadOnly() && s.artifactStore != nil +} + +func (s *Server) importPeerArtifacts( + ctx context.Context, local *db.DB, store artifact.ArtifactStore, +) (artifact.ImportResult, error) { + localOrigin := s.cfg.ArtifactOriginID + if localOrigin == "" { + var err error + localOrigin, err = artifact.EnsureOrigin(local) + if err != nil { + return artifact.ImportResult{}, internalError("artifact import origin", err) + } + } + coordinator := artifact.NewStoreImportCoordinator(local, store, localOrigin) + res, err := coordinator.Finalize(ctx) + if err != nil { + return artifact.ImportResult{}, artifactRouteError("import peer artifacts", err) + } + if res.Changed() && s.broadcaster != nil { + // Imports add sessions and apply curation metadata, so live + // clients need the session-index refresh that only the + // "sessions" scope triggers; "messages" additionally + // invalidates hydrated session details and cached signal + // detail. Emit "sessions" last so a coalesced burst resolves + // to the index refresh. + s.broadcaster.Emit("messages") + s.broadcaster.Emit("sessions") + } + return res, nil +} + +func artifactRouteError(logPrefix string, err error) error { + switch { + case errors.Is(err, artifact.ErrArtifactInvalid): + return apiError(http.StatusBadRequest, err.Error()) + case errors.Is(err, artifact.ErrArtifactNotFound): + return apiError(http.StatusNotFound, "artifact not found") + case errors.Is(err, artifact.ErrArtifactConflict): + return apiError(http.StatusConflict, "artifact conflict") + default: + return internalError(logPrefix, err) + } +} diff --git a/internal/server/huma_routes_metadata_internal_test.go b/internal/server/huma_routes_metadata_internal_test.go index a658370a6..4515db2bb 100644 --- a/internal/server/huma_routes_metadata_internal_test.go +++ b/internal/server/huma_routes_metadata_internal_test.go @@ -1,12 +1,23 @@ package server import ( + "bytes" "context" + "crypto/sha256" + "encoding/hex" + "errors" + "slices" + "sync" + "sync/atomic" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" + "go.kenn.io/agentsview/internal/dbtest" "go.kenn.io/agentsview/internal/service" ) @@ -15,6 +26,25 @@ type statsSpyService struct { got service.StatsFilter } +func newArtifactHandlerTestServer( + t *testing.T, database *db.DB, cfg config.Config, +) *Server { + t.Helper() + if cfg.DataDir == "" { + cfg.DataDir = t.TempDir() + } + repository, err := artifact.OpenRepository(t.Context(), cfg.DataDir) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + store := repository.Content() + return &Server{ + db: database, + cfg: cfg, + artifactStore: store, + artifactOps: artifactOperationLifetime{store: store}, + } +} + func (s *statsSpyService) Stats( _ context.Context, f service.StatsFilter, ) (*service.SessionStats, error) { @@ -36,3 +66,993 @@ func TestHumaGetSessionStatsUsesServerGitHubToken(t *testing.T) { require.NoError(t, err) assert.Equal(t, "server-token", spy.got.GHToken) } + +func TestHumaBatchDeleteRestoresOnlyUnpublishedNewDeletions(t *testing.T) { + tests := []struct { + name string + failure error + wantDeleted []string + }{ + { + name: "ordinary append failure restores current and later", + failure: errors.New("artifact write failed"), + wantDeleted: []string{"already-trashed", "s1"}, + }, + { + name: "published failure keeps current and restores later", + failure: &artifact.MetadataPublishedError{ + Err: errors.New("replay bookkeeping failed"), + }, + wantDeleted: []string{"already-trashed", "s1", "s2"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + database := dbtest.OpenTestDB(t) + for _, id := range []string{"already-trashed", "s1", "s2", "s3"} { + dbtest.SeedSession(t, database, id, "alpha") + } + require.NoError(t, database.SoftDeleteSession("already-trashed")) + + var appended []string + srv := &Server{ + db: database, + metadataAppend: func( + _ context.Context, input artifact.MetadataEventInput, + ) error { + appended = append(appended, input.SessionID) + if input.SessionID == "s2" { + return tt.failure + } + return nil + }, + } + in := &batchDeleteInput{} + in.Body.SessionIDs = []string{"already-trashed", "s1", "s2", "s3"} + + _, err := srv.humaBatchDeleteSessions(context.Background(), in) + + require.Error(t, err) + assert.Equal(t, []string{"s1", "s2"}, appended, + "publication must stop at the first failed event") + for _, id := range []string{"already-trashed", "s1", "s2", "s3"} { + session, getErr := database.GetSessionFull(context.Background(), id) + require.NoError(t, getErr) + require.NotNil(t, session) + assert.Equal(t, containsString(tt.wantDeleted, id), session.DeletedAt != nil, + "unexpected trash state for %s", id) + } + }) + } +} + +func containsString(values []string, want string) bool { + return slices.Contains(values, want) +} + +type batchRestoreStore struct { + db.Store + restored []string + failID string + failErr error +} + +func (s *batchRestoreStore) RestoreSession(id string) (int64, error) { + s.restored = append(s.restored, id) + if id == s.failID { + return 0, s.failErr + } + return 1, nil +} + +func TestRestoreBatchDeletedSessionsJoinsFailuresAndContinues(t *testing.T) { + cause := errors.New("metadata append failed") + restoreErr := errors.New("restore failed") + store := &batchRestoreStore{failID: "s2", failErr: restoreErr} + srv := &Server{db: store} + + err := srv.restoreBatchDeletedSessions([]string{"s2", "s3"}, cause) + + assert.ErrorIs(t, err, cause) + assert.ErrorIs(t, err, restoreErr) + assert.Equal(t, []string{"s2", "s3"}, store.restored, + "a restoration failure must not prevent later sessions from being restored") +} + +func TestRestoreLifecycleWaitsForFailedRestoreCompensation(t *testing.T) { + const origin = "desk-a1b2c3" + database := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, database, "s1", "alpha") + require.NoError(t, database.SoftDeleteSession("s1")) + recordSoftDeleteReplayState(t, database, origin, "s1") + + firstAppendStarted := make(chan struct{}) + secondAppendStarted := make(chan struct{}) + secondLockAttempted := make(chan struct{}) + releaseFirstAppend := make(chan struct{}) + var appendCalls atomic.Int32 + var lockAttempts atomic.Int32 + srv := &Server{ + db: database, + cfg: config.Config{ArtifactOriginID: origin}, + beforeSessionLifecycleLock: func() { + if lockAttempts.Add(1) == 2 { + close(secondLockAttempted) + } + }, + metadataAppend: func( + _ context.Context, input artifact.MetadataEventInput, + ) error { + if input.Op != artifact.MetadataOpRestore { + return errors.New("unexpected metadata op") + } + switch appendCalls.Add(1) { + case 1: + close(firstAppendStarted) + <-releaseFirstAppend + return errors.New("artifact write failed") + case 2: + close(secondAppendStarted) + return nil + default: + return errors.New("unexpected extra metadata append") + } + }, + } + + firstResult := make(chan error, 1) + go func() { + _, err := srv.humaRestoreSession(context.Background(), &idPathInput{ID: "s1"}) + firstResult <- err + }() + <-firstAppendStarted + + secondResult := make(chan error, 1) + go func() { + _, err := srv.humaRestoreSession(context.Background(), &idPathInput{ID: "s1"}) + secondResult <- err + }() + <-secondLockAttempted + + interleaved := false + select { + case <-secondAppendStarted: + interleaved = true + case <-time.After(500 * time.Millisecond): + } + close(releaseFirstAppend) + firstErr := <-firstResult + secondErr := <-secondResult + + assert.False(t, interleaved, + "a second restore must not publish before the first restore compensates") + require.Error(t, firstErr) + require.NoError(t, secondErr) + assert.Equal(t, int32(2), appendCalls.Load()) + session, err := database.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, session) + assert.Nil(t, session.DeletedAt, + "the later successful restore must remain visible locally") +} + +func TestPermanentDeleteLifecycleExcludesConcurrentRestore(t *testing.T) { + const origin = "desk-a1b2c3" + database := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, database, "s1", "alpha") + require.NoError(t, database.SoftDeleteSession("s1")) + recordSoftDeleteReplayState(t, database, origin, "s1") + + firstAppendStarted := make(chan struct{}) + secondAppendStarted := make(chan struct{}) + secondLockAttempted := make(chan struct{}) + releaseFirstAppend := make(chan struct{}) + var appendCalls atomic.Int32 + var lockAttempts atomic.Int32 + srv := &Server{ + db: database, + cfg: config.Config{ArtifactOriginID: origin}, + beforeSessionLifecycleLock: func() { + if lockAttempts.Add(1) == 2 { + close(secondLockAttempted) + } + }, + metadataAppend: func( + _ context.Context, input artifact.MetadataEventInput, + ) error { + switch appendCalls.Add(1) { + case 1: + if input.Op != artifact.MetadataOpPurge { + return errors.New("unexpected first metadata op") + } + close(firstAppendStarted) + <-releaseFirstAppend + return nil + case 2: + if input.Op != artifact.MetadataOpRestore { + return errors.New("unexpected second metadata op") + } + close(secondAppendStarted) + return nil + default: + return errors.New("unexpected extra metadata append") + } + }, + } + + purgeResult := make(chan error, 1) + go func() { + _, err := srv.humaPermanentDeleteSession( + context.Background(), &idPathInput{ID: "s1"}, + ) + purgeResult <- err + }() + <-firstAppendStarted + + restoreResult := make(chan error, 1) + go func() { + _, err := srv.humaRestoreSession(context.Background(), &idPathInput{ID: "s1"}) + restoreResult <- err + }() + <-secondLockAttempted + + interleaved := false + select { + case <-secondAppendStarted: + interleaved = true + case <-time.After(500 * time.Millisecond): + } + close(releaseFirstAppend) + purgeErr := <-purgeResult + restoreErr := <-restoreResult + + assert.False(t, interleaved, + "restore must not publish while a purge has reserved the trashed session") + require.NoError(t, purgeErr) + require.Error(t, restoreErr) + assert.Equal(t, int32(1), appendCalls.Load()) + session, err := database.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + assert.Nil(t, session, + "a durable purge must delete the local session before restore can run") +} + +func TestPermanentDeleteLifecycleExcludesConcurrentPeerRestore(t *testing.T) { + const ( + localOrigin = "desk-a1b2c3" + peerOrigin = "peer-b2c3d4" + ) + database := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, database, "s1", "alpha") + require.NoError(t, database.SoftDeleteSession("s1")) + recordSoftDeleteReplayState(t, database, localOrigin, "s1") + + firstAppendStarted := make(chan struct{}) + secondLockAttempted := make(chan struct{}) + releaseFirstAppend := make(chan struct{}) + var appendCalls atomic.Int32 + var lockAttempts atomic.Int32 + srv := newArtifactHandlerTestServer(t, database, config.Config{ + ArtifactOriginID: localOrigin, + }) + srv.beforeSessionLifecycleLock = func() { + if lockAttempts.Add(1) == 2 { + close(secondLockAttempted) + } + } + srv.metadataAppend = func( + _ context.Context, input artifact.MetadataEventInput, + ) error { + if appendCalls.Add(1) != 1 || input.Op != artifact.MetadataOpPurge { + return errors.New("unexpected metadata append") + } + close(firstAppendStarted) + <-releaseFirstAppend + return nil + } + + purgeResult := make(chan error, 1) + go func() { + _, err := srv.humaPermanentDeleteSession( + context.Background(), &idPathInput{ID: "s1"}, + ) + purgeResult <- err + }() + <-firstAppendStarted + + body, name := peerRestoreArtifact( + peerOrigin, artifact.MetadataSessionGID(localOrigin, "s1"), + ) + peerResult := make(chan error, 1) + go func() { + _, err := srv.humaPostArtifact(context.Background(), &artifactPostInput{ + Origin: peerOrigin, + Kind: artifact.KindMeta, + Name: name, + Body: bytes.NewReader(body), + }) + peerResult <- err + }() + + peerCompletedBeforeRelease := false + var peerErr error + select { + case <-secondLockAttempted: + select { + case peerErr = <-peerResult: + peerCompletedBeforeRelease = true + case <-time.After(500 * time.Millisecond): + } + case peerErr = <-peerResult: + peerCompletedBeforeRelease = true + case <-time.After(2 * time.Second): + require.FailNow(t, "peer restore reached neither lifecycle lock nor completion") + } + close(releaseFirstAppend) + purgeErr := <-purgeResult + if !peerCompletedBeforeRelease { + peerErr = <-peerResult + } + + assert.False(t, peerCompletedBeforeRelease, + "peer restore import must wait for the durable purge to delete locally") + require.NoError(t, purgeErr) + require.NoError(t, peerErr) + assert.Equal(t, int32(1), appendCalls.Load()) + session, err := database.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + assert.Nil(t, session, + "peer restore must not interleave between purge publication and deletion") +} + +func TestArtifactContentPostDefersDatabaseImport(t *testing.T) { + ctx := context.Background() + const peerOrigin = "peer-b2c3d4" + peerDB := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, peerDB, "s1", "alpha") + peerStore := exportHumaArtifactFixture(t, ctx, peerDB, peerOrigin) + segmentRef := oneHumaArtifactRef(t, peerStore, peerOrigin, artifact.KindSegments) + segmentWire, segmentData := humaWireArtifact(t, peerStore, segmentRef) + require.NotEmpty(t, segmentData) + + local := dbtest.OpenTestDB(t) + server := newArtifactHandlerTestServer(t, local, config.Config{}) + _, err := server.humaPostArtifact(ctx, &artifactPostInput{ + Origin: peerOrigin, + Kind: artifact.KindSegments, + Name: segmentWire.Name, + Body: bytes.NewReader(segmentData), + }) + require.NoError(t, err) + + localOrigin, err := artifact.StoredOrigin(local) + require.NoError(t, err) + assert.Empty(t, localOrigin, + "content-only upload must not trigger a full import or enroll the receiver") +} + +func exportHumaArtifactFixture( + t *testing.T, ctx context.Context, database *db.DB, origin string, +) artifact.ArtifactStore { + t.Helper() + repository, err := artifact.OpenRepository(ctx, t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + _, err = artifact.ExportToStore(ctx, database, repository.Content(), artifact.ExportOptions{ + Origin: origin, + Full: true, + }) + require.NoError(t, err) + return repository.Content() +} + +func oneHumaArtifactRef( + t *testing.T, store artifact.ArtifactStore, origin string, kind artifact.Kind, +) artifact.Ref { + t.Helper() + entries := collectArtifactEntries(t, store, origin, kind, 10) + require.Len(t, entries, 1) + return entries[0].Ref +} + +func humaWireArtifact( + t *testing.T, store artifact.ArtifactStore, ref artifact.Ref, +) (artifact.WireRef, []byte) { + t.Helper() + _, reader, err := store.Open(t.Context(), ref) + require.NoError(t, err) + defer func() { require.NoError(t, reader.Close()) }() + wire, err := artifact.ToWireRef(ref) + require.NoError(t, err) + var body bytes.Buffer + require.NoError(t, artifact.EncodeWire(t.Context(), ref, reader, &body)) + require.NoError(t, reader.Verify()) + return wire, body.Bytes() +} + +func TestArtifactDependencyPostRetriesDeferredCheckpointImport(t *testing.T) { + ctx := context.Background() + const peerOrigin = "peer-b2c3d4" + peerDB := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, peerDB, "s1", "alpha") + peerStore := exportHumaArtifactFixture(t, ctx, peerDB, peerOrigin) + checkpoint := oneHumaArtifactRef(t, peerStore, peerOrigin, artifact.KindCheckpoints) + manifest := oneHumaArtifactRef(t, peerStore, peerOrigin, artifact.KindManifests) + segment := oneHumaArtifactRef(t, peerStore, peerOrigin, artifact.KindSegments) + + local := dbtest.OpenTestDB(t) + server := newArtifactHandlerTestServer(t, local, config.Config{}) + post := func(ref artifact.Ref) { + t.Helper() + wire, data := humaWireArtifact(t, peerStore, ref) + require.NotEmpty(t, data) + _, postErr := server.humaPostArtifact(ctx, &artifactPostInput{ + Origin: peerOrigin, Kind: string(ref.Kind), Name: wire.Name, Body: bytes.NewReader(data), + }) + require.NoError(t, postErr) + } + + post(checkpoint) + post(manifest) + got, err := local.GetSession(ctx, peerOrigin+"~s1") + require.NoError(t, err) + assert.Nil(t, got) + + post(segment) + got, err = local.GetSession(ctx, peerOrigin+"~s1") + require.NoError(t, err) + require.NotNil(t, got) + assert.Equal(t, "alpha", got.Project) +} + +func peerRestoreArtifact(origin, sessionGID string) ([]byte, string) { + const hlc = "2026-07-10T010203.000000001Z-00000000000000000000" + body := []byte(`{"hlc":"` + hlc + `","op":"restore","origin":"` + origin + + `","session_gid":"` + sessionGID + `","v":1}` + "\n") + sum := sha256.Sum256(body) + hash := hex.EncodeToString(sum[:]) + return body, hlc + "-" + hash + ".json" +} + +func recordSoftDeleteReplayState(t *testing.T, database *db.DB, origin, sessionID string) { + t.Helper() + _, err := database.RecordLocalMetadataProjection(context.Background(), db.MetadataProjection{ + EventOrigin: origin, + OrderKey: "0001-soft-delete", + HLC: "0001", + ArtifactHash: "soft-delete", + SessionGID: artifact.MetadataSessionGID(origin, sessionID), + LocalSessionID: sessionID, + Field: "deleted_at", + Op: artifact.MetadataOpSoftDelete, + Value: artifact.MetadataOpSoftDelete, + }) + require.NoError(t, err) +} + +type curationAppendGate struct { + firstStarted chan struct{} + secondStarted chan struct{} + releaseFirst chan struct{} + releaseOnce sync.Once + calls atomic.Int32 + inputs [2]artifact.MetadataEventInput +} + +func newCurationAppendGate() *curationAppendGate { + return &curationAppendGate{ + firstStarted: make(chan struct{}), + secondStarted: make(chan struct{}), + releaseFirst: make(chan struct{}), + } +} + +func (g *curationAppendGate) append( + _ context.Context, input artifact.MetadataEventInput, +) error { + call := g.calls.Add(1) + if call <= int32(len(g.inputs)) { + g.inputs[call-1] = input + } + switch call { + case 1: + close(g.firstStarted) + <-g.releaseFirst + return errors.New("artifact write failed") + case 2: + close(g.secondStarted) + return nil + default: + return errors.New("unexpected extra metadata append") + } +} + +func (g *curationAppendGate) release() { + g.releaseOnce.Do(func() { close(g.releaseFirst) }) +} + +func curationMetadataRecorder(t *testing.T, database *db.DB) *artifact.MetadataRecorder { + t.Helper() + repository, err := artifact.OpenRepository(t.Context(), t.TempDir()) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, repository.Close()) }) + return artifact.NewMetadataRecorder(database, artifact.MetadataRecorderOptions{ + Store: repository.Content(), + Origin: "desk-a1b2c3", + }) +} + +func waitForSecondCurationRequest( + t *testing.T, + secondLockAttempted <-chan struct{}, + secondAppendStarted <-chan struct{}, +) bool { + t.Helper() + select { + case <-secondLockAttempted: + select { + case <-secondAppendStarted: + return true + case <-time.After(500 * time.Millisecond): + return false + } + case <-secondAppendStarted: + return true + case <-time.After(2 * time.Second): + require.FailNow(t, "second curation request reached neither lifecycle lock nor metadata append") + return false + } +} + +func TestRenameLifecycleWaitsForFailedAppendCompensation(t *testing.T) { + database := dbtest.OpenTestDB(t) + original := "original" + dbtest.SeedSession(t, database, "s1", "alpha", func(session *db.Session) { + session.DisplayName = &original + }) + + gate := newCurationAppendGate() + t.Cleanup(gate.release) + secondLockAttempted := make(chan struct{}) + var lockAttempts atomic.Int32 + srv := &Server{ + db: database, + beforeSessionLifecycleLock: func() { + if lockAttempts.Add(1) == 2 { + close(secondLockAttempted) + } + }, + metadataAppend: gate.append, + } + + firstName := "first" + firstResult := make(chan error, 1) + go func() { + _, err := srv.humaRenameSession(context.Background(), &renameSessionInput{ + ID: "s1", + Body: renameRequest{DisplayName: &firstName}, + }) + firstResult <- err + }() + <-gate.firstStarted + + secondName := "second" + secondResult := make(chan error, 1) + go func() { + _, err := srv.humaRenameSession(context.Background(), &renameSessionInput{ + ID: "s1", + Body: renameRequest{DisplayName: &secondName}, + }) + secondResult <- err + }() + + interleaved := waitForSecondCurationRequest( + t, secondLockAttempted, gate.secondStarted, + ) + gate.release() + firstErr := <-firstResult + secondErr := <-secondResult + + assert.False(t, interleaved, + "a later rename must not publish before the failed rename compensates") + require.Error(t, firstErr) + require.NoError(t, secondErr) + assert.Equal(t, int32(2), gate.calls.Load()) + assert.Equal(t, artifact.MetadataOpRename, gate.inputs[0].Op) + assert.Equal(t, artifact.MetadataOpRename, gate.inputs[1].Op) + assert.JSONEq(t, `{"display_name":"second"}`, string(gate.inputs[1].Value)) + session, err := database.GetSession(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, session) + require.NotNil(t, session.DisplayName) + assert.Equal(t, "second", *session.DisplayName, + "SQLite must retain the value represented by the later durable rename") +} + +func TestStarLifecycleWaitsForFailedAppendCompensation(t *testing.T) { + database := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, database, "s1", "alpha") + + gate := newCurationAppendGate() + t.Cleanup(gate.release) + secondLockAttempted := make(chan struct{}) + var lockAttempts atomic.Int32 + srv := &Server{ + db: database, + metadata: curationMetadataRecorder(t, database), + beforeSessionLifecycleLock: func() { + if lockAttempts.Add(1) == 2 { + close(secondLockAttempted) + } + }, + metadataAppend: gate.append, + } + + firstResult := make(chan error, 1) + go func() { + _, err := srv.humaStarSession(context.Background(), &idPathInput{ID: "s1"}) + firstResult <- err + }() + <-gate.firstStarted + + secondResult := make(chan error, 1) + go func() { + _, err := srv.humaStarSession(context.Background(), &idPathInput{ID: "s1"}) + secondResult <- err + }() + + interleaved := waitForSecondCurationRequest( + t, secondLockAttempted, gate.secondStarted, + ) + gate.release() + firstErr := <-firstResult + secondErr := <-secondResult + + assert.False(t, interleaved, + "a later star must not publish before the failed star compensates") + require.Error(t, firstErr) + require.NoError(t, secondErr) + assert.Equal(t, int32(2), gate.calls.Load()) + assert.Equal(t, artifact.MetadataOpStar, gate.inputs[0].Op) + assert.Equal(t, artifact.MetadataOpStar, gate.inputs[1].Op) + starred, err := database.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"s1"}, starred, + "SQLite must retain the star represented by the later durable event") +} + +func seedCurationPinMessage(t *testing.T, database *db.DB) int64 { + t.Helper() + dbtest.SeedSession(t, database, "s1", "alpha") + message := dbtest.UserMsg("s1", 0, "investigate") + message.SourceUUID = "message-a1b2c3" + dbtest.SeedMessages(t, database, message) + messages, err := database.GetAllMessages(context.Background(), "s1") + require.NoError(t, err) + require.Len(t, messages, 1) + return messages[0].ID +} + +type metadataPinPointLookupStore struct { + db.Store + message *db.Message + fullReads int + pointReads int + lookupSessionID string + lookupMessageID int64 +} + +func (s *metadataPinPointLookupStore) GetAllMessages( + _ context.Context, _ string, +) ([]db.Message, error) { + s.fullReads++ + return nil, errors.New("full transcript read must not be used for pin metadata") +} + +func (s *metadataPinPointLookupStore) GetMessageForMetadataPin( + _ context.Context, sessionID string, messageID int64, +) (*db.Message, error) { + s.pointReads++ + s.lookupSessionID = sessionID + s.lookupMessageID = messageID + return s.message, nil +} + +func TestMetadataPinUsesPointMessageLookup(t *testing.T) { + store := &metadataPinPointLookupStore{message: &db.Message{ + ID: 42, + SessionID: "s1", + Ordinal: 7, + SourceUUID: "message-a1b2c3", + }} + srv := &Server{db: store} + note := "remember" + + pin, err := srv.metadataPinForMessage( + context.Background(), "s1", 42, ¬e, + ) + require.NoError(t, err) + require.NotNil(t, pin) + assert.Equal(t, "message-a1b2c3", pin.SourceUUID) + assert.Equal(t, 7, pin.Ordinal) + require.NotNil(t, pin.Note) + assert.Equal(t, "remember", *pin.Note) + assert.Equal(t, 1, store.pointReads) + assert.Zero(t, store.fullReads) + assert.Equal(t, "s1", store.lookupSessionID) + assert.Equal(t, int64(42), store.lookupMessageID) +} + +func TestPinLifecycleWaitsForFailedAppendCompensation(t *testing.T) { + database := dbtest.OpenTestDB(t) + messageID := seedCurationPinMessage(t, database) + + gate := newCurationAppendGate() + t.Cleanup(gate.release) + secondLockAttempted := make(chan struct{}) + var lockAttempts atomic.Int32 + srv := &Server{ + db: database, + metadata: curationMetadataRecorder(t, database), + beforeSessionLifecycleLock: func() { + if lockAttempts.Add(1) == 2 { + close(secondLockAttempted) + } + }, + metadataAppend: gate.append, + } + + firstNote := "first" + firstResult := make(chan error, 1) + go func() { + _, err := srv.humaPinMessage(context.Background(), &pinMessageInput{ + ID: "s1", + MessageID: messageID, + Body: pinRequest{Note: &firstNote}, + }) + firstResult <- err + }() + <-gate.firstStarted + + secondNote := "second" + secondResult := make(chan error, 1) + go func() { + _, err := srv.humaPinMessage(context.Background(), &pinMessageInput{ + ID: "s1", + MessageID: messageID, + Body: pinRequest{Note: &secondNote}, + }) + secondResult <- err + }() + + interleaved := waitForSecondCurationRequest( + t, secondLockAttempted, gate.secondStarted, + ) + gate.release() + firstErr := <-firstResult + secondErr := <-secondResult + + assert.False(t, interleaved, + "a later pin must not publish before the failed pin compensates") + require.Error(t, firstErr) + require.NoError(t, secondErr) + assert.Equal(t, int32(2), gate.calls.Load()) + assert.Equal(t, artifact.MetadataOpPin, gate.inputs[0].Op) + assert.Equal(t, artifact.MetadataOpPin, gate.inputs[1].Op) + require.NotNil(t, gate.inputs[1].Pin) + require.NotNil(t, gate.inputs[1].Pin.Note) + assert.Equal(t, "second", *gate.inputs[1].Pin.Note) + pins, err := database.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + require.Len(t, pins, 1, + "SQLite must retain the pin represented by the later durable event") + require.NotNil(t, pins[0].Note) + assert.Equal(t, "second", *pins[0].Note) +} + +func assertCurationMutationWaitsForLifecycle( + t *testing.T, + srv *Server, + wantOp string, + mutate func() error, + assertBefore func(), + assertAfter func(), +) { + t.Helper() + lockAttempted := make(chan struct{}) + appended := make(chan artifact.MetadataEventInput, 2) + var lockAttempts atomic.Int32 + srv.beforeSessionLifecycleLock = func() { + if lockAttempts.Add(1) == 1 { + close(lockAttempted) + } + } + srv.metadataAppend = func( + _ context.Context, input artifact.MetadataEventInput, + ) error { + appended <- input + return nil + } + + srv.sessionLifecycleMu.Lock() + locked := true + defer func() { + if locked { + srv.sessionLifecycleMu.Unlock() + } + }() + + result := make(chan error, 1) + go func() { result <- mutate() }() + + completedBeforeLock := false + var mutationErr error + select { + case <-lockAttempted: + case mutationErr = <-result: + completedBeforeLock = true + case <-time.After(2 * time.Second): + require.FailNow(t, "curation mutation reached neither lifecycle lock nor completion") + } + assert.False(t, completedBeforeLock, + "curation mutation must wait for the shared lifecycle boundary") + assertBefore() + + srv.sessionLifecycleMu.Unlock() + locked = false + if !completedBeforeLock { + mutationErr = <-result + } + require.NoError(t, mutationErr) + assertAfter() + assert.Equal(t, int32(1), lockAttempts.Load()) + select { + case input := <-appended: + assert.Equal(t, wantOp, input.Op) + case <-time.After(2 * time.Second): + require.FailNow(t, "curation mutation did not append metadata") + } + select { + case input := <-appended: + assert.Fail(t, "curation mutation appended extra metadata", "op: %s", input.Op) + default: + } +} + +func TestUnstarMutationWaitsForSessionLifecycle(t *testing.T) { + database := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, database, "s1", "alpha") + starred, err := database.StarSession("s1") + require.NoError(t, err) + require.True(t, starred) + srv := &Server{db: database} + + assertStarred := func(want []string) func() { + return func() { + ids, listErr := database.ListStarredSessionIDs(context.Background()) + require.NoError(t, listErr) + assert.ElementsMatch(t, want, ids) + } + } + assertCurationMutationWaitsForLifecycle( + t, + srv, + artifact.MetadataOpUnstar, + func() error { + _, handlerErr := srv.humaUnstarSession( + context.Background(), &idPathInput{ID: "s1"}, + ) + return handlerErr + }, + assertStarred([]string{"s1"}), + assertStarred([]string{}), + ) +} + +func TestBulkStarMutationWaitsForSessionLifecycle(t *testing.T) { + database := dbtest.OpenTestDB(t) + dbtest.SeedSession(t, database, "s1", "alpha") + srv := &Server{db: database} + in := &bulkStarInput{} + in.Body.SessionIDs = []string{"s1"} + + assertStarred := func(want []string) func() { + return func() { + ids, err := database.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.ElementsMatch(t, want, ids) + } + } + assertCurationMutationWaitsForLifecycle( + t, + srv, + artifact.MetadataOpStar, + func() error { + _, err := srv.humaBulkStar(context.Background(), in) + return err + }, + assertStarred([]string{}), + assertStarred([]string{"s1"}), + ) +} + +func TestBulkStarEmptyInputDoesNotWaitForSessionLifecycle(t *testing.T) { + lockAttempted := make(chan struct{}) + var appendCalls atomic.Int32 + srv := &Server{ + beforeSessionLifecycleLock: func() { close(lockAttempted) }, + metadataAppend: func( + _ context.Context, _ artifact.MetadataEventInput, + ) error { + appendCalls.Add(1) + return nil + }, + } + srv.sessionLifecycleMu.Lock() + locked := true + defer func() { + if locked { + srv.sessionLifecycleMu.Unlock() + } + }() + + result := make(chan error, 1) + go func() { + _, err := srv.humaBulkStar(context.Background(), &bulkStarInput{}) + result <- err + }() + + select { + case err := <-result: + require.NoError(t, err) + case <-lockAttempted: + srv.sessionLifecycleMu.Unlock() + locked = false + <-result + require.FailNow(t, "empty bulk star waited for the lifecycle boundary") + case <-time.After(2 * time.Second): + require.FailNow(t, "empty bulk star did not complete") + } + assert.Equal(t, int32(0), appendCalls.Load()) +} + +func TestUnpinMutationWaitsForSessionLifecycle(t *testing.T) { + database := dbtest.OpenTestDB(t) + messageID := seedCurationPinMessage(t, database) + note := "keep" + _, err := database.PinMessage("s1", messageID, ¬e) + require.NoError(t, err) + srv := &Server{ + db: database, + metadata: curationMetadataRecorder(t, database), + } + + assertPinned := func(want bool) func() { + return func() { + pins, listErr := database.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, listErr) + if !want { + assert.Empty(t, pins) + return + } + require.Len(t, pins, 1) + require.NotNil(t, pins[0].Note) + assert.Equal(t, "keep", *pins[0].Note) + } + } + assertCurationMutationWaitsForLifecycle( + t, + srv, + artifact.MetadataOpUnpin, + func() error { + _, handlerErr := srv.humaUnpinMessage(context.Background(), &messagePathInput{ + ID: "s1", + MessageID: messageID, + }) + return handlerErr + }, + assertPinned(true), + assertPinned(false), + ) +} diff --git a/internal/server/huma_routes_pins.go b/internal/server/huma_routes_pins.go index 13eb81ce2..f0cdfec40 100644 --- a/internal/server/huma_routes_pins.go +++ b/internal/server/huma_routes_pins.go @@ -2,8 +2,11 @@ package server import ( "context" + "errors" + "fmt" "net/http" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/db" ) @@ -63,9 +66,25 @@ func (s *Server) humaListSessionPins( } func (s *Server) humaPinMessage( - _ context.Context, + ctx context.Context, in *pinMessageInput, ) (*createdOutput[pinMessageResponse], error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + + var prior *db.PinnedMessage + var pin *artifact.MetadataPin + if s.metadata != nil { + var err error + prior, err = s.findPinnedMessage(ctx, in.ID, in.MessageID) + if err != nil { + return nil, internalError("pin message prior state", err) + } + pin, err = s.metadataPinForMessage(ctx, in.ID, in.MessageID, in.Body.Note) + if err != nil { + return nil, internalError("pin message metadata lookup", err) + } + } id, err := s.db.PinMessage(in.ID, in.MessageID, in.Body.Note) if err != nil { if handled := handleHumaReadOnly(err); handled != nil { @@ -77,6 +96,20 @@ func (s *Server) humaPinMessage( return nil, apiError(http.StatusBadRequest, "message does not belong to this session") } + if pin != nil { + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpPin, + Pin: pin, + }); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("pin message metadata event", err) + } + return nil, internalError("pin message metadata event", + s.restorePinState(in.ID, in.MessageID, prior, err)) + } + } return &createdOutput[pinMessageResponse]{ Status: http.StatusCreated, Body: pinMessageResponse{ID: id}, @@ -84,14 +117,82 @@ func (s *Server) humaPinMessage( } func (s *Server) humaUnpinMessage( - _ context.Context, + ctx context.Context, in *messagePathInput, ) (*noContentOutput, error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + + var prior *db.PinnedMessage + var pin *artifact.MetadataPin + if s.metadata != nil { + var err error + prior, err = s.findPinnedMessage(ctx, in.ID, in.MessageID) + if err != nil { + return nil, internalError("unpin message prior state", err) + } + pin, err = s.metadataPinForMessage(ctx, in.ID, in.MessageID, nil) + if err != nil { + return nil, internalError("unpin message metadata lookup", err) + } + } if err := s.db.UnpinMessage(in.ID, in.MessageID); err != nil { if handled := handleHumaReadOnly(err); handled != nil { return nil, handled } return nil, internalError("unpin message", err) } + if pin != nil { + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpUnpin, + Pin: pin, + }); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("unpin message metadata event", err) + } + return nil, internalError("unpin message metadata event", + s.restorePinState(in.ID, in.MessageID, prior, err)) + } + } return &noContentOutput{Status: http.StatusNoContent}, nil } + +// findPinnedMessage returns the current pinned_messages row for the +// message, or nil when the message is not pinned. +func (s *Server) findPinnedMessage( + ctx context.Context, sessionID string, messageID int64, +) (*db.PinnedMessage, error) { + pins, err := s.db.ListPinnedMessages(ctx, sessionID, "") + if err != nil { + return nil, err + } + for i := range pins { + if pins[i].MessageID == messageID { + return &pins[i], nil + } + } + return nil, nil +} + +// restorePinState puts the pinned_messages row for the message back to +// prior after a pre-publish metadata failure, so local pin state never +// diverges from the durable ledger. It returns baseErr joined with any +// restore failure. +func (s *Server) restorePinState( + sessionID string, messageID int64, prior *db.PinnedMessage, baseErr error, +) error { + if prior != nil { + if _, err := s.db.PinMessage(sessionID, messageID, prior.Note); err != nil { + return errors.Join(baseErr, + fmt.Errorf("restore pin after metadata failure: %w", err)) + } + return baseErr + } + if err := s.db.UnpinMessage(sessionID, messageID); err != nil { + return errors.Join(baseErr, + fmt.Errorf("remove pin after metadata failure: %w", err)) + } + return baseErr +} diff --git a/internal/server/huma_routes_sessions.go b/internal/server/huma_routes_sessions.go index 02e3816f7..cc5587ebe 100644 --- a/internal/server/huma_routes_sessions.go +++ b/internal/server/huma_routes_sessions.go @@ -15,6 +15,7 @@ import ( "time" "github.com/danielgtaylor/huma/v2" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/export" "go.kenn.io/agentsview/internal/parser" @@ -35,6 +36,7 @@ func (s *Server) registerSessionRoutes() { get(s, group, "/sessions/{id}/activity", "Get session activity", s.humaGetSessionActivity) get(s, group, "/sessions/{id}/timing", "Get session timing", s.humaSessionTiming) get(s, group, "/sessions/{id}/usage", "Get session usage", s.humaSessionUsage) + get(s, group, "/sessions/{id}/metadata-conflicts", "List session metadata conflicts", s.humaListMetadataConflicts) stream(s, group, http.MethodGet, "/sessions/{id}/watch", "Watch session events", s.humaWatchSession) stream(s, group, http.MethodGet, "/events", "Watch server events", s.humaEvents) raw(s, group, http.MethodGet, "/sessions/{id}/export", "Export session as HTML", s.humaExportSession) @@ -575,6 +577,43 @@ type emptyTrashResponse struct { Deleted int `json:"deleted"` } +type metadataConflictsResponse struct { + Conflicts []db.MetadataConflict `json:"conflicts"` +} + +func (s *Server) humaListMetadataConflicts( + ctx context.Context, + in *idPathInput, +) (*jsonOutput[metadataConflictsResponse], error) { + session, err := s.db.GetSessionFull(ctx, in.ID) + if err != nil { + return nil, internalError("metadata conflict session lookup", err) + } + if session == nil { + return nil, apiError(http.StatusNotFound, "session not found") + } + gids := []string{in.ID} + if localDB, ok := s.db.(*db.DB); ok && !strings.Contains(in.ID, "~") { + origin, err := artifact.StoredOrigin(localDB) + if err != nil { + return nil, internalError("read artifact origin", err) + } + if origin != "" { + gids = append(gids, artifact.MetadataSessionGID(origin, in.ID)) + } + } + conflicts, err := s.db.ListMetadataConflicts(ctx, gids) + if err != nil { + return nil, internalError("list metadata conflicts", err) + } + if conflicts == nil { + conflicts = []db.MetadataConflict{} + } + return &jsonOutput[metadataConflictsResponse]{ + Body: metadataConflictsResponse{Conflicts: conflicts}, + }, nil +} + func (s *Server) humaGetSessionDir( ctx context.Context, in *idPathInput, @@ -685,6 +724,9 @@ func (s *Server) humaRenameSession( ctx context.Context, in *renameSessionInput, ) (*jsonOutput[*db.Session], error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + session, err := s.db.GetSession(ctx, in.ID) if err != nil { return nil, internalError("rename session lookup", err) @@ -702,6 +744,27 @@ func (s *Server) humaRenameSession( } return nil, internalError("rename session", err) } + value, err := renameMetadataValue(displayName) + if err != nil { + return nil, internalError("rename session metadata value", err) + } + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpRename, + Value: value, + }); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("rename session metadata event", err) + } + if restoreErr := s.db.RenameSession(in.ID, session.DisplayName); restoreErr != nil { + return nil, internalError( + "rename session metadata event", + errors.Join(err, fmt.Errorf("restore display name after metadata failure: %w", restoreErr)), + ) + } + return nil, internalError("rename session metadata event", err) + } updated, err := s.db.GetSession(ctx, in.ID) if err != nil { @@ -717,6 +780,9 @@ func (s *Server) humaDeleteSession( ctx context.Context, in *idPathInput, ) (*noContentOutput, error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + session, err := s.db.GetSessionFull(ctx, in.ID) if err != nil { return nil, internalError("delete session lookup", err) @@ -730,6 +796,32 @@ func (s *Server) humaDeleteSession( } return nil, internalError("soft delete session", err) } + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpSoftDelete, + }); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("soft delete session metadata event", err) + } + // Only undo a trashing this request performed: SoftDeleteSession + // is a no-op on an already-trashed session, and restoring one of + // those would revert an earlier, ledger-backed deletion. + if session.DeletedAt == nil { + if n, restoreErr := s.db.RestoreSession(in.ID); restoreErr != nil { + return nil, internalError( + "soft delete session metadata event", + errors.Join(err, fmt.Errorf("restore session after metadata failure: %w", restoreErr)), + ) + } else if n == 0 { + return nil, internalError( + "soft delete session metadata event", + errors.Join(err, fmt.Errorf("restore session after metadata failure: session %q not in trash", in.ID)), + ) + } + } + return nil, internalError("soft delete session metadata event", err) + } s.notifySessionMutation() return &noContentOutput{Status: http.StatusNoContent}, nil } @@ -748,27 +840,118 @@ func (s *Server) notifySessionMutation() { } } +type trashedSessionIDStore interface { + TrashedSessionIDs(ids []string) ([]string, error) +} + +type excludedSessionStore interface { + IsSessionExcluded(id string) bool +} + func (s *Server) humaBatchDeleteSessions( - _ context.Context, + ctx context.Context, in *batchDeleteInput, ) (*noContentOutput, error) { if len(in.Body.SessionIDs) == 0 { return &noContentOutput{Status: http.StatusNoContent}, nil } - if _, err := s.db.SoftDeleteSessions(in.Body.SessionIDs); err != nil { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + ctx, releaseArtifactStore, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return nil, internalError("batch delete artifact store", err) + } + defer releaseArtifactStore() + + deletedIDs, err := s.db.SoftDeleteSessionsReturningIDs(in.Body.SessionIDs) + if err != nil { if handled := handleHumaReadOnly(err); handled != nil { return nil, handled } return nil, internalError("batch delete sessions", err) } + newlyDeleted := make(map[string]struct{}, len(deletedIDs)) + for _, id := range deletedIDs { + newlyDeleted[id] = struct{}{} + } + for i, id := range deletedIDs { + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: id, + Op: artifact.MetadataOpSoftDelete, + }); err != nil { + rollbackStart := i + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + rollbackStart++ + } + rollback := make([]string, 0, len(deletedIDs)-rollbackStart) + for _, rollbackID := range deletedIDs[rollbackStart:] { + if _, ok := newlyDeleted[rollbackID]; ok { + rollback = append(rollback, rollbackID) + } + } + return nil, internalError( + "batch delete session metadata event", + s.restoreBatchDeletedSessions(rollback, err), + ) + } + } + store, ok := s.db.(trashedSessionIDStore) + if !ok { + return &noContentOutput{Status: http.StatusNoContent}, nil + } + trashedIDs, err := store.TrashedSessionIDs(in.Body.SessionIDs) + if err != nil { + return nil, internalError("batch delete retry lookup", err) + } + for _, id := range trashedIDs { + if _, ok := newlyDeleted[id]; ok { + continue + } + if err := s.ensureLocalMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: id, + Op: artifact.MetadataOpSoftDelete, + }, "deleted_at", artifact.MetadataOpSoftDelete); err != nil { + return nil, internalError("batch delete metadata repair", err) + } + } s.notifySessionMutation() return &noContentOutput{Status: http.StatusNoContent}, nil } +func (s *Server) restoreBatchDeletedSessions(ids []string, cause error) error { + errs := []error{cause} + for _, id := range ids { + n, err := s.db.RestoreSession(id) + if err != nil { + errs = append(errs, + fmt.Errorf("restore session %s after metadata failure: %w", id, err)) + continue + } + if n == 0 { + errs = append(errs, + fmt.Errorf("restore session %s after metadata failure: session not in trash", id)) + } + } + return errors.Join(errs...) +} + func (s *Server) humaRestoreSession( - _ context.Context, + ctx context.Context, in *idPathInput, ) (*noContentOutput, error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + ctx, releaseArtifactStore, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return nil, internalError("restore session artifact store", err) + } + defer releaseArtifactStore() + + metadataInput := artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpRestore, + } n, err := s.db.RestoreSession(in.ID) if err != nil { if handled := handleHumaReadOnly(err); handled != nil { @@ -777,26 +960,122 @@ func (s *Server) humaRestoreSession( return nil, internalError("restore session", err) } if n == 0 { + session, err := s.db.GetSessionFull(ctx, in.ID) + if err != nil { + return nil, internalError("restore session retry lookup", err) + } + if session != nil && session.DeletedAt == nil { + op, ok, err := s.metadataReplayStateOp(ctx, in.ID, "deleted_at") + if err != nil { + return nil, internalError("restore session metadata retry lookup", err) + } + if !ok || op != artifact.MetadataOpSoftDelete { + return nil, apiError(http.StatusNotFound, "session not found or not in trash") + } + _, err = s.repairLocalMetadataEvent(ctx, metadataInput) + if err != nil { + return nil, internalError("restore session metadata repair", err) + } + op, ok, err = s.metadataReplayStateOp(ctx, in.ID, "deleted_at") + if err != nil { + return nil, internalError("restore session metadata retry lookup", err) + } + if !ok || op != artifact.MetadataOpRestore { + if err := s.appendMetadataEvent(ctx, metadataInput); err != nil { + return nil, internalError("restore session metadata event", err) + } + } + return &noContentOutput{Status: http.StatusNoContent}, nil + } return nil, apiError(http.StatusNotFound, "session not found or not in trash") } + if err := s.appendMetadataEvent(ctx, metadataInput); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("restore session metadata event", err) + } + deletedIDs, rollbackErr := s.db.SoftDeleteSessionsReturningIDs([]string{in.ID}) + if rollbackErr != nil { + return nil, internalError( + "restore session metadata event", + errors.Join(err, fmt.Errorf("return session to trash after metadata failure: %w", rollbackErr)), + ) + } + if len(deletedIDs) != 1 || deletedIDs[0] != in.ID { + return nil, internalError( + "restore session metadata event", + errors.Join(err, fmt.Errorf("return session %q to trash after metadata failure: session no longer visible", in.ID)), + ) + } + return nil, internalError("restore session metadata event", err) + } s.notifySessionMutation() return &noContentOutput{Status: http.StatusNoContent}, nil } func (s *Server) humaPermanentDeleteSession( - _ context.Context, + ctx context.Context, in *idPathInput, ) (*noContentOutput, error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + ctx, releaseArtifactStore, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return nil, internalError("permanent delete artifact store", err) + } + defer releaseArtifactStore() + + metadataInput := artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpPurge, + } + session, err := s.db.GetSessionFull(ctx, in.ID) + if err != nil { + return nil, internalError("permanent delete session lookup", err) + } + if session == nil { + if store, ok := s.db.(excludedSessionStore); ok && store.IsSessionExcluded(in.ID) { + repaired, err := s.repairLocalMetadataEvent(ctx, metadataInput) + if err != nil { + return nil, internalError("permanent delete session metadata repair", err) + } + if repaired == 0 { + if err := s.appendMetadataEvent(ctx, metadataInput); err != nil { + return nil, internalError("permanent delete session metadata event", err) + } + } + return &noContentOutput{Status: http.StatusNoContent}, nil + } + return nil, apiError(http.StatusConflict, "session not found or not in trash") + } + if session.DeletedAt == nil { + return nil, apiError(http.StatusConflict, "session not found or not in trash") + } + metadataErr := s.ensureLocalMetadataEvent( + ctx, metadataInput, "purge", artifact.MetadataOpPurge, + ) + if metadataErr != nil { + var publishedErr *artifact.MetadataPublishedError + if !errors.As(metadataErr, &publishedErr) { + return nil, internalError("permanent delete session metadata event", metadataErr) + } + } n, err := s.db.DeleteSessionIfTrashed(in.ID) if err != nil { if handled := handleHumaReadOnly(err); handled != nil { return nil, handled } - return nil, internalError("permanent delete session", err) + return nil, internalError( + "permanent delete session", + errors.Join(metadataErr, err), + ) } if n == 0 { return nil, apiError(http.StatusConflict, "session not found or not in trash") } + if metadataErr != nil { + return nil, internalError("permanent delete session metadata event", metadataErr) + } s.notifySessionMutation() return &noContentOutput{Status: http.StatusNoContent}, nil } @@ -816,6 +1095,9 @@ func (s *Server) humaEmptyTrash( _ context.Context, _ *emptyInput, ) (*jsonOutput[emptyTrashResponse], error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + count, err := s.db.EmptyTrash() if err != nil { if handled := handleHumaReadOnly(err); handled != nil { diff --git a/internal/server/huma_routes_starred.go b/internal/server/huma_routes_starred.go index 1235123b4..b2861820c 100644 --- a/internal/server/huma_routes_starred.go +++ b/internal/server/huma_routes_starred.go @@ -2,7 +2,12 @@ package server import ( "context" + "errors" + "fmt" "net/http" + "slices" + + "go.kenn.io/agentsview/internal/artifact" ) func (s *Server) registerStarredRoutes() { @@ -39,9 +44,23 @@ func (s *Server) humaListStarred( } func (s *Server) humaStarSession( - _ context.Context, + ctx context.Context, in *idPathInput, ) (*noContentOutput, error) { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + + // Prior state decides rollback: StarSession reports success for an + // already-starred session too, and that star must survive a failed + // metadata append. + wasStarred := false + if s.metadata != nil { + var err error + wasStarred, err = s.sessionStarred(ctx, in.ID) + if err != nil { + return nil, internalError("star session prior state", err) + } + } ok, err := s.db.StarSession(in.ID) if err != nil { if handled := handleHumaReadOnly(err); handled != nil { @@ -52,34 +71,172 @@ func (s *Server) humaStarSession( if !ok { return nil, apiError(http.StatusNotFound, "session not found") } + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpStar, + }); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("star session metadata event", err) + } + if !wasStarred { + if _, removeErr := s.db.UnstarSession(in.ID); removeErr != nil { + return nil, internalError( + "star session metadata event", + errors.Join(err, fmt.Errorf("remove star after metadata failure: %w", removeErr)), + ) + } + } + return nil, internalError("star session metadata event", err) + } return &noContentOutput{Status: http.StatusNoContent}, nil } +// sessionStarred reports whether the session is currently starred. +func (s *Server) sessionStarred(ctx context.Context, id string) (bool, error) { + ids, err := s.db.ListStarredSessionIDs(ctx) + if err != nil { + return false, err + } + if slices.Contains(ids, id) { + return true, nil + } + return false, nil +} + func (s *Server) humaUnstarSession( - _ context.Context, + ctx context.Context, in *idPathInput, ) (*noContentOutput, error) { - if err := s.db.UnstarSession(in.ID); err != nil { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + + removed, err := s.db.UnstarSession(in.ID) + if err != nil { if handled := handleHumaReadOnly(err); handled != nil { return nil, handled } return nil, internalError("unstar session", err) } + if !removed { + if _, err := s.repairLocalMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpUnstar, + }); err != nil { + return nil, internalError("unstar session metadata repair", err) + } + return &noContentOutput{Status: http.StatusNoContent}, nil + } + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: in.ID, + Op: artifact.MetadataOpUnstar, + }); err != nil { + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + return nil, internalError("unstar session metadata event", err) + } + if restored, restoreErr := s.db.StarSession(in.ID); restoreErr != nil { + return nil, internalError( + "unstar session metadata event", + errors.Join(err, fmt.Errorf("restore star after metadata failure: %w", restoreErr)), + ) + } else if !restored { + return nil, internalError( + "unstar session metadata event", + errors.Join(err, fmt.Errorf("restore star after metadata failure: session %q not found", in.ID)), + ) + } + return nil, internalError("unstar session metadata event", err) + } return &noContentOutput{Status: http.StatusNoContent}, nil } func (s *Server) humaBulkStar( - _ context.Context, + ctx context.Context, in *bulkStarInput, ) (*noContentOutput, error) { if len(in.Body.SessionIDs) == 0 { return &noContentOutput{Status: http.StatusNoContent}, nil } - if err := s.db.BulkStarSessions(in.Body.SessionIDs); err != nil { + s.lockSessionLifecycle() + defer s.sessionLifecycleMu.Unlock() + ctx, releaseArtifactStore, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return nil, internalError("bulk star artifact store", err) + } + defer releaseArtifactStore() + + starred, err := s.db.BulkStarSessions(in.Body.SessionIDs) + if err != nil { if handled := handleHumaReadOnly(err); handled != nil { return nil, handled } return nil, internalError("bulk star", err) } + newlyStarred := make(map[string]struct{}, len(starred)) + // Emit one star event per session actually starred so localStorage star + // migration converges through artifact sync, matching single-session star. + for i, id := range starred { + newlyStarred[id] = struct{}{} + if err := s.appendMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: id, + Op: artifact.MetadataOpStar, + }); err != nil { + // Stars whose events are already in the ledger stay; the + // rest were created by this request without a ledger event + // to sync them, so they are removed. A published error + // means the failed event itself is durably recorded, so + // its star stays too. + rollback := starred[i:] + var publishedErr *artifact.MetadataPublishedError + if errors.As(err, &publishedErr) { + rollback = starred[i+1:] + } + return nil, internalError("bulk star metadata event", + s.rollbackBulkStar(rollback, err)) + } + } + starredIDs, err := s.db.ListStarredSessionIDs(ctx) + if err != nil { + return nil, internalError("bulk star metadata repair", err) + } + starredNow := make(map[string]struct{}, len(starredIDs)) + for _, id := range starredIDs { + starredNow[id] = struct{}{} + } + seenRetry := map[string]struct{}{} + for _, id := range in.Body.SessionIDs { + if _, ok := newlyStarred[id]; ok { + continue + } + if _, ok := starredNow[id]; !ok { + continue + } + if _, ok := seenRetry[id]; ok { + continue + } + seenRetry[id] = struct{}{} + if err := s.ensureLocalMetadataEvent(ctx, artifact.MetadataEventInput{ + SessionID: id, + Op: artifact.MetadataOpStar, + }, "starred", artifact.MetadataOpStar); err != nil { + return nil, internalError("bulk star metadata repair", err) + } + } return &noContentOutput{Status: http.StatusNoContent}, nil } + +// rollbackBulkStar removes stars this request created but never +// recorded in the ledger, so local state does not run ahead of the +// artifact log. The returned error joins the append failure with any +// rollback failures. +func (s *Server) rollbackBulkStar(ids []string, cause error) error { + errs := []error{cause} + for _, id := range ids { + if _, err := s.db.UnstarSession(id); err != nil { + errs = append(errs, + fmt.Errorf("remove star %s after metadata failure: %w", id, err)) + } + } + return errors.Join(errs...) +} diff --git a/internal/server/metadata_events.go b/internal/server/metadata_events.go new file mode 100644 index 000000000..63ceba199 --- /dev/null +++ b/internal/server/metadata_events.go @@ -0,0 +1,155 @@ +package server + +import ( + "context" + "encoding/json" + "fmt" + + "go.kenn.io/agentsview/internal/artifact" +) + +type metadataArtifactLeaseKey struct{} + +type metadataArtifactLease struct { + server *Server +} + +func (s *Server) acquireMetadataArtifactLease( + ctx context.Context, +) (context.Context, func(), error) { + if s.metadata == nil { + return ctx, func() {}, nil + } + if lease, ok := ctx.Value(metadataArtifactLeaseKey{}).(*metadataArtifactLease); ok && + lease.server == s { + return ctx, func() {}, nil + } + _, release, err := s.artifactOps.acquire() + if err != nil { + return ctx, nil, err + } + return context.WithValue(ctx, metadataArtifactLeaseKey{}, &metadataArtifactLease{ + server: s, + }), release, nil +} + +func (s *Server) appendMetadataEvent( + ctx context.Context, + input artifact.MetadataEventInput, +) error { + if artifact.MetadataEventsSuppressed(ctx) { + return nil + } + if s.metadataAppend != nil { + return s.metadataAppend(ctx, input) + } + if s.metadata == nil { + return nil + } + ctx, release, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return fmt.Errorf("acquiring artifact store for metadata append: %w", err) + } + defer release() + _, err = s.metadata.Append(ctx, input) + return err +} + +func (s *Server) repairLocalMetadataEvent( + ctx context.Context, + input artifact.MetadataEventInput, +) (int, error) { + if s.metadata == nil { + return 0, nil + } + ctx, release, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return 0, fmt.Errorf("acquiring artifact store for metadata repair: %w", err) + } + defer release() + return s.metadata.RepairLocalSessionMetadata(ctx, input.SessionID, input.Op) +} + +func (s *Server) ensureLocalMetadataEvent( + ctx context.Context, + input artifact.MetadataEventInput, + field string, + wantOp string, +) error { + if s.metadata == nil { + if s.metadataAppend != nil { + return s.appendMetadataEvent(ctx, input) + } + return nil + } + ctx, release, err := s.acquireMetadataArtifactLease(ctx) + if err != nil { + return fmt.Errorf("acquiring artifact store for metadata ensure: %w", err) + } + defer release() + if _, err := s.repairLocalMetadataEvent(ctx, input); err != nil { + return err + } + op, ok, err := s.metadataReplayStateOp(ctx, input.SessionID, field) + if err != nil { + return err + } + if ok && op == wantOp { + return nil + } + return s.appendMetadataEvent(ctx, input) +} + +type metadataReplayStateStore interface { + MetadataReplayStateOp(ctx context.Context, sessionGID string, field string) (string, bool, error) +} + +func (s *Server) metadataReplayStateOp( + ctx context.Context, + sessionID string, + field string, +) (string, bool, error) { + store, ok := s.db.(metadataReplayStateStore) + if !ok { + return "", false, nil + } + origin := s.localArtifactOrigin() + if origin == "" { + return "", false, nil + } + return store.MetadataReplayStateOp(ctx, artifact.MetadataSessionGID(origin, sessionID), field) +} + +func renameMetadataValue(displayName *string) (json.RawMessage, error) { + data, err := json.Marshal(struct { + DisplayName *string `json:"display_name"` + }{DisplayName: displayName}) + if err != nil { + return nil, err + } + return json.RawMessage(data), nil +} + +func (s *Server) metadataPinForMessage( + ctx context.Context, + sessionID string, + messageID int64, + note *string, +) (*artifact.MetadataPin, error) { + msg, err := s.db.GetMessageForMetadataPin(ctx, sessionID, messageID) + if err != nil { + return nil, fmt.Errorf("loading message for metadata pin: %w", err) + } + if msg == nil { + return nil, nil + } + pin := &artifact.MetadataPin{ + SourceUUID: msg.SourceUUID, + Ordinal: msg.Ordinal, + } + if note != nil { + noteCopy := *note + pin.Note = ¬eCopy + } + return pin, nil +} diff --git a/internal/server/metadata_events_test.go b/internal/server/metadata_events_test.go new file mode 100644 index 000000000..9787b7737 --- /dev/null +++ b/internal/server/metadata_events_test.go @@ -0,0 +1,1026 @@ +package server_test + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "sort" + "strings" + "testing" + + _ "github.com/mattn/go-sqlite3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/agentsview/internal/artifact" + "go.kenn.io/agentsview/internal/config" + "go.kenn.io/agentsview/internal/db" +) + +type recordedMetadataEvent struct { + Version int `json:"v"` + HLC string `json:"hlc"` + Origin string `json:"origin"` + SessionGID string `json:"session_gid"` + Op string `json:"op"` + Value map[string]any `json:"value,omitempty"` + Pin *recordedMetadataPin `json:"pin,omitempty"` +} + +type recordedMetadataPin struct { + SourceUUID string `json:"source_uuid,omitempty"` + Ordinal int `json:"ordinal,omitempty"` + Note *string `json:"note,omitempty"` +} + +func withArtifactOrigin(origin string) setupOption { + return func(c *config.Config) { c.ArtifactOriginID = origin } +} + +func TestMetadataEventsAppendForUserMutations(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + te.seedMessages(t, "s1", 2, func(i int, m *db.Message) { + if i == 1 { + m.SourceUUID = "uuid-answer" + } + }) + msgs, err := te.db.GetAllMessages(context.Background(), "s1") + require.NoError(t, err) + require.Len(t, msgs, 2) + messageID := msgs[1].ID + + w := te.patch(t, "/api/v1/sessions/s1/rename", `{"display_name":"Pinned investigation"}`) + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + + w = te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + w = te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + w = te.post(t, fmt.Sprintf("/api/v1/sessions/s1/messages/%d/pin", messageID), `{"note":"remember"}`) + require.Equal(t, http.StatusCreated, w.Code, "body: %s", w.Body.String()) + + w = te.del(t, fmt.Sprintf("/api/v1/sessions/s1/messages/%d/pin", messageID)) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + w = te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + w = te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + w = te.del(t, "/api/v1/sessions/s1/permanent") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 9) + assert.Equal(t, []string{ + artifact.MetadataOpRename, + artifact.MetadataOpStar, + artifact.MetadataOpUnstar, + artifact.MetadataOpPin, + artifact.MetadataOpUnpin, + artifact.MetadataOpSoftDelete, + artifact.MetadataOpRestore, + artifact.MetadataOpSoftDelete, + artifact.MetadataOpPurge, + }, metadataOps(events)) + for _, event := range events { + assert.Equal(t, 1, event.Version) + assert.NotEmpty(t, event.HLC) + assert.Equal(t, "desk-a1b2c3", event.Origin) + assert.Equal(t, "desk-a1b2c3~s1", event.SessionGID) + } + assert.Equal(t, "Pinned investigation", events[0].Value["display_name"]) + require.NotNil(t, events[3].Pin) + assert.Equal(t, "uuid-answer", events[3].Pin.SourceUUID) + assert.Equal(t, 1, events[3].Pin.Ordinal) + require.NotNil(t, events[3].Pin.Note) + assert.Equal(t, "remember", *events[3].Pin.Note) + require.NotNil(t, events[4].Pin) + assert.Equal(t, "uuid-answer", events[4].Pin.SourceUUID) + assert.Equal(t, 1, events[4].Pin.Ordinal) + assert.Nil(t, events[4].Pin.Note) +} + +func TestMetadataEventsAppendWithoutSyncEngine(t *testing.T) { + te := setupNoSyncMode(t) + require.NoError(t, artifact.AdoptOrigin(te.db, "nosync-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.patch(t, "/api/v1/sessions/s1/rename", `{"display_name":"No sync title"}`) + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpRename, events[0].Op) + assert.Equal(t, "nosync-a1b2c3", events[0].Origin) + assert.Equal(t, "nosync-a1b2c3~s1", events[0].SessionGID) +} + +func TestMetadataEventsNotRecordedWithoutOptIn(t *testing.T) { + te := setup(t) + te.seedSession(t, "s1", "alpha", 2) + + w := te.patch(t, "/api/v1/sessions/s1/rename", `{"display_name":"Local only"}`) + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + w = te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + assert.Empty(t, readMetadataEvents(t, te), + "curation must not write ledger events before the machine opts into artifact sync") + origin, err := artifact.StoredOrigin(te.db) + require.NoError(t, err) + assert.Empty(t, origin, + "curation must not mint an artifact origin before opt-in") +} + +func TestMetadataEventsEmptyTrashStaysLocalOnly(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + for _, id := range []string{"s1", "s2"} { + te.seedSession(t, id, "alpha", 2) + w := te.del(t, "/api/v1/sessions/"+id) + require.Equal(t, http.StatusNoContent, w.Code, "delete %s body: %s", id, w.Body.String()) + } + + w := te.del(t, "/api/v1/trash") + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 2) + assert.Equal(t, []string{ + artifact.MetadataOpSoftDelete, + artifact.MetadataOpSoftDelete, + }, metadataOps(events)) +} + +func TestMetadataEventsBatchDeleteRecordsNewAndUnrecordedTrashedSessions(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + for _, id := range []string{"s1", "s2", "s3"} { + te.seedSession(t, id, "alpha", 2) + } + require.NoError(t, te.db.SoftDeleteSession("s3")) + + w := te.requestJSON(t, http.MethodPost, "/api/v1/sessions/batch-delete", + `{"session_ids":["s1","s2","s3","missing"]}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 3) + assert.Equal(t, []string{ + artifact.MetadataOpSoftDelete, + artifact.MetadataOpSoftDelete, + artifact.MetadataOpSoftDelete, + }, metadataOps(events)) + assert.ElementsMatch(t, []string{ + "desk-a1b2c3~s1", + "desk-a1b2c3~s2", + "desk-a1b2c3~s3", + }, []string{events[0].SessionGID, events[1].SessionGID, events[2].SessionGID}) +} + +func TestMetadataEventsBatchDeleteRetriesAlreadyDeletedSessions(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + for _, id := range []string{"s1", "s2"} { + te.seedSession(t, id, "alpha", 2) + require.NoError(t, te.db.SoftDeleteSession(id)) + } + + w := te.requestJSON(t, http.MethodPost, "/api/v1/sessions/batch-delete", + `{"session_ids":["s1","s2","missing"]}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 2) + assert.Equal(t, []string{ + artifact.MetadataOpSoftDelete, + artifact.MetadataOpSoftDelete, + }, metadataOps(events)) + assert.ElementsMatch(t, []string{ + "desk-a1b2c3~s1", + "desk-a1b2c3~s2", + }, []string{events[0].SessionGID, events[1].SessionGID}) +} + +func TestMetadataEventsBatchDeleteRepairsPublishedFailureOnRetry(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + for _, id := range []string{"s1", "s2"} { + te.seedSession(t, id, "alpha", 2) + } + execTestDDL(t, te, ` +CREATE TRIGGER fail_s1_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +WHEN NEW.session_gid = 'desk-a1b2c3~s1' +BEGIN + SELECT RAISE(FAIL, 'forced s1 metadata replay failure'); +END; +`) + + w := te.requestJSON(t, http.MethodPost, "/api/v1/sessions/batch-delete", + `{"session_ids":["s1","s2"]}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + + s1, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, s1) + assert.NotNil(t, s1.DeletedAt, + "published failure keeps the session whose artifact is durable in trash") + s2, err := te.db.GetSessionFull(context.Background(), "s2") + require.NoError(t, err) + require.NotNil(t, s2) + assert.Nil(t, s2.DeletedAt, + "later unpublished sessions are restored after the failure") + assert.Equal(t, 0, serverMetadataTableCount( + t, te, "metadata_replay_state", + "session_gid = 'desk-a1b2c3~s1' AND field = 'deleted_at'", + )) + firstEvents := readMetadataEvents(t, te) + require.Len(t, firstEvents, 1) + assert.Equal(t, "desk-a1b2c3~s1", firstEvents[0].SessionGID) + + execTestDDL(t, te, `DROP TRIGGER fail_s1_metadata_replay_state_insert`) + w = te.requestJSON(t, http.MethodPost, "/api/v1/sessions/batch-delete", + `{"session_ids":["s1","s2"]}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + for _, id := range []string{"s1", "s2"} { + sess, getErr := te.db.GetSessionFull(context.Background(), id) + require.NoError(t, getErr) + require.NotNil(t, sess) + assert.NotNil(t, sess.DeletedAt, "retry must trash %s", id) + assert.Equal(t, artifact.MetadataOpSoftDelete, + serverMetadataReplayOp(t, te, "desk-a1b2c3~"+id, "deleted_at")) + } + events := readMetadataEvents(t, te) + require.Len(t, events, 2, + "retry must repair the first artifact instead of publishing a duplicate") + counts := map[string]int{} + for _, event := range events { + counts[event.SessionGID]++ + } + assert.Equal(t, map[string]int{ + "desk-a1b2c3~s1": 1, + "desk-a1b2c3~s2": 1, + }, counts) +} + +func TestMetadataEventsUnstarOnlyRecordsRemovedStars(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/missing/star") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + w = te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Empty(t, readMetadataEvents(t, te)) + + ok, err := te.db.StarSession("s1") + require.NoError(t, err) + require.True(t, ok) + w = te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpUnstar, events[0].Op) + assert.Equal(t, "desk-a1b2c3~s1", events[0].SessionGID) +} + +func TestMetadataEventsUnstarRestoresStarWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + ok, err := te.db.StarSession("s1") + require.NoError(t, err) + require.True(t, ok) + + breakMetadataArtifactStore(t, te) + + w := te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"s1"}, ids) + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_replay_state", "session_gid = 'desk-a1b2c3~s1'")) + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_applied_events", "origin = 'desk-a1b2c3'")) + + repairMetadataArtifactStore(t, te) + w = te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + ids, err = te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Empty(t, ids) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpUnstar, events[0].Op) + assert.Equal(t, "desk-a1b2c3~s1", events[0].SessionGID) +} + +func TestMetadataEventsUnstarDoesNotRestoreStarWhenArtifactPublished(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + ok, err := te.db.StarSession("s1") + require.NoError(t, err) + require.True(t, ok) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w := te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Empty(t, ids) + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_replay_state", "session_gid = 'desk-a1b2c3~s1'")) + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_applied_events", "origin = 'desk-a1b2c3'")) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpUnstar, events[0].Op) + assert.Equal(t, "desk-a1b2c3~s1", events[0].SessionGID) +} + +func TestMetadataEventsNoopUnstarRepairsPublishedArtifactState(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpStar, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "starred")) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w = te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpStar, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "starred")) + unstarOrderKey := metadataEventOrderKey(t, te, artifact.MetadataOpUnstar) + + execTestDDL(t, te, `DROP TRIGGER fail_metadata_replay_state_insert`) + w = te.del(t, "/api/v1/sessions/s1/star") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpUnstar, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "starred")) + + remoteHLC, remoteHash := splitMetadataOrderKey(t, unstarOrderKey) + _, err := te.db.ApplyMetadataProjection(context.Background(), db.MetadataProjection{ + EventOrigin: "peer-b2c3d4", + OrderKey: unstarOrderKey, + HLC: remoteHLC, + ArtifactHash: remoteHash, + SessionGID: "desk-a1b2c3~s1", + LocalSessionID: "s1", + Field: "starred", + Op: artifact.MetadataOpStar, + Value: artifact.MetadataOpStar, + }) + require.NoError(t, err) + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Empty(t, ids) +} + +func TestMetadataEventsStarRollsBackWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + breakMetadataArtifactStore(t, te) + + w := te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Empty(t, ids, "failed metadata publish must roll back the star") + + repairMetadataArtifactStore(t, te) + w = te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpStar, events[0].Op) + + // A pre-existing star survives a later failed re-star: the failed + // append changed nothing this request needs to undo. + // (breakMetadataArtifactStore also wipes recorded events.) + breakMetadataArtifactStore(t, te) + w = te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + ids, err = te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"s1"}, ids, + "failed re-star must not remove the pre-existing star") +} + +func TestMetadataEventsStarKeepsStarWhenArtifactPublished(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w := te.put(t, "/api/v1/sessions/s1/star", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"s1"}, ids, + "published metadata event must keep the local star") + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpStar, events[0].Op) +} + +func TestMetadataEventsBulkStarRollsBackWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + te.seedSession(t, "s2", "alpha", 2) + ok, err := te.db.StarSession("s1") + require.NoError(t, err) + require.True(t, ok) + + breakMetadataArtifactStore(t, te) + w := te.requestJSON(t, http.MethodPost, "/api/v1/starred/bulk", + `{"session_ids":["s1","s2"]}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + + ids, err := te.db.ListStarredSessionIDs(context.Background()) + require.NoError(t, err) + assert.Equal(t, []string{"s1"}, ids, + "failed metadata publish must roll back only the stars this request created") + assert.Empty(t, readMetadataEvents(t, te)) +} + +func TestMetadataEventsRenameRestoresNameWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + w := te.patch(t, "/api/v1/sessions/s1/rename", `{"display_name":"keep"}`) + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + + breakMetadataArtifactStore(t, te) + w = te.patch(t, "/api/v1/sessions/s1/rename", `{"display_name":"replace"}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + + session, err := te.db.GetSession(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, session) + require.NotNil(t, session.DisplayName) + assert.Equal(t, "keep", *session.DisplayName, + "failed metadata publish must restore the prior display name") + // breakMetadataArtifactStore also wiped the first rename's event, so + // no events remain after the rolled-back rename. + assert.Empty(t, readMetadataEvents(t, te)) +} + +func TestMetadataEventsDeleteRestoresSessionWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + breakMetadataArtifactStore(t, te) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + session, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, session) + assert.Nil(t, session.DeletedAt, + "failed metadata publish must restore the soft-deleted session") + + repairMetadataArtifactStore(t, te) + w = te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpSoftDelete, events[0].Op) +} + +func seedPinTestMessage(t *testing.T, te *testEnv) int64 { + t.Helper() + te.seedSession(t, "s1", "alpha", 2) + te.seedMessages(t, "s1", 2) + msgs, err := te.db.GetAllMessages(context.Background(), "s1") + require.NoError(t, err) + require.Len(t, msgs, 2) + return msgs[1].ID +} + +func breakMetadataArtifactStore(t *testing.T, te *testEnv) { + t.Helper() + require.NotNil(t, te.artifactFault) + for _, entry := range listMetadataEntries(t, te) { + require.NoError(t, te.artifactStore.Trash(t.Context(), entry.Ref)) + } + te.artifactFault.setCreateError( + errors.New("forced metadata artifact write failure"), + ) +} + +func repairMetadataArtifactStore(t *testing.T, te *testEnv) { + t.Helper() + require.NotNil(t, te.artifactFault) + te.artifactFault.setCreateError(nil) +} + +func TestMetadataEventsPinRollsBackWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + messageID := seedPinTestMessage(t, te) + breakMetadataArtifactStore(t, te) + pinPath := fmt.Sprintf("/api/v1/sessions/s1/messages/%d/pin", messageID) + + w := te.post(t, pinPath, `{"note":"remember"}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + pins, err := te.db.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + assert.Empty(t, pins, "failed metadata publish must roll back the pin") + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_replay_state", "session_gid = 'desk-a1b2c3~s1'")) + assert.Equal(t, 0, serverMetadataTableCount(t, te, "metadata_applied_events", "origin = 'desk-a1b2c3'")) + + repairMetadataArtifactStore(t, te) + w = te.post(t, pinPath, `{"note":"remember"}`) + require.Equal(t, http.StatusCreated, w.Code, "body: %s", w.Body.String()) + pins, err = te.db.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + require.Len(t, pins, 1) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpPin, events[0].Op) +} + +func TestMetadataEventsRepinRestoresPriorNoteWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + messageID := seedPinTestMessage(t, te) + pinPath := fmt.Sprintf("/api/v1/sessions/s1/messages/%d/pin", messageID) + + w := te.post(t, pinPath, `{"note":"keep"}`) + require.Equal(t, http.StatusCreated, w.Code, "body: %s", w.Body.String()) + + breakMetadataArtifactStore(t, te) + w = te.post(t, pinPath, `{"note":"replace"}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + + pins, err := te.db.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + require.Len(t, pins, 1) + require.NotNil(t, pins[0].Note) + assert.Equal(t, "keep", *pins[0].Note, + "failed re-pin must restore the prior note") +} + +func TestMetadataEventsUnpinRestoresPinWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + messageID := seedPinTestMessage(t, te) + pinPath := fmt.Sprintf("/api/v1/sessions/s1/messages/%d/pin", messageID) + + w := te.post(t, pinPath, `{"note":"remember"}`) + require.Equal(t, http.StatusCreated, w.Code, "body: %s", w.Body.String()) + + breakMetadataArtifactStore(t, te) + w = te.del(t, pinPath) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + pins, err := te.db.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + require.Len(t, pins, 1, "failed metadata publish must restore the pin") + require.NotNil(t, pins[0].Note) + assert.Equal(t, "remember", *pins[0].Note) + + repairMetadataArtifactStore(t, te) + w = te.del(t, pinPath) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + pins, err = te.db.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + assert.Empty(t, pins) + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpUnpin, events[0].Op) +} + +func TestMetadataEventsPinKeepsPinWhenArtifactPublished(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + messageID := seedPinTestMessage(t, te) + pinPath := fmt.Sprintf("/api/v1/sessions/s1/messages/%d/pin", messageID) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w := te.post(t, pinPath, `{"note":"remember"}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + + pins, err := te.db.ListPinnedMessages(context.Background(), "s1", "") + require.NoError(t, err) + require.Len(t, pins, 1, + "published metadata event must keep the local pin") + + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpPin, events[0].Op) +} + +func TestMetadataEventsPermanentDeleteRetriesExcludedSession(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w = te.del(t, "/api/v1/sessions/s1/permanent") + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + got, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + assert.Nil(t, got) + assert.True(t, te.db.IsSessionExcluded("s1")) + deleteMetadataEventsByOp(t, te, artifact.MetadataOpPurge) + + execTestDDL(t, te, `DROP TRIGGER fail_metadata_replay_state_insert`) + w = te.del(t, "/api/v1/sessions/s1/permanent") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + assert.Equal(t, []string{ + artifact.MetadataOpSoftDelete, + artifact.MetadataOpPurge, + }, metadataOps(events)) + assert.Equal(t, artifact.MetadataOpPurge, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "purge")) +} + +func TestMetadataEventsPermanentDeleteRetainsSessionWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + breakMetadataArtifactStore(t, te) + + w = te.del(t, "/api/v1/sessions/s1/permanent") + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + session, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, session, "purge must not remove the only local copy before publication") + assert.NotNil(t, session.DeletedAt) + assert.False(t, te.db.IsSessionExcluded("s1")) + assert.Empty(t, readMetadataEvents(t, te)) + + repairMetadataArtifactStore(t, te) + w = te.del(t, "/api/v1/sessions/s1/permanent") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + session, err = te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + assert.Nil(t, session) + assert.True(t, te.db.IsSessionExcluded("s1")) + events := readMetadataEvents(t, te) + require.Len(t, events, 1) + assert.Equal(t, artifact.MetadataOpPurge, events[0].Op) +} + +func TestMetadataEventsRestoreRepairsPublishedArtifactState(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpSoftDelete, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + restored, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, restored) + assert.Nil(t, restored.DeletedAt) + assert.Equal(t, artifact.MetadataOpSoftDelete, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + restoreOrderKey := metadataEventOrderKey(t, te, artifact.MetadataOpRestore) + + execTestDDL(t, te, `DROP TRIGGER fail_metadata_replay_state_insert`) + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpRestore, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + + restoreHLC, _ := splitMetadataOrderKey(t, restoreOrderKey) + olderHash := strings.Repeat("0", 64) + _, err = te.db.ApplyMetadataProjection(context.Background(), db.MetadataProjection{ + EventOrigin: "peer-b2c3d4", + OrderKey: restoreHLC + "-" + olderHash, + HLC: restoreHLC, + ArtifactHash: olderHash, + SessionGID: "desk-a1b2c3~s1", + LocalSessionID: "s1", + Field: "deleted_at", + Op: artifact.MetadataOpSoftDelete, + Value: artifact.MetadataOpSoftDelete, + }) + require.NoError(t, err) + restored, err = te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, restored) + assert.Nil(t, restored.DeletedAt) +} + +func TestMetadataEventsRestoreRetrashesSessionWhenArtifactWriteFails(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + breakMetadataArtifactStore(t, te) + + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + session, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, session) + assert.NotNil(t, session.DeletedAt, + "failed pre-publication restore must return the session to trash") + assert.Equal(t, artifact.MetadataOpSoftDelete, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + assert.Empty(t, readMetadataEvents(t, te)) +} + +func TestMetadataEventsRestoreRetriesWithoutPublishedArtifact(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + restored, err := te.db.GetSessionFull(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, restored) + assert.Nil(t, restored.DeletedAt) + deleteMetadataEventsByOp(t, te, artifact.MetadataOpRestore) + + execTestDDL(t, te, `DROP TRIGGER fail_metadata_replay_state_insert`) + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + events := readMetadataEvents(t, te) + assert.Equal(t, []string{ + artifact.MetadataOpSoftDelete, + artifact.MetadataOpRestore, + }, metadataOps(events)) + assert.Equal(t, artifact.MetadataOpRestore, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) +} + +func TestMetadataEventsRestoreRetryPublishesWhenOlderRestoreLoses(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + w := te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + olderRestore := metadataEventOrderKey(t, te, artifact.MetadataOpRestore) + + w = te.del(t, "/api/v1/sessions/s1") + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpSoftDelete, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + + execTestDDL(t, te, ` +CREATE TRIGGER fail_metadata_replay_state_insert +BEFORE INSERT ON metadata_replay_state +BEGIN + SELECT RAISE(FAIL, 'forced metadata replay failure'); +END; +`) + + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusInternalServerError, w.Code, "body: %s", w.Body.String()) + assert.Equal(t, artifact.MetadataOpSoftDelete, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + for _, key := range metadataEventOrderKeys(t, te, artifact.MetadataOpRestore) { + if key != olderRestore { + deleteMetadataEventOrderKey(t, te, key) + } + } + + execTestDDL(t, te, `DROP TRIGGER fail_metadata_replay_state_insert`) + w = te.post(t, "/api/v1/sessions/s1/restore", `{}`) + require.Equal(t, http.StatusNoContent, w.Code, "body: %s", w.Body.String()) + + assert.Equal(t, artifact.MetadataOpRestore, + serverMetadataReplayOp(t, te, "desk-a1b2c3~s1", "deleted_at")) + assert.Len(t, metadataEventOrderKeys(t, te, artifact.MetadataOpRestore), 2) +} + +func TestMetadataEventsSuppressedDuringReplay(t *testing.T) { + te := setup(t, withArtifactOrigin("desk-a1b2c3")) + te.seedSession(t, "s1", "alpha", 2) + + ctx := artifact.WithMetadataEventSuppression(context.Background()) + req := httptest.NewRequest( + http.MethodPatch, + "/api/v1/sessions/s1/rename", + strings.NewReader(`{"display_name":"Replay name"}`), + ).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", "http://127.0.0.1:0") + w := httptest.NewRecorder() + te.handler.ServeHTTP(w, req) + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + + renamed, err := te.db.GetSession(context.Background(), "s1") + require.NoError(t, err) + require.NotNil(t, renamed) + require.NotNil(t, renamed.DisplayName) + assert.Equal(t, "Replay name", *renamed.DisplayName) + assert.Empty(t, readMetadataEvents(t, te)) +} + +// execTestDDL runs schema DDL (failure-injection triggers) on the test +// database through a short-lived write connection. The server's Reader() pool +// opens with mode=ro, so tests cannot install triggers through it. +func execTestDDL(t *testing.T, te *testEnv, stmt string) { + t.Helper() + raw, err := sql.Open("sqlite3", "file:"+te.db.Path()+"?_busy_timeout=5000") + require.NoError(t, err) + defer func() { require.NoError(t, raw.Close()) }() + _, err = raw.Exec(stmt) + require.NoError(t, err) +} + +func readMetadataEvents(t *testing.T, te *testEnv) []recordedMetadataEvent { + t.Helper() + entries := listMetadataEntries(t, te) + events := make([]recordedMetadataEvent, 0, len(entries)) + for _, entry := range entries { + data := readArtifactEntry(t, te, entry) + var event recordedMetadataEvent + require.NoError(t, json.Unmarshal(data, &event), "artifact %s", entry.Ref.Name) + events = append(events, event) + } + return events +} + +func listMetadataEntries(t *testing.T, te *testEnv) []artifact.Entry { + t.Helper() + if te.artifactStore == nil { + return nil + } + origins := collectArtifactOrigins(t, te.artifactStore, 64) + + var entries []artifact.Entry + for _, origin := range origins { + entries = append(entries, + collectArtifactEntries(t, te.artifactStore, origin, artifact.KindMeta, 64)...) + } + sort.Slice(entries, func(i, j int) bool { + if entries[i].Ref.Origin != entries[j].Ref.Origin { + return entries[i].Ref.Origin < entries[j].Ref.Origin + } + return entries[i].Ref.Name < entries[j].Ref.Name + }) + return entries +} + +func readArtifactEntry( + t *testing.T, te *testEnv, entry artifact.Entry, +) []byte { + t.Helper() + _, reader, err := te.artifactStore.Open(t.Context(), entry.Ref) + require.NoError(t, err) + data, readErr := io.ReadAll(reader) + require.NoError(t, readErr) + require.NoError(t, reader.Verify()) + require.NoError(t, reader.Close()) + return data +} + +func metadataEventOrderKey(t *testing.T, te *testEnv, op string) string { + t.Helper() + keys := metadataEventOrderKeys(t, te, op) + require.NotEmpty(t, keys, "metadata event op %s not found", op) + return keys[0] +} + +func metadataEventOrderKeys(t *testing.T, te *testEnv, op string) []string { + t.Helper() + keys := make([]string, 0) + for _, entry := range listMetadataEntries(t, te) { + data := readArtifactEntry(t, te, entry) + var event recordedMetadataEvent + require.NoError(t, json.Unmarshal(data, &event)) + if event.Op == op { + keys = append(keys, strings.TrimSuffix(entry.Ref.Name, ".json")) + } + } + return keys +} + +func splitMetadataOrderKey(t *testing.T, orderKey string) (string, string) { + t.Helper() + idx := strings.LastIndex(orderKey, "-") + require.NotEqual(t, -1, idx, "order key %q missing hash suffix", orderKey) + return orderKey[:idx], orderKey[idx+1:] +} + +func deleteMetadataEventOrderKey(t *testing.T, te *testEnv, orderKey string) { + t.Helper() + var matches []artifact.Entry + for _, entry := range listMetadataEntries(t, te) { + if entry.Ref.Name == orderKey+".json" { + matches = append(matches, entry) + } + } + require.Len(t, matches, 1) + require.NoError(t, te.artifactStore.Trash(t.Context(), matches[0].Ref)) +} + +func deleteMetadataEventsByOp(t *testing.T, te *testEnv, op string) { + t.Helper() + for _, entry := range listMetadataEntries(t, te) { + data := readArtifactEntry(t, te, entry) + var event recordedMetadataEvent + require.NoError(t, json.Unmarshal(data, &event)) + if event.Op == op { + require.NoError(t, te.artifactStore.Trash(t.Context(), entry.Ref)) + } + } +} + +func metadataOps(events []recordedMetadataEvent) []string { + ops := make([]string, len(events)) + for i, event := range events { + ops[i] = event.Op + } + return ops +} + +func serverMetadataReplayOp(t *testing.T, te *testEnv, sessionGID, field string) string { + t.Helper() + var op string + err := te.db.Reader().QueryRow( + `SELECT op FROM metadata_replay_state WHERE session_gid = ? AND field = ?`, + sessionGID, field, + ).Scan(&op) + require.NoError(t, err) + return op +} + +func serverMetadataTableCount(t *testing.T, te *testEnv, table, where string) int { + t.Helper() + var count int + err := te.db.Reader().QueryRow("SELECT COUNT(*) FROM " + table + " WHERE " + where).Scan(&count) + require.NoError(t, err) + return count +} diff --git a/internal/server/openapi.go b/internal/server/openapi.go index ad845be6a..109592d8b 100644 --- a/internal/server/openapi.go +++ b/internal/server/openapi.go @@ -17,8 +17,9 @@ func OpenAPISpec(version VersionInfo, opts ...Option) *huma.OpenAPI { cfg: config.Config{ WriteTimeout: 30 * time.Second, }, - mux: http.NewServeMux(), - version: version, + mux: http.NewServeMux(), + version: version, + documentArtifactRoutes: true, } for _, opt := range opts { opt(s) diff --git a/internal/server/server.go b/internal/server/server.go index 74d6fe3c1..50d2a1d15 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -2,6 +2,7 @@ package server import ( "context" + "errors" "fmt" "io" "io/fs" @@ -19,6 +20,7 @@ import ( "github.com/danielgtaylor/huma/v2" "github.com/danielgtaylor/huma/v2/adapters/humago" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/config" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/insight" @@ -56,20 +58,29 @@ const ( // Server is the HTTP server that serves the SPA and REST API. type Server struct { - mu gosync.RWMutex - cfg config.Config - db db.Store - engine *sync.Engine - onDemandEngine *sync.Engine - sessions service.SessionService - broadcaster *Broadcaster - mux *http.ServeMux - api huma.API - httpSrv *http.Server - version VersionInfo - dataDir string - - httpRemoteCleanupRegistry *remotesync.CleanupRegistry + mu gosync.RWMutex + sessionLifecycleMu gosync.Mutex + artifactImportPending bool + artifactBaselineDone bool + cfg config.Config + db db.Store + engine *sync.Engine + onDemandEngine *sync.Engine + sessions service.SessionService + broadcaster *Broadcaster + metadata *artifact.MetadataRecorder + metadataAppend func(context.Context, artifact.MetadataEventInput) error + mux *http.ServeMux + api huma.API + httpSrv *http.Server + version VersionInfo + dataDir string + + httpRemoteCleanupRegistry *remotesync.CleanupRegistry + artifactCursors *artifactCursorRegistry + artifactStoreCursorsMu gosync.Mutex + artifactStoreCursors *artifactStoreCursorRegistry + artifactStoreCursorsClosed bool // baseCtx, when set, is used as the base context for all // incoming requests. Cancelling it causes SSE handlers to @@ -87,6 +98,14 @@ type Server struct { // handler, used only by tests to guarantee handlers // exceed a short timeout. Zero in production. handlerDelay time.Duration + // beforeSessionLifecycleLock observes lock attempts in concurrency tests. + // Production servers leave it nil. + beforeSessionLifecycleLock func() + // beginArtifactRepositoryReset and republishArtifactRepositoryReset expose + // the reset commit boundary to the daemon lifecycle. Tests may hold the + // post-move publication phase deterministically. + beginArtifactRepositoryReset func(context.Context, string, string, *artifact.Repository, func() error) (*artifact.Repository, artifact.RepositoryResetResult, error) + republishArtifactRepositoryReset func(context.Context, string, *db.DB, string, *artifact.Repository, artifact.RepositoryResetResult) (artifact.RepositoryResetResult, error) // updateCheckFn is the function called to check for // updates. Defaults to update.CheckForUpdate; tests @@ -143,6 +162,240 @@ type Server struct { localResyncRunner LocalResyncRunner ensurePricing func(context.Context, *db.DB) error + + artifactStore artifact.ArtifactStore + artifactRepository *artifact.Repository + artifactOps artifactOperationLifetime + documentArtifactRoutes bool +} + +type artifactOperationLifetime struct { + mu gosync.Mutex + cond *gosync.Cond + store artifact.ArtifactStore + closeStore func(artifact.ArtifactStore) error + active int + resetting bool + resetCommitted bool + resetCancel context.CancelFunc + closing bool + closed bool + closeErr error +} + +func (l *artifactOperationLifetime) setStore( + store artifact.ArtifactStore, closeStore func(artifact.ArtifactStore) error, +) { + l.mu.Lock() + defer l.mu.Unlock() + l.store = store + l.closeStore = closeStore +} + +func (l *artifactOperationLifetime) close(store artifact.ArtifactStore) error { + if store == nil { + return nil + } + if l.closeStore != nil { + return l.closeStore(store) + } + closer, ok := any(store).(io.Closer) + if !ok { + return nil + } + return closer.Close() +} + +func (l *artifactOperationLifetime) acquire() (artifact.ArtifactStore, func(), error) { + l.mu.Lock() + if l.resetting || l.closing || l.closed || l.store == nil { + l.mu.Unlock() + return nil, nil, errors.New("artifact store is not available") + } + l.active++ + store := l.store + l.mu.Unlock() + return store, l.release, nil +} + +func (l *artifactOperationLifetime) release() { + var store artifact.ArtifactStore + l.mu.Lock() + if l.active > 0 { + l.active-- + } + if l.cond != nil { + l.cond.Broadcast() + } + if l.closing && !l.resetting && l.active == 0 && !l.closed { + l.closed = true + store = l.store + } + l.mu.Unlock() + if store != nil { + err := l.close(store) + l.mu.Lock() + l.closeErr = errors.Join(l.closeErr, err) + l.mu.Unlock() + } +} + +func (l *artifactOperationLifetime) beginReset( + ctx context.Context, +) (artifact.ArtifactStore, context.Context, error) { + if ctx == nil { + return nil, nil, errors.New("artifact store reset context is required") + } + l.mu.Lock() + if l.resetting || l.closing || l.closed || l.store == nil { + l.mu.Unlock() + return nil, nil, errors.New("artifact store is not available for reset") + } + resetCtx, cancel := context.WithCancel(ctx) + l.resetting = true + l.resetCommitted = false + l.resetCancel = cancel + if l.cond == nil { + l.cond = gosync.NewCond(&l.mu) + } + stopWake := context.AfterFunc(resetCtx, func() { + l.mu.Lock() + l.cond.Broadcast() + l.mu.Unlock() + }) + defer stopWake() + for l.active > 0 && resetCtx.Err() == nil && !l.closing { + l.cond.Wait() + } + cause := resetCtx.Err() + if l.closing || l.closed { + cause = errors.Join(cause, errors.New("artifact store is closing")) + } + if cause != nil { + l.resetting = false + l.resetCommitted = false + l.resetCancel = nil + l.cond.Broadcast() + var closeStore artifact.ArtifactStore + if l.closing && l.active == 0 && !l.closed { + l.closed = true + closeStore = l.store + } + l.mu.Unlock() + if closeStore != nil { + closeErr := l.close(closeStore) + l.mu.Lock() + l.closeErr = errors.Join(l.closeErr, closeErr) + l.mu.Unlock() + cause = errors.Join(cause, closeErr) + } + cancel() + return nil, nil, cause + } + store := l.store + l.mu.Unlock() + return store, resetCtx, nil +} + +// commitResetMutation revalidates an admitted reset at Docbank's final +// pre-move boundary. The closing check and mutation commit are serialized, but +// the mutex is released before the filesystem move and SQLite republish. +func (l *artifactOperationLifetime) commitResetMutation( + ctx context.Context, +) error { + if ctx == nil { + return errors.New("artifact store reset context is required") + } + l.mu.Lock() + defer l.mu.Unlock() + if err := ctx.Err(); err != nil { + return err + } + if !l.resetting || l.closing || l.closed || l.store == nil { + return errors.New("artifact store is closing before reset mutation") + } + l.resetCommitted = true + return nil +} + +func (l *artifactOperationLifetime) setResetStore(store artifact.ArtifactStore) error { + l.mu.Lock() + defer l.mu.Unlock() + if !l.resetting || !l.resetCommitted { + return errors.New("artifact store reset mutation is not committed") + } + l.store = store + return nil +} + +func (l *artifactOperationLifetime) finishReset(store artifact.ArtifactStore) error { + var closeStore artifact.ArtifactStore + var cancel context.CancelFunc + l.mu.Lock() + if !l.resetting { + l.mu.Unlock() + return errors.New("artifact store reset is not active") + } + l.store = store + l.resetting = false + l.resetCommitted = false + cancel = l.resetCancel + l.resetCancel = nil + if l.closing && !l.closed { + l.closed = true + closeStore = store + } + if l.cond != nil { + l.cond.Broadcast() + } + l.mu.Unlock() + if cancel != nil { + cancel() + } + if closeStore == nil { + return nil + } + err := l.close(closeStore) + l.mu.Lock() + l.closeErr = errors.Join(l.closeErr, err) + l.mu.Unlock() + return err +} + +func (l *artifactOperationLifetime) closeWhenIdle() error { + var store artifact.ArtifactStore + var cancel context.CancelFunc + l.mu.Lock() + l.closing = true + cancel = l.resetCancel + if l.cond != nil { + l.cond.Broadcast() + } + if !l.resetting && l.active == 0 && !l.closed { + l.closed = true + store = l.store + } + existingErr := l.closeErr + l.mu.Unlock() + if cancel != nil { + cancel() + } + if store == nil { + return existingErr + } + err := l.close(store) + l.mu.Lock() + l.closeErr = errors.Join(l.closeErr, err) + joined := l.closeErr + l.mu.Unlock() + return joined +} + +func (s *Server) lockSessionLifecycle() { + if s.beforeSessionLifecycleLock != nil { + s.beforeSessionLifecycleLock() + } + s.sessionLifecycleMu.Lock() } // New creates a new Server. @@ -168,15 +421,19 @@ func New( } s := &Server{ - cfg: cfg, - db: database, - engine: engine, - sessions: sessions, - mux: http.NewServeMux(), - httpRemoteCleanupRegistry: new(remotesync.CleanupRegistry), - insightLogDrainTimeout: defaultInsightLogDrainTimeout, - insightLogStopWaitTimeout: defaultInsightLogStopWaitTimeout, - ensurePricing: pricingrefresh.EnsureCurrent, + cfg: cfg, + db: database, + engine: engine, + sessions: sessions, + mux: http.NewServeMux(), + httpRemoteCleanupRegistry: new(remotesync.CleanupRegistry), + artifactCursors: newArtifactCursorRegistry(), + artifactStoreCursors: newArtifactStoreCursorRegistry(), + insightLogDrainTimeout: defaultInsightLogDrainTimeout, + insightLogStopWaitTimeout: defaultInsightLogStopWaitTimeout, + ensurePricing: pricingrefresh.EnsureCurrent, + beginArtifactRepositoryReset: artifact.BeginRepositoryReset, + republishArtifactRepositoryReset: artifact.RepublishRepositoryReset, generateStreamFunc: func( ctx context.Context, agent, prompt string, onLog insight.LogFunc, @@ -194,6 +451,13 @@ func New( for _, opt := range opts { opt(s) } + s.artifactOps.setStore(s.artifactStore, s.closeOwnedArtifactStore) + if local, ok := database.(*db.DB); ok && !local.ReadOnly() && s.artifactStore != nil { + s.metadata = artifact.NewMetadataRecorder(local, artifact.MetadataRecorderOptions{ + Origin: cfg.ArtifactOriginID, + Store: s.artifactStore, + }) + } if s.version.APIVersion == 0 { s.version.APIVersion = APIVersion } @@ -237,6 +501,24 @@ func WithDataDir(dir string) Option { return func(s *Server) { s.dataDir = dir } } +// WithArtifactStore gives the server ownership of the one logical artifact +// store shared by every artifact route. Shutdown closes it after active +// artifact operations release their leases. +func WithArtifactStore(store artifact.ArtifactStore) Option { + return func(s *Server) { s.artifactStore = store } +} + +// WithArtifactRepository gives the server ownership of the concrete local +// repository while exposing only its content boundary to route operations. +func WithArtifactRepository(repository *artifact.Repository) Option { + return func(s *Server) { + s.artifactRepository = repository + if repository != nil { + s.artifactStore = repository.Content() + } + } +} + // WithBaseContext sets the base context for all incoming HTTP // requests. When this context is cancelled, request contexts // are also cancelled, causing long-lived handlers (SSE) to @@ -423,6 +705,10 @@ func (s *Server) routes() { } func (s *Server) handleSPA(w http.ResponseWriter, r *http.Request) { + if strings.HasPrefix(r.URL.Path, "/api/v1/artifacts") { + http.NotFound(w, r) + return + } // Try to serve the exact file path := strings.TrimPrefix(r.URL.Path, "/") if path == "" { @@ -977,7 +1263,12 @@ func localInterfaceIPs() map[string]bool { } // ListenAndServe starts the HTTP server. -func (s *Server) ListenAndServe() error { +func (s *Server) ListenAndServe() (retErr error) { + defer func() { + if closeErr := s.closeArtifactResources(); closeErr != nil { + retErr = errors.Join(retErr, closeErr) + } + }() addr := fmt.Sprintf("%s:%d", s.cfg.Host, s.cfg.Port) listenCtx := context.Background() if s.baseCtx != nil { @@ -998,7 +1289,12 @@ func (s *Server) ListenAndServe() error { } // Serve starts the HTTP server on an existing listener. -func (s *Server) Serve(ln net.Listener) error { +func (s *Server) Serve(ln net.Listener) (retErr error) { + defer func() { + if closeErr := s.closeArtifactResources(); closeErr != nil { + retErr = errors.Join(retErr, closeErr) + } + }() addr := ln.Addr().String() srv := &http.Server{ Addr: addr, @@ -1019,6 +1315,25 @@ func (s *Server) Serve(ln net.Listener) error { return srv.Serve(ln) } +func (s *Server) closeArtifactResources() error { + if s.artifactCursors != nil { + s.artifactCursors.close() + } + s.closeArtifactStoreCursorRegistry() + return s.artifactOps.closeWhenIdle() +} + +func (s *Server) closeOwnedArtifactStore(store artifact.ArtifactStore) error { + if s.artifactRepository != nil { + return s.artifactRepository.Close() + } + closer, ok := any(store).(io.Closer) + if !ok { + return nil + } + return closer.Close() +} + // Shutdown gracefully shuts down the HTTP server, then closes the // server-owned on-demand sync engine (if one was lazily created) so // its pending debounced signal recomputes flush while the DB is @@ -1038,7 +1353,7 @@ func (s *Server) Shutdown(ctx context.Context) error { if engine != nil { engine.Close() } - return err + return errors.Join(err, s.closeArtifactResources()) } // FindAvailablePort finds an available port starting from the diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 27cfa39b8..918878908 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -8,6 +8,7 @@ import ( "encoding/json" "errors" "fmt" + "io" "mime/multipart" "net" "net/http" @@ -26,6 +27,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.kenn.io/agentsview/internal/artifact" "go.kenn.io/agentsview/internal/config" "go.kenn.io/agentsview/internal/db" "go.kenn.io/agentsview/internal/dbtest" @@ -51,13 +53,65 @@ const ( // testEnv sets up a server with a temporary database. type testEnv struct { - srv *server.Server - handler http.Handler - db *db.DB - engine *sync.Engine - broadcaster *server.Broadcaster - claudeDir string - dataDir string + srv *server.Server + handler http.Handler + db *db.DB + engine *sync.Engine + broadcaster *server.Broadcaster + claudeDir string + dataDir string + artifactStore artifact.ArtifactStore + artifactFault *faultInjectArtifactStore +} + +type faultInjectArtifactStore struct { + artifact.ArtifactStore + mu stdlibsync.RWMutex + createErr error +} + +func (s *faultInjectArtifactStore) Create( + ctx context.Context, + ref artifact.Ref, + identity artifact.Identity, + mediaType string, + src io.Reader, +) (artifact.CreateResult, error) { + s.mu.RLock() + err := s.createErr + s.mu.RUnlock() + if err != nil { + return artifact.CreateResult{}, err + } + return s.ArtifactStore.Create(ctx, ref, identity, mediaType, src) +} + +func (s *faultInjectArtifactStore) setCreateError(err error) { + s.mu.Lock() + s.createErr = err + s.mu.Unlock() +} + +func (s *faultInjectArtifactStore) Origins( + ctx context.Context, +) (artifact.OriginIterator, error) { + return s.ArtifactStore.Origins(ctx) +} + +func (s *faultInjectArtifactStore) Entries( + ctx context.Context, origin string, kind artifact.Kind, +) (artifact.EntryIterator, error) { + return s.ArtifactStore.Entries(ctx, origin, kind) +} + +func (s *faultInjectArtifactStore) Quarantined( + ctx context.Context, +) (artifact.QuarantineIterator, error) { + return s.ArtifactStore.(artifact.ArtifactQuarantineStore).Quarantined(ctx) +} + +func (s *faultInjectArtifactStore) TrashQuarantined(ctx context.Context, token string) error { + return s.ArtifactStore.(artifact.ArtifactQuarantineStore).TrashQuarantined(ctx, token) } // setupOption customizes the config used by setup. @@ -98,6 +152,21 @@ func setupWithServerOpts( return setupWithServerOptsAndDBTemplate(t, srvOpts, nil, opts...) } +func setupArtifact( + t *testing.T, + opts ...setupOption, +) *testEnv { + return setupArtifactWithServerOpts(t, nil, opts...) +} + +func setupArtifactWithServerOpts( + t *testing.T, + srvOpts []server.Option, + opts ...setupOption, +) *testEnv { + return setupWithServerOptsAndDBTemplateMode(t, srvOpts, nil, true, opts...) +} + func setupWithDBTemplate( t *testing.T, dbFiles map[string][]byte, @@ -111,6 +180,16 @@ func setupWithServerOptsAndDBTemplate( srvOpts []server.Option, dbFiles map[string][]byte, opts ...setupOption, +) *testEnv { + return setupWithServerOptsAndDBTemplateMode(t, srvOpts, dbFiles, false, opts...) +} + +func setupWithServerOptsAndDBTemplateMode( + t *testing.T, + srvOpts []server.Option, + dbFiles map[string][]byte, + withArtifacts bool, + opts ...setupOption, ) *testEnv { t.Helper() dir := tempDirWithRetryCleanup(t) @@ -124,6 +203,9 @@ func setupWithServerOptsAndDBTemplate( for _, opt := range opts { opt(&cfg) } + if cfg.ArtifactOriginID != "" { + withArtifacts = true + } if dbFiles != nil { writeDBTemplateFiles(t, cfg.DBPath, dbFiles) } @@ -149,19 +231,38 @@ func setupWithServerOptsAndDBTemplate( Emitter: broadcaster, } engine := sync.NewEngine(database, engineCfg) + var artifactStore artifact.ArtifactStore + var artifactFault *faultInjectArtifactStore + if withArtifacts { + repository, err := artifact.OpenRepository(t.Context(), cfg.DataDir) + require.NoError(t, err) + artifactFault = &faultInjectArtifactStore{ + ArtifactStore: repository.Content(), + } + artifactStore = artifactFault + srvOpts = append(srvOpts, + server.WithArtifactRepository(repository), + server.WithArtifactStore(artifactStore), + ) + t.Cleanup(func() { + require.NoError(t, repository.Close()) + }) + } // Prepend so caller-provided srvOpts can still override. srvOpts = append([]server.Option{server.WithBroadcaster(broadcaster)}, srvOpts...) srv := server.New(cfg, database, engine, srvOpts...) return &testEnv{ - srv: srv, - handler: wrapTestHandler(cfg, srv.Handler()), - db: database, - engine: engine, - broadcaster: broadcaster, - claudeDir: claudeDir, - dataDir: dir, + srv: srv, + handler: wrapTestHandler(cfg, srv.Handler()), + db: database, + engine: engine, + broadcaster: broadcaster, + claudeDir: claudeDir, + dataDir: dir, + artifactStore: artifactStore, + artifactFault: artifactFault, } } @@ -261,18 +362,27 @@ func setupNoSyncMode(t *testing.T) *testEnv { WriteTimeout: 30 * time.Second, } broadcaster := server.NewBroadcaster(0) + repository, err := artifact.OpenRepository(t.Context(), cfg.DataDir) + require.NoError(t, err) + artifactFault := &faultInjectArtifactStore{ + ArtifactStore: repository.Content(), + } + t.Cleanup(func() { require.NoError(t, repository.Close()) }) srv := server.New( cfg, database, nil, server.WithBroadcaster(broadcaster), + server.WithArtifactStore(artifactFault), ) return &testEnv{ - srv: srv, - handler: wrapTestHandler(cfg, srv.Handler()), - db: database, - engine: nil, - broadcaster: broadcaster, - dataDir: dir, + srv: srv, + handler: wrapTestHandler(cfg, srv.Handler()), + db: database, + engine: nil, + broadcaster: broadcaster, + dataDir: dir, + artifactStore: artifactFault, + artifactFault: artifactFault, } } @@ -825,6 +935,79 @@ func TestOpenAPIEndpointDocumentsEnumsAndRequestBodies(t *testing.T) { assert.Equal(t, []string{"auto", "custom", "clipboard"}, mode.Enum) } +func TestArtifactOpenAPIDocumentsStreamingAndPagination(t *testing.T) { + te := setupArtifact(t) + w := te.get(t, "/api/openapi.json") + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + + type schema struct { + Ref string `json:"$ref"` + Type any `json:"type"` + Format string `json:"format"` + Properties map[string]schema `json:"properties"` + } + type mediaType struct { + Schema schema `json:"schema"` + } + type operation struct { + Parameters []struct { + Name string `json:"name"` + In string `json:"in"` + } `json:"parameters"` + RequestBody *struct { + Content map[string]mediaType `json:"content"` + } `json:"requestBody"` + Responses map[string]struct { + Content map[string]mediaType `json:"content"` + } `json:"responses"` + } + var spec struct { + Paths map[string]map[string]operation `json:"paths"` + Components struct { + Schemas map[string]schema `json:"schemas"` + } `json:"components"` + } + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &spec)) + + origins := spec.Paths["/api/v1/artifacts/origins"]["get"] + params := map[string]string{} + for _, param := range origins.Parameters { + params[param.Name] = param.In + } + assert.Equal(t, "query", params["cursor"]) + assert.Equal(t, "query", params["limit"]) + originsSchema := origins.Responses["200"].Content["application/json"].Schema + require.True(t, strings.HasPrefix(originsSchema.Ref, "#/components/schemas/")) + originsSchema = spec.Components.Schemas[strings.TrimPrefix( + originsSchema.Ref, "#/components/schemas/")] + assert.Contains(t, originsSchema.Properties, "next_cursor") + + path := spec.Paths["/api/v1/artifacts/{origin}/{kind}/{name}"] + getOp, ok := path["get"] + require.True(t, ok, "raw artifact GET missing from OpenAPI") + getSchema := getOp.Responses["200"].Content["application/octet-stream"].Schema + assert.Equal(t, "string", getSchema.Type) + assert.Equal(t, "binary", getSchema.Format) + assert.Contains(t, getOp.Responses, "404") + assert.Equal(t, "string", getOp.Responses["404"].Content["text/plain"].Schema.Type) + + checkpointOp, ok := spec.Paths["/api/v1/artifacts/{origin}/checkpoint"]["get"] + require.True(t, ok, "latest checkpoint GET missing from OpenAPI") + checkpointSchema := checkpointOp.Responses["200"].Content["application/octet-stream"].Schema + assert.Equal(t, "string", checkpointSchema.Type) + assert.Equal(t, "binary", checkpointSchema.Format) + + postOp, ok := path["post"] + require.True(t, ok, "raw artifact POST missing from OpenAPI") + require.NotNil(t, postOp.RequestBody) + postSchema := postOp.RequestBody.Content["application/octet-stream"].Schema + assert.Equal(t, "string", postSchema.Type) + assert.Equal(t, "binary", postSchema.Format) + assert.NotEmpty(t, postOp.Responses["200"].Content["application/json"].Schema.Ref) + assert.Contains(t, postOp.Responses, "400") + assert.Equal(t, "string", postOp.Responses["400"].Content["text/plain"].Schema.Type) +} + func TestSearchContentSemanticGETRequiresIntentHeader(t *testing.T) { te := setup(t) te.db.SetVectorSearcher(fakeTransientVectorSearcher{}) diff --git a/internal/server/session_mgmt_test.go b/internal/server/session_mgmt_test.go index ea4512d82..df27bda53 100644 --- a/internal/server/session_mgmt_test.go +++ b/internal/server/session_mgmt_test.go @@ -2,6 +2,7 @@ package server_test import ( "context" + "database/sql" "net/http" "testing" @@ -19,6 +20,10 @@ type emptyTrashHandlerResponse struct { Deleted int `json:"deleted"` } +type metadataConflictsHandlerResponse struct { + Conflicts []db.MetadataConflict `json:"conflicts"` +} + func TestSessionManagementRenameHandler(t *testing.T) { te := setup(t) te.seedSession(t, "s1", "alpha", 2) @@ -94,3 +99,40 @@ func TestSessionManagementEmptyTrashHandler(t *testing.T) { trash := decode[trashHandlerResponse](t, w) assert.Empty(t, trash.Sessions) } + +func TestSessionManagementMetadataConflictsHandler(t *testing.T) { + te := setup(t) + te.seedSession(t, "s1", "alpha", 2) + require.NoError(t, te.db.SetSyncState("artifact_origin_id", "desktop-d4e5f6")) + require.NoError(t, te.db.Update(func(tx *sql.Tx) error { + _, err := tx.Exec( + `INSERT INTO metadata_conflicts + (session_gid, field, winning_order_key, losing_order_key, + winning_origin, losing_origin, winning_op, losing_op, + winning_value, losing_value) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + "desktop-d4e5f6~s1", + "display_name", + "2026-06-14T01:02:03.000000002Z-00000000000000000000-bbb", + "2026-06-14T01:02:03.000000002Z-00000000000000000000-aaa", + "desktop-d4e5f6", + "laptop-a1b2c3", + "rename", + "rename", + `{"display_name":"Winner"}`, + `{"display_name":"Other"}`, + ) + return err + })) + + w := te.get(t, "/api/v1/sessions/s1/metadata-conflicts") + require.Equal(t, http.StatusOK, w.Code, "body: %s", w.Body.String()) + resp := decode[metadataConflictsHandlerResponse](t, w) + require.Len(t, resp.Conflicts, 1) + assert.Equal(t, "display_name", resp.Conflicts[0].Field) + assert.Equal(t, `{"display_name":"Winner"}`, resp.Conflicts[0].WinningValue) + + w = te.get(t, "/api/v1/sessions/missing/metadata-conflicts") + require.Equal(t, http.StatusNotFound, w.Code, "body: %s", w.Body.String()) + assertErrorResponse(t, w, "session not found") +} diff --git a/internal/sync/engine.go b/internal/sync/engine.go index 4074478d5..0b7054c39 100644 --- a/internal/sync/engine.go +++ b/internal/sync/engine.go @@ -2109,6 +2109,29 @@ func (e *Engine) resyncBuildLocked( return stats, err } + // Copy the artifact metadata replay tables so peer metadata + // events keep their durable LWW state across the swap. Losing + // this state would let already-applied peer events replay and + // overwrite newer local curation, so failure aborts the swap + // just like the sync state copy above. + if err := newDB.CopyMetadataReplayFrom(origPath); err != nil { + log.Printf("resync: copy metadata replay state: %v", err) + stats.Aborted = true + stats.Warnings = append(stats.Warnings, + "metadata replay copy failed, aborting swap: "+err.Error(), + ) + newDB.Close() + removeTempDB(tempPath) + restoreSkipCache() + if rerr := origDB.Reopen(); rerr != nil { + log.Printf("resync: recovery reopen: %v", rerr) + } + e.mu.Lock() + e.lastSyncStats = stats + e.mu.Unlock() + return stats, err + } + // Copy insights into newDB from the quiesced old DB file. tInsights := time.Now() reportResyncPhase( diff --git a/internal/sync/engine_integration_test.go b/internal/sync/engine_integration_test.go index 2ffa85f6e..7e874f04f 100644 --- a/internal/sync/engine_integration_test.go +++ b/internal/sync/engine_integration_test.go @@ -3716,7 +3716,9 @@ func TestSyncEngineHashSkip(t *testing.T) { different := testjsonl.NewSessionBuilder(). AddClaudeUser(tsZero, "msg2"). String() - os.WriteFile(path, []byte(different), 0o644) + require.NoError(t, os.WriteFile(path, []byte(different), 0o644), "rewrite changed session") + future := time.Unix(0, mtime).Add(time.Second) + require.NoError(t, os.Chtimes(path, future, future), "advance changed file mtime") // Third sync — mtime changed → re-synced runSyncAndAssert(t, env.engine, sync.SyncStats{TotalSessions: 1 + 0, Synced: 1, Skipped: 0}) @@ -6634,6 +6636,8 @@ func TestSyncPathsOpenCodeStorageChildUpdateAdvancesSessionMtime( `{"id":"part-a1","sessionID":"oc-storage-mtime","messageID":"msg-a1","type":"text","text":"updated reply","time":{"created":1704067201000}}`, ), 0o644) require.NoError(t, err, "rewrite part") + childMtime := time.Unix(0, initialMtime).Add(time.Second) + require.NoError(t, os.Chtimes(partPath, childMtime, childMtime), "advance part mtime") err = os.Chtimes( sessionPath, time.Unix(0, sessionMtime), @@ -7408,6 +7412,8 @@ func TestSyncPathsMiMoCodeStorageIgnoresStaleSessionSkipCache(t *testing.T) { t, sessionID, "msg-a1", "part-a1", "rewritten mimo reply", 1704067201000, ) + childMtime := sessionMtime.Add(time.Second) + require.NoError(t, os.Chtimes(partPath, childMtime, childMtime), "advance part mtime") require.NoError(t, os.Chtimes(sessionPath, sessionMtime, sessionMtime), "restore session mtime") @@ -8768,9 +8774,7 @@ func TestResyncAllPreservesTrashedSessionData(t *testing.T) { t.Fatalf("orphan health score = %v, want 94", sess.HealthScore) } qs := sess.StoredQualitySignals() - if qs == nil { - t.Fatal("orphan quality signals were not preserved") - } + require.NotNil(t, qs, "orphan quality signals were not preserved") if qs.Version != db.CurrentQualitySignalVersion || qs.ShortPromptCount != 1 || qs.MissingSuccessCriteriaCount != 1 { @@ -13289,6 +13293,67 @@ func TestResyncAllPreservesPGPushMarkerID(t *testing.T) { assert.Equal(t, "marker-123", got) } +func TestResyncAllPreservesArtifactMetadataState(t *testing.T) { + env := setupTestEnv(t) + ctx := context.Background() + + content := testjsonl.NewSessionBuilder(). + AddClaudeUser(tsEarly, "hello"). + AddClaudeAssistant(tsZeroS5, "hi"). + String() + env.writeClaudeSession(t, "proj", "sess.jsonl", content) + env.engine.SyncAll(ctx, nil) + + artifactState := map[string]string{ + "artifact_origin_id": "laptop-a1b2c3", + "artifact_metadata_hlc": "hlc-42", + "artifact_import:peer-b4c5d6:peer-b4c5d6~sess-9": "hash-imp", + "artifact_export:laptop-a1b2c3:sess.jsonl": "hash-exp", + } + for key, value := range artifactState { + require.NoError(t, env.db.SetSyncState(key, value), + "SetSyncState %s", key) + } + + projection := db.MetadataProjection{ + EventOrigin: "peer-b4c5d6", + OrderKey: "0000000001", + HLC: "hlc-1", + ArtifactHash: "hash-1", + SessionGID: "peer-b4c5d6~sess-9", + LocalSessionID: "peer-b4c5d6~sess-9", + Field: "display_name", + Op: "rename", + Value: `{"display_name":"peer name"}`, + } + _, err := env.db.RecordLocalMetadataProjection(ctx, projection) + require.NoError(t, err, "RecordLocalMetadataProjection") + + stats := env.engine.ResyncAll(ctx, nil) + require.False(t, stats.Aborted, "ResyncAll aborted: %+v", stats) + + for key, want := range artifactState { + got, err := env.db.GetSyncState(key) + require.NoError(t, err, "GetSyncState %s after resync", key) + assert.Equal(t, want, got, + "artifact sync state %s must survive resync", key) + } + + applied, err := env.db.MetadataEventApplied( + ctx, projection.EventOrigin, projection.OrderKey, + ) + require.NoError(t, err, "MetadataEventApplied after resync") + assert.True(t, applied, + "applied peer metadata event must survive resync") + + op, ok, err := env.db.MetadataReplayStateOp( + ctx, projection.SessionGID, projection.Field, + ) + require.NoError(t, err, "MetadataReplayStateOp after resync") + require.True(t, ok, "metadata replay state must survive resync") + assert.Equal(t, "rename", op) +} + func TestOpenCodeExcludedSessionsAreSkipped(t *testing.T) { env := setupSingleAgentTestEnv(t, parser.AgentOpenCode)