diff --git a/docs/specifications/agent-loop/conformance.md b/docs/specifications/agent-loop/conformance.md index 7bb3b9e..90f6c8f 100644 --- a/docs/specifications/agent-loop/conformance.md +++ b/docs/specifications/agent-loop/conformance.md @@ -38,7 +38,7 @@ | Opt-in parent-history forking per profile | MAY | [`subagents.md`](subagents.md#context-isolation-default-fresh) | | Configurable `max_concurrent_subagents` cap | MUST | [`subagents.md`](subagents.md#concurrency-limits) | | Parent turn joins on all spawned children before proceeding | MUST | [`subagents.md`](subagents.md#concurrency-limits) | -| Queryable parent→children session index | MUST | [`subagents.md`](subagents.md#session-hierarchy-bookkeeping) | +| Child lookup via a `session_meta.parent_session_id` scan (no separate index) | MUST | [`subagents.md`](subagents.md#session-hierarchy-bookkeeping) | | Structural (schema-level) depth-cap enforcement | MUST | [`subagents.md`](subagents.md#depth-limits) | | Default-deny recursion (child spawning further children) | MUST (default) | [`subagents.md`](subagents.md#tool-scoping-at-spawn) | | Cascading cancellation to in-flight children | MUST | [`subagents.md`](subagents.md#cancellation-propagation) | diff --git a/docs/specifications/agent-loop/subagents.md b/docs/specifications/agent-loop/subagents.md index 75252cf..4b787b9 100644 --- a/docs/specifications/agent-loop/subagents.md +++ b/docs/specifications/agent-loop/subagents.md @@ -33,7 +33,7 @@ See [`kernel-callbacks.md`](../kernel-callbacks.md) for `RunSession` as a kernel Fresh context per subagent, with only the prompt string crossing the boundary, is the dominant pattern across the surveyed harnesses — a strong, nearly-universal convergence. A lone surveyed outlier forks the parent's full conversation history into the child thread instead. -The kernel MUST default to fresh context: a child session's initial history contains only the `prompt` string (plus whatever its `context-assemble` hook chain independently contributes per its own profile) — the parent's message history, tool state, and in-flight context MUST NOT be implicitly visible to the child. A profile MAY opt into history forking as an explicit, named alternative (e.g. a `fork_parent_history = true` profile setting), but this MUST be opt-in, not the default, given how strongly comparable systems converge against it. Only `final_message` — the child's last message — crosses back into the parent's `tool_result`; intermediate turns are never visible to the parent's model context, though they remain queryable in the state backend via the parent/child event linkage ([Session-hierarchy bookkeeping](#session-hierarchy-bookkeeping)) for replay and audit. +The kernel MUST default to fresh context: a child session's initial history contains only the `prompt` string (plus whatever its `context-assemble` hook chain independently contributes per its own profile) — the parent's message history, tool state, and in-flight context MUST NOT be implicitly visible to the child. A profile MAY opt into history forking as an explicit, named alternative (e.g. a `fork_parent_history = true` profile setting), but this MUST be opt-in, not the default, given how strongly comparable systems converge against it. Only `final_message` — the child's last message — crosses back into the parent's `tool_result`; intermediate turns are never visible to the parent's model context, though they remain queryable in the state backend via each session's own `session_meta.parent_session_id` row ([Session-hierarchy bookkeeping](#session-hierarchy-bookkeeping)) for replay and audit. ## Concurrency limits @@ -43,7 +43,7 @@ A parent turn's apply step ([`turn-algorithm.md`](turn-algorithm.md#the-runturn- ## Session-hierarchy bookkeeping -Per the state backend's event envelope, every event carries `parent_session_id`. The kernel MUST additionally maintain a queryable parent→children index (not just per-event linkage) — a first-class, graph-queryable session-spawn structure is the kind of thing that "enables restart, audit, and cost attribution" once a harness has more than a handful of sub-agent spawns to reason about. At minimum the kernel MUST record, per spawn: parent session id, child session id, the profile name used, and the resulting depth (see [Depth limits](#depth-limits)). +Parent linkage lives in `session_meta.parent_session_id` ([`../state-backend.md#session_meta`](../state-backend.md#session_meta)), not on every event. There is no separate queryable parent→children index: finding a session's children, or reconstructing a tree, means scanning `session_meta` across the sessions directory ([`../state-backend.md#cross-session-queries`](../state-backend.md#cross-session-queries)) — a first-class, graph-queryable session-spawn structure is the kind of thing that "enables restart, audit, and cost attribution" once a harness has more than a handful of sub-agent spawns to reason about, and this scan-based approach is deliberately what fills that role here. At minimum `session_meta` MUST record, per spawn: parent session id, child session id, the profile name used, and the resulting depth (see [Depth limits](#depth-limits)). ## Depth limits @@ -64,7 +64,7 @@ Scoping child tools to the task's privilege level at spawn time is a convergent ## Cancellation propagation -If a session is cancelled (the model-provider stream cancellation described in [`provider/README.md`](../provider/README.md#transport--lifecycle), or an explicit user interrupt), the kernel MUST cascade-cancel every in-flight child `RunSession` reachable from that session via the parent→children index ([Session-hierarchy bookkeeping](#session-hierarchy-bookkeeping)), recursively through any grandchildren — cancellation MUST NOT leave orphaned child sessions running after their parent has been torn down. Each cancelled child session MUST reach `session-end` with `status = cancelled`, and its `final_message` (if any partial content exists) MUST still be recorded to the state backend for audit, even though it never reaches a waiting parent (the parent is itself being cancelled). +If a session is cancelled (the model-provider stream cancellation described in [`provider/README.md`](../provider/README.md#transport--lifecycle), or an explicit user interrupt), the kernel MUST cascade-cancel every in-flight child `RunSession` reachable from that session, recursively through any grandchildren. This is a live mechanism, not a state-backend lookup: the kernel propagates cancellation through its own in-memory tracking of currently-running child sessions (the same in-memory session state depth-budget threading and cost-rollup use, [`../state-backend.md#live-vs-post-hoc-tree-walking`](../state-backend.md#live-vs-post-hoc-tree-walking)) — the state backend plays no role in deciding *which* sessions to cancel. Cancellation MUST NOT leave orphaned child sessions running after their parent has been torn down. Each cancelled child session MUST reach `session-end` with `status = cancelled`, durably recorded to its own `session_meta` row, and its `final_message` (if any partial content exists) MUST still be recorded to the state backend for audit, even though it never reaches a waiting parent (the parent is itself being cancelled). ## Inter-session communication diff --git a/docs/specifications/agent-loop/turn-algorithm.md b/docs/specifications/agent-loop/turn-algorithm.md index d4ec4ce..985ff15 100644 --- a/docs/specifications/agent-loop/turn-algorithm.md +++ b/docs/specifications/agent-loop/turn-algorithm.md @@ -78,7 +78,7 @@ remaining_cost_budget_usd(child) = min(remaining_cost_budget_usd(parent), child_profile.max_cost_usd ?? +inf) ``` -Every `usage` event's computed `cost_usd` MUST be atomically subtracted from `remaining_cost_budget_usd` at **every** session on the path from where the spend occurred up to the root — the kernel's session-tree bookkeeping ([`subagents.md#session-hierarchy-bookkeeping`](subagents.md#session-hierarchy-bookkeeping)'s parent→children index) is what makes this walk possible. When any session's `remaining_cost_budget_usd` reaches `<= 0`, that session's bound has fired ([Limit-reached behavior](#limit-reached-behavior)), regardless of whether the spend that tripped it happened in that session directly or in a descendant. `RunSessionRequest` ([`subagents.md#data-types`](subagents.md#data-types)) carries a `remaining_cost_budget_usd` field alongside `remaining_depth`, computed identically at spawn time. +Every `usage` event's computed `cost_usd` MUST be atomically subtracted from `remaining_cost_budget_usd` at **every** session on the path from where the spend occurred up to the root — the kernel's own in-memory session-tree state — the same live tracking depth-budget threading and cancellation propagation use ([`../state-backend.md#live-vs-post-hoc-tree-walking`](../state-backend.md#live-vs-post-hoc-tree-walking)), not a state-backend read — is what makes this walk possible. When any session's `remaining_cost_budget_usd` reaches `<= 0`, that session's bound has fired ([Limit-reached behavior](#limit-reached-behavior)), regardless of whether the spend that tripped it happened in that session directly or in a descendant. `RunSessionRequest` ([`subagents.md#data-types`](subagents.md#data-types)) carries a `remaining_cost_budget_usd` field alongside `remaining_depth`, computed identically at spawn time. ### Limit-reached behavior diff --git a/docs/specifications/state-backend.md b/docs/specifications/state-backend.md index eb5474b..df183b3 100644 --- a/docs/specifications/state-backend.md +++ b/docs/specifications/state-backend.md @@ -32,7 +32,7 @@ CREATE TABLE events ( id TEXT NOT NULL UNIQUE, -- stable event identifier, independent of storage timestamp TEXT NOT NULL, -- wall-clock, display only, not ordering-authoritative kind TEXT NOT NULL, -- see "The kind enum" below for the authoritative enum - producer_category TEXT NOT NULL, + producer_category TEXT NOT NULL, -- provider | tool | context | memory | frontend | widget producer_name TEXT NOT NULL, producer_version TEXT NOT NULL, schema_version TEXT NOT NULL, diff --git a/go.mod b/go.mod index 9a99009..4262821 100644 --- a/go.mod +++ b/go.mod @@ -6,6 +6,7 @@ require ( github.com/hashicorp/go-hclog v1.6.3 github.com/hashicorp/go-plugin v1.8.0 github.com/hashicorp/hcl/v2 v2.24.0 + github.com/oklog/ulid/v2 v2.1.2 github.com/zclconf/go-cty v1.19.0 go.opentelemetry.io/contrib/bridges/otelslog v0.19.0 go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.69.0 @@ -27,6 +28,7 @@ require ( go.opentelemetry.io/otel/trace v1.44.0 google.golang.org/grpc v1.82.1 google.golang.org/protobuf v1.36.11 + modernc.org/sqlite v1.54.0 ) require ( @@ -35,6 +37,7 @@ require ( github.com/apparentlymart/go-textseg/v17 v17.0.1 // indirect github.com/cenkalti/backoff/v5 v5.0.3 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect github.com/fatih/color v1.13.0 // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect @@ -43,9 +46,11 @@ require ( github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect github.com/hashicorp/yamux v0.1.2 // indirect github.com/mattn/go-colorable v0.1.12 // indirect - github.com/mattn/go-isatty v0.0.17 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect github.com/mitchellh/go-wordwrap v1.0.1 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect github.com/oklog/run v1.1.0 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.44.0 // indirect go.opentelemetry.io/proto/otlp v1.10.0 // indirect @@ -58,6 +63,9 @@ require ( google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect google.golang.org/grpc/cmd/protoc-gen-go-grpc v1.6.2 // indirect + modernc.org/libc v1.74.1 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect ) tool ( diff --git a/go.sum b/go.sum index dc1f869..474305f 100644 --- a/go.sum +++ b/go.sum @@ -13,6 +13,8 @@ github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XL github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/fatih/color v1.13.0 h1:8LOYc1KYPPmyKMuN8QV2DNRWNbLo6LZ0iLs8+mlH53w= github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A= @@ -26,6 +28,8 @@ github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk= @@ -34,6 +38,8 @@ github.com/hashicorp/go-hclog v1.6.3 h1:Qr2kF+eVWjTiYmU7Y31tYlP1h0q/X3Nl3tPGdaB1 github.com/hashicorp/go-hclog v1.6.3/go.mod h1:W4Qnvbt70Wk/zYJryRzDRU/4r0kIg0PVHBcfoyhpF5M= github.com/hashicorp/go-plugin v1.8.0 h1:ie8S6RRY8RvB2usYZv+AAZ/wBvx2AU5p5QeP5j/FORs= github.com/hashicorp/go-plugin v1.8.0/go.mod h1:BExt6KEaIYx804z8k4gRzRLEvxKVb+kn0NMcihqOqb8= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/hashicorp/hcl/v2 v2.24.0 h1:2QJdZ454DSsYGoaE6QheQZjtKZSUs9Nh2izTWiwQxvE= github.com/hashicorp/hcl/v2 v2.24.0/go.mod h1:oGoO1FIQYfn/AgyOhlg9qLC6/nOJPX3qGbkZpYAcqfM= github.com/hashicorp/yamux v0.1.2 h1:XtB8kyFOyHXYVFnwT5C3+Bdo8gArse7j2AQ0DA0Uey8= @@ -45,14 +51,21 @@ github.com/mattn/go-colorable v0.1.12 h1:jF+Du6AlPIjs2BiUiQlKOX0rt3SujHxPnksPKZb github.com/mattn/go-colorable v0.1.12/go.mod h1:u5H1YNBxpqRaxsYJYSkiCWKzEfiAb1Gb520KVy5xxl4= github.com/mattn/go-isatty v0.0.12/go.mod h1:cbi8OIDigv2wuxKPP5vlRcQ1OAZbq2CE4Kysco4FUpU= github.com/mattn/go-isatty v0.0.14/go.mod h1:7GGIvUiUoEMVVmxf/4nioHXj79iQHKdU27kJ6hsGG94= -github.com/mattn/go-isatty v0.0.17 h1:BTarxUcIeDqL27Mc+vyvdWYSL28zpIhv3RoTdsLMPng= -github.com/mattn/go-isatty v0.0.17/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mitchellh/go-wordwrap v1.0.1 h1:TLuKupo69TCn6TQSyGxwI1EblZZEsQ0vMlAFQflz0v0= github.com/mitchellh/go-wordwrap v1.0.1/go.mod h1:R62XHJLzvMFRBbcrT7m7WgmE1eOyTSsCt+hzestvNj0= +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/run v1.1.0 h1:GEenZ1cK0+q0+wsJew9qUg/DyD8k3JzYsZAi5gYi2mA= github.com/oklog/run v1.1.0/go.mod h1:sVPdnTZT1zYwAJeCMu2Th4T21pA3FPOQRfWjQlk7DVU= +github.com/oklog/ulid/v2 v2.1.2 h1:IEclFb9JNvzYA6MW2SCxbLzcHTVsfqm3PrqGQJH5zec= +github.com/oklog/ulid/v2 v2.1.2/go.mod h1:rcEKHmBBKfef9DhnvX7y1HZBYxjXb0cP5ExxNsTT1QQ= +github.com/pborman/getopt v0.0.0-20170112200414-7148bc3a4c30/go.mod h1:85jBQOZwpVEaDAr341tbn15RS4fCAsIst0qp7i8ex1o= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +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/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/testify v1.7.2/go.mod h1:R6va5+xMeoiuVRoj+gSkQ7d3FALtqAAGI1FQKckRals= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= @@ -120,7 +133,7 @@ golang.org/x/sys v0.0.0-20200223170610-d5e6a3e2c0ae/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220503163025-988cb79eb6c6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= @@ -142,3 +155,31 @@ google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +modernc.org/cc/v4 v4.29.0 h1:CXgwL8cvxmyzBQZzbSl/6xFtMCryb6u8IOqDci39cgc= +modernc.org/cc/v4 v4.29.0/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI= +modernc.org/ccgo/v4 v4.34.6 h1:sBgfIwyN0TQ9C5hwIeuqyeAKyMWnbvj2fvpF4L11uzU= +modernc.org/ccgo/v4 v4.34.6/go.mod h1:SZ8YcN9NG7XVsQYdm6jYBvi8PQP1qi+kqB6OhjqI3Fk= +modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM= +modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU= +modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= +modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= +modernc.org/gc/v3 v3.1.4 h1:2g65LGVSmFQrXeITAw97x7hCRvZFcyE1uDP+7Vng7JI= +modernc.org/gc/v3 v3.1.4/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= +modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= +modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= +modernc.org/libc v1.74.1 h1:bdR4VTKFMC4966QSNZ05XLGI/VwzVa2kTUX51Dm0riQ= +modernc.org/libc v1.74.1/go.mod h1:uH4t5bOx3G3g9Xcmj10YKlTcVISlRDwv8VoQJG9n8Os= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg= +modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= +modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= +modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= +modernc.org/sqlite v1.54.0 h1:JCxR4qwkJvOaqAoYcgDoO25Nc+ROg6EJ2LfBVzdrgog= +modernc.org/sqlite v1.54.0/go.mod h1:4ntCLuNmnH8+GNqjka1wNg7KJd5/Hi5FYp8K+XQ7GZw= +modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= +modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= +modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= +modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM= diff --git a/internal/statebackend/CLAUDE.md b/internal/statebackend/CLAUDE.md new file mode 100644 index 0000000..3d4996e --- /dev/null +++ b/internal/statebackend/CLAUDE.md @@ -0,0 +1,8 @@ +# internal/statebackend — agent notes + +- **`schema.go`'s `schemaStatements` are the spec's DDL, verbatim** — column names, types, constraints, and index definitions must match [`docs/specifications/state-backend.md#schema`](../../docs/specifications/state-backend.md#schema) exactly, including `AUTOINCREMENT` on every `sequence` column. A schema change starts in the spec, not here; this file follows. +- **`events` (and therefore `cost_ledger`/`plan_items`/`producers`, which all reference it) is append-only by design — there is no delete or update path anywhere in this package, and none should be added.** The spec's retention model is whole-file pruning by an explicit operator action, never row-level deletion (`docs/specifications/state-backend.md#retention--pruning`). Don't add a "delete event" or "prune old rows" method here — that's a future CLI-level, whole-file concern, not something this package's API should expose. +- **`sequence` is the only ordering authority, everywhere** — `Events`, `CostLedger`, and `PlanItems` all `ORDER BY sequence`; never add an `ORDER BY timestamp` or compare `Event.Timestamp` values to decide ordering (`.claude/rules/determinism.md`). `timestamp` columns exist for display only. +- **`Store.List`/`Store.Children`/recovery's `recoverTable` all bypass `Open`'s migration path on purpose** — they read `session_meta` (or salvage rows) directly via a raw `openDB`, never through `Open`. A metadata scan or a recovery pass must not have the side effect of silently migrating every file it touches. +- **`mapAppendEventError` is the only place a raw `*sqlite.Error` is inspected** (translating a `UNIQUE` violation on `events.id` into `ErrDuplicateEventID`) — event IDs are caller-supplied, and this package enforces their uniqueness purely via the DB constraint rather than a pre-check query. Don't add a `SELECT ... WHERE id = ?` existence check ahead of the insert; it would be redundant and racy against the single-writer invariant this package already gives you for free. +- **`recoveryTableSpecs`' `columns` always includes the table's own primary key** so salvaged rows keep their original `sequence`/identity rather than being renumbered on reinsert — foreign keys from `cost_ledger`/`plan_items`/`producers` into `events.sequence` depend on this. diff --git a/internal/statebackend/README.md b/internal/statebackend/README.md new file mode 100644 index 0000000..68e4eab --- /dev/null +++ b/internal/statebackend/README.md @@ -0,0 +1,43 @@ +# internal/statebackend + +The kernel state backend from [`docs/specifications/state-backend.md`](../../docs/specifications/state-backend.md) — sqlite-per-session persistence for durability, replay, and audit. Unlike the plugin categories, this is a kernel foundation: not pluggable, not configurable, one sqlite file per session at `$XDG_STATE_HOME/agent/sessions/.sqlite`. + +## What this package does + +- `statebackend.go` — `Store`: manages the sessions directory. `Create` makes a new session file (schema applied, `PRAGMA user_version` stamped, initial `session_meta` row inserted) and returns an open `Session`. `Open` opens an existing one, running corruption recovery and schema migration first. `List`/`Children` scan `session_meta` across every file for cross-session queries — there is no separate index. +- `schema.go` — the five-table DDL, reproduced verbatim from the spec, plus the ordered `migrationStep` slice `Open` walks when a file's `user_version` is older than `currentSchemaVersion`. +- `event.go` — `Event`, `CostEntry`, `PlanItem`: the Go mirrors of the `events`, `cost_ledger`, and `plan_items` columns, plus the encode/decode helpers between their proto enum fields and the lowercase TEXT vocabulary the spec stores on disk. +- `session.go` — `Session`'s write path: `AppendEvent`, `AppendMessage`, `AppendPlan`, `SetStatus`, `Close`. The kernel is this file's sole writer; every append runs in one transaction so an event row never exists without its accompanying `cost_ledger`/`plan_items`/`producers` rows. +- `query.go` — `Session`'s read path: `Meta`, `Events` (a sequence-ordered `iter.Seq2[Event, error]` — replay's entry point), `Producers`, `TotalCostUSD`, `CostLedger`, `PlanItems`. +- `integrity.go` — `PRAGMA integrity_check` on every `Open`, and the salvage recovery path when it fails. +- `sessionid.go` — `NewSessionID`/`ValidateSessionID`: session IDs are canonical uppercase ULIDs, sortable chronologically by filename alone. + +## Public API sketch + +```go +store, err := statebackend.NewStore(dir, statebackend.WithLogger(logger), statebackend.WithTelemetry(prov)) + +sess, err := store.Create(ctx, statebackend.SessionMeta{SessionID: statebackend.NewSessionID(time.Now()), Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING}) +seq, err := sess.AppendEvent(ctx, statebackend.Event{ID: "...", Kind: kernelv1.EventKind_EVENT_KIND_TOOL_CALL, Producer: ref, Payload: b}) +err = sess.SetStatus(ctx, sessionv1.SessionStatus_SESSION_STATUS_COMPLETED, &endedAt) +err = sess.Close() + +sess, err = store.Open(ctx, sessionID) +for ev, err := range sess.Events(ctx) { ... } // sequence-ordered replay +meta, err := sess.Meta(ctx) +metas, err := store.Children(ctx, sessionID) +``` + +`Session` write methods return `ErrClosed` once `Close` has been called. `Store.Open`/`Store.Create` reject a malformed or non-canonical session ID via `ValidateSessionID`. + +## Recovery behavior + +On `Open`, `checkIntegrity` runs `PRAGMA integrity_check`. If the file can't be opened at all, or the check reports problems, the damaged original is renamed to `.sqlite.corrupt` — never deleted — and salvage begins: a fresh, schema-correct file is built table-by-table from the damaged one, tolerating per-row scan/insert failures (a bad row is skipped, not fatal to the rest of the table). Recovery only counts as successful once the single `session_meta` row is itself salvaged; every other table is best-effort on top of that. A successful recovery installs the new file at the original path and returns an open handle to it, after logging a `WARN` with per-table salvaged/skipped counts — recovery is never silent. If `session_meta` itself can't be salvaged, `Open` returns `ErrUnrecoverable` and the `.corrupt` file is left in place for manual inspection. + +## Testing notes + +- Unit tests are the default tier (`go-testing.md`) — in-memory sqlite files under `t.TempDir()`, no external fixtures. +- Concurrency-sensitive paths (every write method, corruption recovery) run under `go test -race`, per `.claude/rules/go-testing.md`'s hard requirement for anything touching the state backend. +- `event_fuzz_test.go`'s `FuzzEventRoundTrip` and `sessionid_fuzz_test.go`'s `FuzzValidateSessionID` are the two fuzz targets — event append/scan round-tripping and ULID validation, respectively. +- `integrity_test.go` covers the corruption-recovery path directly: deliberately truncated/corrupted files, partial-table salvage, and the `ErrUnrecoverable` case. +- Replay-adjacent assertions compare `sequence`, never wall-clock time, per `.claude/rules/determinism.md`. diff --git a/internal/statebackend/doc.go b/internal/statebackend/doc.go new file mode 100644 index 0000000..da1f03d --- /dev/null +++ b/internal/statebackend/doc.go @@ -0,0 +1,65 @@ +// Package statebackend implements the kernel state backend specified in +// docs/specifications/state-backend.md: sqlite-per-session persistence for +// everything the kernel needs to durably reconstruct a session after the +// fact. Each session gets its own sqlite file holding an append-only event +// log (events), a single mutable summary row (session_meta), a structured +// spend ledger (cost_ledger), and a plan/apply audit trail (plan_items) — +// see docs/specifications/state-backend.md#schema for the full five-table +// shape (the fifth table, producers, is the distinct-plugin manifest +// AppendEvent/AppendMessage/AppendPlan maintain alongside events). +// +// Store (statebackend.go) manages the directory of per-session files +// (docs/specifications/state-backend.md#file-layout); Session is an open +// handle to one file, with append methods in session.go — the kernel's +// sole-writer path (AppendEvent, AppendMessage, AppendPlan, SetStatus) — +// and read methods in query.go (Meta, Events, Producers, CostLedger, +// PlanItems). Replay reconstructs a session by iterating Events in +// sequence order: sequence is the only ordering authority anywhere in +// this package (docs/specifications/state-backend.md#ordering--concurrency, +// .claude/rules/determinism.md), never wall-clock time. Corruption +// recovery (integrity.go) and schema migration (schema.go) both run +// inside Store.Open, before any other operation touches a file, per +// docs/specifications/state-backend.md#schema-migration--corruption-recovery. +// +// # Design decisions +// +// The spec leaves some choices open; this package makes them as follows, +// coordinator-approved: +// +// - Every TEXT timestamp column is stored as RFC 3339, UTC, millisecond +// precision (formatTimestamp/parseTimestamp, statebackend.go). This is +// display-only, exactly as the spec requires — sequence remains the +// sole ordering authority; no code path in this package compares +// timestamps to order events. +// - session_meta.depth stores the remaining_depth value carried on the +// RunSessionRequest at spawn time +// (docs/specifications/agent-loop/subagents.md#depth-limits), not a +// freshly re-derived value — a cache that avoids a parent-chain walk +// when scanning. +// - cost_ledger and plan_items declare their event_sequence foreign key +// with no ON DELETE clause: events is append-only and its rows are +// never deleted, so no delete-cascade behavior is needed or defined. +// - Event.ID is caller-supplied, not generated by this package; +// uniqueness within a session's file is enforced here (a UNIQUE +// constraint on events.id, surfaced to callers as +// ErrDuplicateEventID). +// - Session.Close runs PRAGMA wal_checkpoint(TRUNCATE), folding the WAL +// back into the main file rather than leaving a -wal sidecar behind. +// - The Store directory is created mode 0700, and each session file is +// created mode 0600 — session state is treated as sensitive by +// default, readable only by the user running the kernel. Enforced on +// Unix-like systems only; Windows ACLs are out of scope. +// - cost_ledger.cost_usd stores exactly the cost_usd value the caller +// computed (the model provider protocol's own cost computation, +// docs/specifications/provider/protocol.md#cost-computation); this +// package never recomputes a cost or token figure itself +// (.claude/rules/determinism.md). +// - events.producer_category and producers.category store the lowercase +// plugin-category vocabulary: provider, tool, context, memory, +// frontend, widget — the same names docs/specifications/ uses as its +// own per-category directory names. +// - Store.List and Store.Children read each file's session_meta row +// directly (scan, statebackend.go), bypassing Open's schema-version +// check and migration path — a metadata scan MUST NOT have the side +// effect of migrating every file it merely lists. +package statebackend diff --git a/internal/statebackend/errors.go b/internal/statebackend/errors.go new file mode 100644 index 0000000..1367ea0 --- /dev/null +++ b/internal/statebackend/errors.go @@ -0,0 +1,24 @@ +package statebackend + +import "errors" + +// ErrSchemaTooNew is returned when a session file's schema version is newer than this kernel understands. +var ErrSchemaTooNew = errors.New("statebackend: session file schema is newer than this kernel supports") + +// ErrNotFound is returned when a requested resource does not exist. +var ErrNotFound = errors.New("statebackend: not found") + +// ErrDuplicateEventID is returned when an event with the same ID is attempted to be inserted. +var ErrDuplicateEventID = errors.New("statebackend: duplicate event id") + +// ErrInvalidKind is returned when an event's kind is invalid or unspecified. +var ErrInvalidKind = errors.New("statebackend: invalid event kind") + +// ErrInvalidDecision is returned when a plan item's decision is invalid or unspecified. +var ErrInvalidDecision = errors.New("statebackend: invalid plan decision") + +// ErrUnrecoverable is returned when a session file is corrupted and recovery failed. +var ErrUnrecoverable = errors.New("statebackend: session file unrecoverable") + +// ErrClosed is returned when an operation is attempted on a closed session store. +var ErrClosed = errors.New("statebackend: closed") diff --git a/internal/statebackend/event.go b/internal/statebackend/event.go new file mode 100644 index 0000000..66bada5 --- /dev/null +++ b/internal/statebackend/event.go @@ -0,0 +1,214 @@ +package statebackend + +import ( + "fmt" + "time" + + commonv1 "github.com/pluggableharness/agent/pkg/common/proto/v1" + kernelv1 "github.com/pluggableharness/agent/pkg/kernel/proto/v1" + planv1 "github.com/pluggableharness/agent/pkg/plan/proto/v1" +) + +// Event mirrors the events table's columns +// (docs/specifications/state-backend.md#events) — the kernel's event +// envelope, verbatim, as an append-only row. Sequence is assigned by +// AppendEvent/AppendMessage/AppendPlan (session.go) and is meaningless on +// an Event passed into one of those calls; it's populated on an Event read +// back from storage (Stage 3's query.go). +type Event struct { + // Sequence is the row's assigned INTEGER PRIMARY KEY AUTOINCREMENT + // value. Ignored on append; only meaningful on a read-back Event. + Sequence int64 + // ID is the stable event identifier, independent of storage — UNIQUE + // within a session's file; a repeat on append returns + // ErrDuplicateEventID. + ID string + // Timestamp is wall-clock, display-only, never ordering-authoritative + // (docs/specifications/state-backend.md#ordering--concurrency). + Timestamp time.Time + // Kind identifies the event envelope's payload shape + // (docs/specifications/state-backend.md#the-kind-enum). + // EVENT_KIND_UNSPECIFIED is rejected on append with ErrInvalidKind. + Kind kernelv1.EventKind + // Producer identifies which plugin produced this event. Required + // (producer_category/name/version are all NOT NULL columns). + Producer *commonv1.ProducerRef + // SchemaVersion is the producer's payload schema version. + SchemaVersion string + // Payload is the opaque event body — the kernel never inspects this + // (docs/specifications/state-backend.md#events). + Payload []byte +} + +// CostEntry mirrors the cost_ledger table's columns +// (docs/specifications/state-backend.md#cost_ledger) — structured spend, +// appended once per completed model turn alongside its message event +// (AppendMessage, session.go). CostUSD and the token counters are stored +// exactly as the caller computed them; this package never recomputes a +// cost/token figure itself (determinism.md). +type CostEntry struct { + ProviderName string + ModelID string + InputTokens int64 + OutputTokens int64 + CacheWriteTokens int64 + CacheReadTokens int64 + CostUSD float64 +} + +// PlanItem mirrors the plan_items table's columns +// (docs/specifications/state-backend.md#plan_items) — one row per plan +// item, appended alongside its plan event (AppendPlan, session.go). +type PlanItem struct { + TurnID string + ToolCallID string + ProviderName string + ToolName string + // Decision is one of PLAN_DECISION_ALLOW/ASK/DENY. + // PLAN_DECISION_UNSPECIFIED and PLAN_DECISION_PENDING are rejected on + // append with ErrInvalidDecision — the spec's decision column only + // ever holds a made decision ("allow | ask | deny"), never a pending + // one (docs/specifications/state-backend.md#plan_items). + Decision planv1.PlanDecision + DecidedBy string +} + +// eventKindText maps EventKind to the exact lowercase snake_case text +// docs/specifications/state-backend.md#the-kind-enum documents — the wire +// enum's own SCREAMING_SNAKE_CASE String() is not what gets stored. +// EVENT_KIND_UNSPECIFIED is deliberately absent: like SessionStatus's zero +// value (statebackend.go), it MUST NOT ever be persisted. +var eventKindText = map[kernelv1.EventKind]string{ + kernelv1.EventKind_EVENT_KIND_MESSAGE: "message", + kernelv1.EventKind_EVENT_KIND_TOOL_CALL: "tool_call", + kernelv1.EventKind_EVENT_KIND_TOOL_RESULT: "tool_result", + kernelv1.EventKind_EVENT_KIND_PLAN: "plan", + kernelv1.EventKind_EVENT_KIND_APPLY: "apply", + kernelv1.EventKind_EVENT_KIND_CONTEXT_CONTRIBUTION: "context_contribution", + kernelv1.EventKind_EVENT_KIND_MEMORY_WRITE: "memory_write", + kernelv1.EventKind_EVENT_KIND_MEMORY_UPDATE: "memory_update", + kernelv1.EventKind_EVENT_KIND_MEMORY_DELETE: "memory_delete", +} + +// eventTextKind is eventKindText inverted, built once from eventKindText +// itself so the two can never drift. +var eventTextKind = func() map[string]kernelv1.EventKind { + m := make(map[string]kernelv1.EventKind, len(eventKindText)) + for kind, text := range eventKindText { + m[text] = kind + } + return m +}() + +// encodeEventKind renders kind as its stored TEXT representation. +// EVENT_KIND_UNSPECIFIED and any unrecognized value return ErrInvalidKind. +func encodeEventKind(kind kernelv1.EventKind) (string, error) { + text, ok := eventKindText[kind] + if !ok { + return "", fmt.Errorf("statebackend: %w: %v", ErrInvalidKind, kind) + } + return text, nil +} + +// decodeEventKind is the inverse of encodeEventKind, used when reading an +// events row back (Stage 3's query.go). +func decodeEventKind(text string) (kernelv1.EventKind, error) { + kind, ok := eventTextKind[text] + if !ok { + return kernelv1.EventKind_EVENT_KIND_UNSPECIFIED, fmt.Errorf("statebackend: %w: %q", ErrInvalidKind, text) + } + return kind, nil +} + +// producerCategoryText maps a plugin Category to the lowercase text this +// package stores in events.producer_category and producers.category. +// state-backend.md's DDL leaves the column's exact text vocabulary +// undocumented (unlike session_meta.status and events.kind, which the spec +// enumerates literally) — this uses the same lowercase category names the +// specifications/ tree itself uses as directory names (provider/, tool/, +// context/, memory/, frontend/, widget/), for consistency with every other +// lowercase-text enum this package stores. CATEGORY_UNSPECIFIED is +// deliberately absent — a producer's category MUST NOT ever be +// unspecified (kernel-callbacks.md's server-derived producer identity is +// always a real category). +var producerCategoryText = map[commonv1.Category]string{ + commonv1.Category_CATEGORY_PROVIDER: "provider", + commonv1.Category_CATEGORY_TOOL: "tool", + commonv1.Category_CATEGORY_CONTEXT: "context", + commonv1.Category_CATEGORY_MEMORY: "memory", + commonv1.Category_CATEGORY_FRONTEND: "frontend", + commonv1.Category_CATEGORY_WIDGET: "widget", +} + +// producerTextCategory is producerCategoryText inverted, built once from +// producerCategoryText itself so the two can never drift. +var producerTextCategory = func() map[string]commonv1.Category { + m := make(map[string]commonv1.Category, len(producerCategoryText)) + for category, text := range producerCategoryText { + m[text] = category + } + return m +}() + +// encodeProducerCategory renders category as its stored TEXT +// representation. CATEGORY_UNSPECIFIED and any unrecognized value are +// rejected. +func encodeProducerCategory(category commonv1.Category) (string, error) { + text, ok := producerCategoryText[category] + if !ok { + return "", fmt.Errorf("statebackend: producer category %v has no stored representation", category) + } + return text, nil +} + +// decodeProducerCategory is the inverse of encodeProducerCategory, used +// when reading an events or producers row back (Stage 3's query.go). +func decodeProducerCategory(text string) (commonv1.Category, error) { + category, ok := producerTextCategory[text] + if !ok { + return commonv1.Category_CATEGORY_UNSPECIFIED, fmt.Errorf("statebackend: unrecognized producer category %q", text) + } + return category, nil +} + +// planDecisionText maps a PlanDecision to the exact lowercase text +// docs/specifications/state-backend.md#plan_items documents +// ("allow | ask | deny"). PLAN_DECISION_UNSPECIFIED and +// PLAN_DECISION_PENDING are deliberately absent — the plan_items table +// only ever holds a made decision. +var planDecisionText = map[planv1.PlanDecision]string{ + planv1.PlanDecision_PLAN_DECISION_ALLOW: "allow", + planv1.PlanDecision_PLAN_DECISION_ASK: "ask", + planv1.PlanDecision_PLAN_DECISION_DENY: "deny", +} + +// planTextDecision is planDecisionText inverted, built once from +// planDecisionText itself so the two can never drift. +var planTextDecision = func() map[string]planv1.PlanDecision { + m := make(map[string]planv1.PlanDecision, len(planDecisionText)) + for decision, text := range planDecisionText { + m[text] = decision + } + return m +}() + +// encodePlanDecision renders decision as its stored TEXT representation. +// PLAN_DECISION_UNSPECIFIED, PLAN_DECISION_PENDING, and any unrecognized +// value return ErrInvalidDecision. +func encodePlanDecision(decision planv1.PlanDecision) (string, error) { + text, ok := planDecisionText[decision] + if !ok { + return "", fmt.Errorf("statebackend: %w: %v", ErrInvalidDecision, decision) + } + return text, nil +} + +// decodePlanDecision is the inverse of encodePlanDecision, used when +// reading a plan_items row back (Stage 3's query.go). +func decodePlanDecision(text string) (planv1.PlanDecision, error) { + decision, ok := planTextDecision[text] + if !ok { + return planv1.PlanDecision_PLAN_DECISION_UNSPECIFIED, fmt.Errorf("statebackend: %w: %q", ErrInvalidDecision, text) + } + return decision, nil +} diff --git a/internal/statebackend/event_fuzz_test.go b/internal/statebackend/event_fuzz_test.go new file mode 100644 index 0000000..a3fcf38 --- /dev/null +++ b/internal/statebackend/event_fuzz_test.go @@ -0,0 +1,109 @@ +package statebackend + +import ( + "bytes" + "context" + "sort" + "testing" + "time" + + kernelv1 "github.com/pluggableharness/agent/pkg/kernel/proto/v1" +) + +// fuzzValidEventKinds is every EventKind that encodeEventKind accepts, +// built once (outside FuzzEventRoundTrip's Fuzz closure) in a stable, +// sorted order so a given kindByte input maps to the same kind across +// runs — reproducibility matters for a saved fuzz failure to replay +// identically. +var fuzzValidEventKinds = func() []kernelv1.EventKind { + kinds := make([]kernelv1.EventKind, 0, len(eventKindText)) + for k := range eventKindText { + kinds = append(kinds, k) + } + sort.Slice(kinds, func(i, j int) bool { return kinds[i] < kinds[j] }) + return kinds +}() + +// FuzzEventRoundTrip appends an arbitrary, always-valid-shaped Event to a +// fresh session and reads it back via Events(), asserting field equality +// and byte-identical payload. kindByte is mapped modulo into +// fuzzValidEventKinds so every fuzzed input produces a kind AppendEvent +// accepts — this fuzzes the append/scan/decode round trip itself, not +// EventKind validation (event_test.go's TestEncodeEventKind_unspecifiedRejected +// already covers the rejection path). Must never panic, per go-testing.md's +// unit-tier speed budget this creates a fresh t.TempDir() Store per +// iteration to keep each session's event IDs collision-free without +// tracking state across iterations. +func FuzzEventRoundTrip(f *testing.F) { + f.Add("evt-1", byte(0), []byte("hello"), "v1") + f.Add("", byte(1), []byte{}, "") + f.Add("evt-unicode-🎉", byte(255), []byte{0x00, 0xFF, 0x7F}, "schema-v2") + f.Add("evt-nul\x00id", byte(37), bytes.Repeat([]byte{0xAA}, 4096), "s") + + f.Fuzz(func(t *testing.T, id string, kindByte byte, payload []byte, schemaVersion string) { + if payload == nil { + // A nil []byte binds as SQL NULL, which the payload BLOB NOT + // NULL column rejects — that's a real, correct constraint + // enforcement, not a round-trip bug, so it's out of scope for + // this fuzz target. Normalize instead of asserting either way. + payload = []byte{} + } + kind := fuzzValidEventKinds[int(kindByte)%len(fuzzValidEventKinds)] + + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + wantTimestamp := time.Now() + ev := Event{ + ID: id, + Timestamp: wantTimestamp, + Kind: kind, + Producer: testProducer(), + SchemaVersion: schemaVersion, + Payload: payload, + } + + seq, err := sess.AppendEvent(context.Background(), ev) + if err != nil { + t.Fatalf("AppendEvent: %v", err) + } + + var got []Event + for e, err := range sess.Events(context.Background()) { + if err != nil { + t.Fatalf("Events: %v", err) + } + got = append(got, e) + } + if len(got) != 1 { + t.Fatalf("Events returned %d events, want 1", len(got)) + } + readBack := got[0] + + if readBack.Sequence != seq { + t.Fatalf("Sequence = %d, want %d", readBack.Sequence, seq) + } + if readBack.ID != id { + t.Fatalf("ID = %q, want %q", readBack.ID, id) + } + if readBack.Kind != kind { + t.Fatalf("Kind = %v, want %v", readBack.Kind, kind) + } + if readBack.SchemaVersion != schemaVersion { + t.Fatalf("SchemaVersion = %q, want %q", readBack.SchemaVersion, schemaVersion) + } + if !bytes.Equal(readBack.Payload, payload) { + t.Fatalf("Payload mismatch: got %d bytes, want %d bytes", len(readBack.Payload), len(payload)) + } + wantTimestampTrunc := wantTimestamp.UTC().Truncate(time.Millisecond) + if !readBack.Timestamp.Equal(wantTimestampTrunc) { + t.Fatalf("Timestamp = %v, want %v (millisecond-truncated)", readBack.Timestamp, wantTimestampTrunc) + } + if readBack.Producer == nil || + readBack.Producer.GetCategory() != ev.Producer.GetCategory() || + readBack.Producer.GetName() != ev.Producer.GetName() || + readBack.Producer.GetVersion() != ev.Producer.GetVersion() { + t.Fatalf("Producer = %+v, want %+v", readBack.Producer, ev.Producer) + } + }) +} diff --git a/internal/statebackend/event_test.go b/internal/statebackend/event_test.go new file mode 100644 index 0000000..3a9e4f7 --- /dev/null +++ b/internal/statebackend/event_test.go @@ -0,0 +1,117 @@ +package statebackend + +import ( + "errors" + "testing" + + commonv1 "github.com/pluggableharness/agent/pkg/common/proto/v1" + kernelv1 "github.com/pluggableharness/agent/pkg/kernel/proto/v1" + planv1 "github.com/pluggableharness/agent/pkg/plan/proto/v1" +) + +func TestEventKind_roundTrip(t *testing.T) { + t.Parallel() + + for kind := range eventKindText { + text, err := encodeEventKind(kind) + if err != nil { + t.Fatalf("encodeEventKind(%v): %v", kind, err) + } + got, err := decodeEventKind(text) + if err != nil { + t.Fatalf("decodeEventKind(%q): %v", text, err) + } + if got != kind { + t.Errorf("round trip %v -> %q -> %v, want %v", kind, text, got, kind) + } + } +} + +func TestEncodeEventKind_unspecifiedRejected(t *testing.T) { + t.Parallel() + if _, err := encodeEventKind(kernelv1.EventKind_EVENT_KIND_UNSPECIFIED); !errors.Is(err, ErrInvalidKind) { + t.Fatalf("encodeEventKind(UNSPECIFIED) err = %v, want ErrInvalidKind", err) + } +} + +func TestDecodeEventKind_unrecognized(t *testing.T) { + t.Parallel() + if _, err := decodeEventKind("not_a_kind"); !errors.Is(err, ErrInvalidKind) { + t.Fatalf("decodeEventKind(garbage) err = %v, want ErrInvalidKind", err) + } +} + +func TestProducerCategory_roundTrip(t *testing.T) { + t.Parallel() + + for category := range producerCategoryText { + text, err := encodeProducerCategory(category) + if err != nil { + t.Fatalf("encodeProducerCategory(%v): %v", category, err) + } + got, err := decodeProducerCategory(text) + if err != nil { + t.Fatalf("decodeProducerCategory(%q): %v", text, err) + } + if got != category { + t.Errorf("round trip %v -> %q -> %v, want %v", category, text, got, category) + } + } +} + +func TestEncodeProducerCategory_unspecifiedRejected(t *testing.T) { + t.Parallel() + if _, err := encodeProducerCategory(commonv1.Category_CATEGORY_UNSPECIFIED); err == nil { + t.Fatal("encodeProducerCategory(UNSPECIFIED) = nil error, want error") + } +} + +func TestDecodeProducerCategory_unrecognized(t *testing.T) { + t.Parallel() + if _, err := decodeProducerCategory("not_a_category"); err == nil { + t.Fatal("decodeProducerCategory(garbage) = nil error, want error") + } +} + +func TestPlanDecision_roundTrip(t *testing.T) { + t.Parallel() + + for decision := range planDecisionText { + text, err := encodePlanDecision(decision) + if err != nil { + t.Fatalf("encodePlanDecision(%v): %v", decision, err) + } + got, err := decodePlanDecision(text) + if err != nil { + t.Fatalf("decodePlanDecision(%q): %v", text, err) + } + if got != decision { + t.Errorf("round trip %v -> %q -> %v, want %v", decision, text, got, decision) + } + } +} + +func TestEncodePlanDecision_unspecifiedRejected(t *testing.T) { + t.Parallel() + if _, err := encodePlanDecision(planv1.PlanDecision_PLAN_DECISION_UNSPECIFIED); !errors.Is(err, ErrInvalidDecision) { + t.Fatalf("encodePlanDecision(UNSPECIFIED) err = %v, want ErrInvalidDecision", err) + } +} + +func TestEncodePlanDecision_pendingRejected(t *testing.T) { + t.Parallel() + // PENDING is a real, valid PlanDecision value elsewhere in the system + // (a plan item awaiting a decision) but the plan_items table only ever + // holds a *made* decision — state-backend.md's decision column is + // documented as "allow | ask | deny" with no "pending" value. + if _, err := encodePlanDecision(planv1.PlanDecision_PLAN_DECISION_PENDING); !errors.Is(err, ErrInvalidDecision) { + t.Fatalf("encodePlanDecision(PENDING) err = %v, want ErrInvalidDecision", err) + } +} + +func TestDecodePlanDecision_unrecognized(t *testing.T) { + t.Parallel() + if _, err := decodePlanDecision("not_a_decision"); !errors.Is(err, ErrInvalidDecision) { + t.Fatalf("decodePlanDecision(garbage) err = %v, want ErrInvalidDecision", err) + } +} diff --git a/internal/statebackend/integrity.go b/internal/statebackend/integrity.go new file mode 100644 index 0000000..687982d --- /dev/null +++ b/internal/statebackend/integrity.go @@ -0,0 +1,302 @@ +package statebackend + +import ( + "context" + "database/sql" + "fmt" + "os" + "strings" + + "github.com/pluggableharness/agent/internal/telemetry" +) + +// recoveryStats records, per table, how many rows a recoverSession call +// salvaged versus skipped — logged as part of the "recovery MUST NOT +// happen silently" requirement +// (docs/specifications/state-backend.md#corruption-recovery). +type recoveryStats struct { + salvaged map[string]int + skipped map[string]int +} + +// recoveryTableSpec describes one table's shape for recoverTable's +// generic table-by-table copy. columns MUST include the table's own +// primary key column(s) explicitly so recovered rows keep their original +// sequence/identity rather than being renumbered — cost_ledger, plan_items, +// and producers all reference events.sequence by foreign key, and sequence +// is the sole ordering authority anywhere in the kernel (determinism.md), +// so a recovered file MUST NOT renumber it. +type recoveryTableSpec struct { + name string + columns []string + orderBy string +} + +// recoveryTableSpecs lists every append-only table in dependency order: +// events first (nothing depends on it having already been inserted, but +// everything else references it), then its three dependents. session_meta +// is handled separately by recoverSessionMeta since it's the one row that +// makes recovery worth attempting at all. +var recoveryTableSpecs = []recoveryTableSpec{ + { + name: "events", + columns: []string{"sequence", "id", "timestamp", "kind", "producer_category", "producer_name", "producer_version", "schema_version", "payload"}, + orderBy: "sequence", + }, + { + name: "producers", + columns: []string{"category", "name", "version", "first_seen_sequence"}, + orderBy: "category, name, version", + }, + { + name: "cost_ledger", + columns: []string{"sequence", "event_sequence", "provider_name", "model_id", "input_tokens", "output_tokens", "cache_write_tokens", "cache_read_tokens", "cost_usd"}, + orderBy: "sequence", + }, + { + name: "plan_items", + columns: []string{"sequence", "event_sequence", "turn_id", "tool_call_id", "provider_name", "tool_name", "decision", "decided_by"}, + orderBy: "sequence", + }, +} + +// checkIntegrity implements the Stage 1 seam for real: it opens +// sessionID's file at path, runs PRAGMA integrity_check, and — if the +// file can't even be opened, or integrity_check reports problems — +// attempts salvage recovery per +// docs/specifications/state-backend.md#corruption-recovery. It returns the +// *sql.DB to use going forward: db opened directly on path when healthy, +// or a fresh handle on the recovered file after a successful salvage. +// ErrUnrecoverable if salvage itself fails. +func (st *Store) checkIntegrity(ctx context.Context, path, sessionID string) (_ *sql.DB, err error) { + ctx, span := st.telemetry.StartStateBackendIntegrityCheck(ctx, sessionID) + defer func() { telemetry.EndSpan(span, err) }() + + db, openErr := openDB(ctx, path) + if openErr == nil { + problems, checkErr := runIntegrityCheck(ctx, db) + if checkErr == nil && len(problems) == 0 { + return db, nil + } + _ = db.Close() + if checkErr != nil { + st.logger.WarnContext(ctx, "statebackend: integrity check could not complete, attempting recovery", "session_id", sessionID, "err", checkErr) + } else { + st.logger.WarnContext(ctx, "statebackend: session file failed integrity check, attempting recovery", "session_id", sessionID, "problem_count", len(problems), "problems", strings.Join(problems, "; ")) + } + } else { + st.logger.WarnContext(ctx, "statebackend: session file could not be opened, attempting recovery", "session_id", sessionID, "err", openErr) + } + + // Either the file couldn't even be opened, or it opened but failed + // integrity_check — both routes attempt the same salvage, per + // state-backend.md's corruption-recovery section, which doesn't + // distinguish the two: both mean "this file's contents can't be + // trusted as-is." + corruptPath := path + ".corrupt" + if renameErr := os.Rename(path, corruptPath); renameErr != nil { + err = fmt.Errorf("statebackend: recover %s: rename damaged file: %w", sessionID, renameErr) + return nil, err + } + + recovered, stats, recErr := st.recoverSession(ctx, corruptPath, path, sessionID) + if recErr != nil { + st.logger.WarnContext(ctx, "statebackend: recovery failed, session flagged unreadable", "session_id", sessionID, "corrupt_path", corruptPath, "err", recErr) + // Double %w (Go 1.20+) keeps recErr in the chain for diagnosis + // while errors.Is(err, ErrUnrecoverable) still holds for callers + // that only care about the category. + err = fmt.Errorf("statebackend: recover %s: %w: %w", sessionID, ErrUnrecoverable, recErr) + return nil, err + } + + st.logger.WarnContext(ctx, "statebackend: session recovered from corruption", + "session_id", sessionID, "corrupt_path", corruptPath, + "events_salvaged", stats.salvaged["events"], "events_skipped", stats.skipped["events"], + "producers_salvaged", stats.salvaged["producers"], "producers_skipped", stats.skipped["producers"], + "cost_ledger_salvaged", stats.salvaged["cost_ledger"], "cost_ledger_skipped", stats.skipped["cost_ledger"], + "plan_items_salvaged", stats.salvaged["plan_items"], "plan_items_skipped", stats.skipped["plan_items"], + ) + return recovered, nil +} + +// runIntegrityCheck runs PRAGMA integrity_check against db +// (docs/specifications/state-backend.md#corruption-recovery names this +// pragma specifically — not quick_check, which trades thoroughness for +// speed) and returns the problem strings sqlite reports. A healthy +// database returns exactly one row containing "ok", which this function +// reports as a nil, empty result. A non-nil error means the check itself +// could not run to completion — the file is damaged badly enough that +// even the diagnostic query fails. +func runIntegrityCheck(ctx context.Context, db *sql.DB) ([]string, error) { + rows, err := db.QueryContext(ctx, "PRAGMA integrity_check") + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + var problems []string + for rows.Next() { + var msg string + if err := rows.Scan(&msg); err != nil { + return nil, err + } + if msg != "ok" { + problems = append(problems, msg) + } + } + if err := rows.Err(); err != nil { + return nil, err + } + return problems, nil +} + +// recoverSession attempts state-backend.md's salvage: a fresh, +// schema-correct file built at dstPath+".recovering", populated +// table-by-table from srcPath (the just-renamed damaged original), +// tolerating per-row read/insert failures. Recovery is only considered +// successful once session_meta — the one row every session file must +// carry — is itself salvaged; every other table is best-effort on top of +// that. On success, the recovered file is installed at dstPath (the +// original session filename, per spec — the .corrupt rename already +// happened in checkIntegrity before this is called) and the returned +// *sql.DB is opened on it. +func (st *Store) recoverSession(ctx context.Context, srcPath, dstPath, sessionID string) (*sql.DB, *recoveryStats, error) { + recoveryPath := dstPath + ".recovering" + _ = os.Remove(recoveryPath) // best-effort: clear any stale attempt left by a previous crashed recovery + + // Pre-create recoveryPath at 0600 ourselves, mirroring Create()'s + // pattern (statebackend.go). Left to sql.Open/modernc's own file + // creation, the file lands at whatever mode the SQLite library + // defaults to (0644 modulo umask), silently breaking this package's + // documented 0600 guarantee for any session that has been through + // recovery — os.Rename below then installs that file, mode and all, + // permanently at the canonical session path. + // #nosec G304 -- recoveryPath is derived from dstPath, itself derived from an already-ValidateSessionID-checked session ID, not attacker-controlled input. + f, err := os.OpenFile(recoveryPath, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0o600) + if err != nil { + return nil, nil, fmt.Errorf("create recovery file: %w", err) + } + if err := f.Close(); err != nil { + _ = os.Remove(recoveryPath) + return nil, nil, fmt.Errorf("create recovery file: %w", err) + } + + recDB, err := openDB(ctx, recoveryPath) + if err != nil { + _ = os.Remove(recoveryPath) + return nil, nil, fmt.Errorf("create recovery file: %w", err) + } + // cleanupRecoveryFile guards a half-built recovery attempt: true until + // the file is fully populated and closed, at which point it becomes a + // complete, valid artifact worth preserving even if the final install + // rename below fails (never destroy salvaged data). + cleanupRecoveryFile := true + defer func() { + if cleanupRecoveryFile { + _ = recDB.Close() + _ = os.Remove(recoveryPath) + } + }() + + if err := initSchema(ctx, recDB); err != nil { + return nil, nil, fmt.Errorf("init recovery schema: %w", err) + } + + // #nosec G304 -- srcPath is the Store's own just-renamed .corrupt file (path+".corrupt"), derived from an already-ValidateSessionID-checked session ID, not attacker-controlled input. + srcDB, err := sql.Open("sqlite", srcPath) + if err != nil { + return nil, nil, fmt.Errorf("open damaged file for reading: %w", err) + } + defer func() { _ = srcDB.Close() }() + + if !recoverSessionMeta(ctx, srcDB, recDB, sessionID) { + return nil, nil, fmt.Errorf("session_meta row unreadable: %w", ErrUnrecoverable) + } + + stats := &recoveryStats{salvaged: make(map[string]int, len(recoveryTableSpecs)), skipped: make(map[string]int, len(recoveryTableSpecs))} + for _, spec := range recoveryTableSpecs { + salvaged, skipped := recoverTable(ctx, srcDB, recDB, spec) + stats.salvaged[spec.name] = salvaged + stats.skipped[spec.name] = skipped + } + + if err := recDB.Close(); err != nil { + return nil, nil, fmt.Errorf("close recovery file: %w", err) + } + cleanupRecoveryFile = false // the recovery file is complete and valid from here on, regardless of what happens below. + + if err := os.Rename(recoveryPath, dstPath); err != nil { + return nil, nil, fmt.Errorf("install recovered file (recovered data preserved at %s): %w", recoveryPath, err) + } + + newDB, err := openDB(ctx, dstPath) + if err != nil { + return nil, nil, fmt.Errorf("reopen recovered file: %w", err) + } + return newDB, stats, nil +} + +// recoverSessionMeta copies session_meta's single row from src into dst, +// reporting whether it succeeded. +func recoverSessionMeta(ctx context.Context, src, dst *sql.DB, sessionID string) bool { + meta, err := querySessionMeta(ctx, src, sessionID) + if err != nil { + return false + } + if err := insertSessionMeta(ctx, dst, meta); err != nil { + return false + } + return true +} + +// recoverTable performs spec's best-effort, per-row-tolerant table copy +// from src to dst per spec: a row that fails to scan, or fails to insert +// (including a foreign-key violation against a parent row that was itself +// skipped), is counted as skipped and the copy continues with the next +// row rather than aborting the whole table. If the SELECT itself can't +// even start (the table's pages are entirely unreadable), the table is +// reported as 0 salvaged, 0 skipped rather than erroring the whole +// recovery — that's exactly the "tolerate per-row/per-table failure" +// posture this function exists for. +func recoverTable(ctx context.Context, src, dst *sql.DB, spec recoveryTableSpec) (salvaged, skipped int) { + // #nosec G201 -- spec's name/columns/orderBy all come from the hardcoded, package-level recoveryTableSpecs slice above, never from caller or row data; the actual row values are always bound as placeholders below, not interpolated. + selectQuery := fmt.Sprintf("SELECT %s FROM %s ORDER BY %s", strings.Join(spec.columns, ", "), spec.name, spec.orderBy) + rows, err := src.QueryContext(ctx, selectQuery) + if err != nil { + return 0, 0 + } + defer func() { _ = rows.Close() }() + + placeholders := make([]string, len(spec.columns)) + for i := range placeholders { + placeholders[i] = "?" + } + // #nosec G201 -- same reasoning as selectQuery above: spec.name/columns are hardcoded, never caller-controlled. + insertQuery := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", spec.name, strings.Join(spec.columns, ", "), strings.Join(placeholders, ", ")) + + for rows.Next() { + values := make([]any, len(spec.columns)) + scanDests := make([]any, len(spec.columns)) + for i := range values { + scanDests[i] = &values[i] + } + if err := rows.Scan(scanDests...); err != nil { + skipped++ + continue + } + if _, err := dst.ExecContext(ctx, insertQuery, values...); err != nil { + skipped++ + continue + } + salvaged++ + } + // rows.Err() being non-nil here means iteration was cut short by a + // read error partway through the table. Everything already salvaged + // above is kept; the unread remainder simply isn't separately + // countable through this API, so it shows up only as + // salvaged+skipped summing to less than the table's original row + // count in the recovery log — never as a hard failure of the whole + // recovery. + + return salvaged, skipped +} diff --git a/internal/statebackend/integrity_test.go b/internal/statebackend/integrity_test.go new file mode 100644 index 0000000..f75ffe4 --- /dev/null +++ b/internal/statebackend/integrity_test.go @@ -0,0 +1,320 @@ +package statebackend + +import ( + "bytes" + "context" + "errors" + "fmt" + "io/fs" + "log/slog" + "os" + "runtime" + "strings" + "testing" + "time" +) + +// newTestStoreWithLogBuffer returns a Store (like newTestStore) whose +// WithLogger writes DEBUG-and-up records to the returned buffer, for tests +// asserting that recovery logs a WARN — matching the bytes.Buffer + +// slog.TextHandler recording-handler convention used throughout this +// package (e.g. statebackend_test.go's TestWithLogger). +func newTestStoreWithLogBuffer(t *testing.T, opts ...Option) (*Store, *bytes.Buffer) { + t.Helper() + var buf bytes.Buffer + logger := slog.New(slog.NewTextHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug})) + st := newTestStore(t, append([]Option{WithLogger(logger)}, opts...)...) + return st, &buf +} + +// corruptRegion overwrites length bytes at offset with a repeating, +// clearly-not-original pattern, simulating localized on-disk corruption +// (e.g. a bad sector, a partial write) without touching the file's +// header — the file must remain openable, but PRAGMA integrity_check must +// find something wrong with its data pages. +func corruptRegion(t *testing.T, path string, offset, length int) { + t.Helper() + f, err := os.OpenFile(path, os.O_RDWR, 0o600) // #nosec G304 -- test helper operating on its own t.TempDir()-rooted fixture path + if err != nil { + t.Fatalf("OpenFile: %v", err) + } + defer func() { _ = f.Close() }() + + garbage := bytes.Repeat([]byte{0xDE, 0xAD, 0xBE, 0xEF}, length/4+1)[:length] + if _, err := f.WriteAt(garbage, int64(offset)); err != nil { + t.Fatalf("WriteAt: %v", err) + } +} + +// corruptHeader overwrites the file's 16-byte "SQLite format 3\000" magic +// header, which sqlite validates on every connection — this reliably +// makes even a basic PRAGMA fail with "file is not a database," simulating +// a file too destroyed to open at all, as opposed to corruptRegion's +// page-level damage that PRAGMA integrity_check specifically detects. +func corruptHeader(t *testing.T, path string) { + t.Helper() + f, err := os.OpenFile(path, os.O_RDWR, 0o600) // #nosec G304 -- test helper operating on its own t.TempDir()-rooted fixture path + if err != nil { + t.Fatalf("OpenFile: %v", err) + } + defer func() { _ = f.Close() }() + if _, err := f.WriteAt(bytes.Repeat([]byte{0xFF}, 16), 0); err != nil { + t.Fatalf("WriteAt: %v", err) + } +} + +func TestStore_Open_healthyFileNoRecovery(t *testing.T) { + t.Parallel() + st, logBuf := newTestStoreWithLogBuffer(t) + meta := testSessionMeta() + created := createSession(t, st, meta) + + for i := range 3 { + if _, err := created.AppendEvent(context.Background(), testEvent(fmt.Sprintf("evt-%d", i))); err != nil { + t.Fatalf("AppendEvent[%d]: %v", i, err) + } + } + + opened, err := st.Open(context.Background(), meta.SessionID) + if err != nil { + t.Fatalf("Open: %v", err) + } + t.Cleanup(func() { _ = opened.Close() }) + + if strings.Contains(logBuf.String(), "recover") { + t.Errorf("unexpected recovery log for a healthy file: %s", logBuf.String()) + } + if _, statErr := os.Stat(st.sessionPath(meta.SessionID) + ".corrupt"); !errors.Is(statErr, fs.ErrNotExist) { + t.Errorf(".corrupt file exists for a healthy file: err=%v", statErr) + } +} + +// TestStore_Open_corruptionTriggersRecovery covers the recoverable-damage +// path: localized page corruption is caught by PRAGMA integrity_check, the +// damaged original is renamed to .corrupt, a fresh usable file with +// salvaged rows takes its place, and recovery is logged at WARN. +func TestStore_Open_corruptionTriggersRecovery(t *testing.T) { + t.Parallel() + st, logBuf := newTestStoreWithLogBuffer(t) + meta := testSessionMeta() + sess, err := st.Create(context.Background(), meta) + if err != nil { + t.Fatalf("Create: %v", err) + } + + // Enough data across several pages that corrupting a couple of + // scattered regions is very likely to hit real, populated structure + // rather than free space. + const eventCount = 8 + for i := range eventCount { + ev := testEvent(fmt.Sprintf("evt-%d", i)) + ev.Payload = bytes.Repeat([]byte{byte(i)}, 2048) + if _, err := sess.AppendEvent(context.Background(), ev); err != nil { + t.Fatalf("AppendEvent[%d]: %v", i, err) + } + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + path := st.sessionPath(meta.SessionID) + info, err := os.Stat(path) + if err != nil { + t.Fatalf("Stat: %v", err) + } + if info.Size() < 32768 { + t.Fatalf("file too small (%d bytes) for a reliable corruption test", info.Size()) + } + // Table root pages are allocated at CREATE TABLE time (schema.go's + // initSchema runs events, then session_meta, then the rest, all + // before any row exists), so session_meta's root page — and its one + // small row — lands early in the file, well before the additional + // data pages events' later, larger-payload rows spill into as the + // file grows. Trashing the back half of the file corrupts real, + // populated event data pages with overwhelming likelihood while + // leaving session_meta's early page untouched, so recovery has + // something to actually salvage. + corruptStart := int(info.Size()) / 2 + corruptRegion(t, path, corruptStart, int(info.Size())-corruptStart) + + opened, err := st.Open(context.Background(), meta.SessionID) + if err != nil { + t.Fatalf("Open (after corruption): %v", err) + } + t.Cleanup(func() { _ = opened.Close() }) + + if _, statErr := os.Stat(path + ".corrupt"); statErr != nil { + t.Errorf(".corrupt file missing after recovery: %v", statErr) + } + + // The recovered file installed at the canonical path must carry the + // same 0600 guarantee every session file does — recoverSession + // pre-creates it itself rather than leaving the mode up to + // sql.Open/modernc's own file creation (whose default is 0644 modulo + // umask). Unix permission bits aren't meaningful on Windows, so this + // assertion only runs on Unix-like systems. + if runtime.GOOS != "windows" { + recoveredInfo, statErr := os.Stat(path) + if statErr != nil { + t.Fatalf("Stat (recovered file): %v", statErr) + } + if perm := recoveredInfo.Mode().Perm(); perm != 0o600 { + t.Errorf("recovered file perm = %o, want %o", perm, 0o600) + } + } + + gotMeta, err := opened.Meta(context.Background()) + if err != nil { + t.Fatalf("Meta (recovered file): %v", err) + } + if gotMeta.SessionID != meta.SessionID { + t.Errorf("recovered SessionID = %q, want %q", gotMeta.SessionID, meta.SessionID) + } + + logOutput := logBuf.String() + if !strings.Contains(logOutput, "recover") { + t.Errorf("no recovery log captured; log = %q", logOutput) + } + if !strings.Contains(logOutput, meta.SessionID) { + t.Errorf("recovery log missing session_id; log = %q", logOutput) + } + if !strings.Contains(logOutput, "level=WARN") { + t.Errorf("recovery log not at WARN level; log = %q", logOutput) + } +} + +// TestStore_Open_fullyDestroyedFileUnrecoverable covers the unsalvageable +// case: the file's own header is destroyed, so neither the initial open +// nor a recovery attempt's own read of the (renamed) damaged file can +// succeed. Open must return ErrUnrecoverable, and the damaged file must +// still be renamed to .corrupt (never deleted, never left at its +// original name). +func TestStore_Open_fullyDestroyedFileUnrecoverable(t *testing.T) { + t.Parallel() + st, logBuf := newTestStoreWithLogBuffer(t) + meta := testSessionMeta() + sess, err := st.Create(context.Background(), meta) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + path := st.sessionPath(meta.SessionID) + corruptHeader(t, path) + + _, err = st.Open(context.Background(), meta.SessionID) + if !errors.Is(err, ErrUnrecoverable) { + t.Fatalf("Open (destroyed file) err = %v, want ErrUnrecoverable", err) + } + + if _, statErr := os.Stat(path + ".corrupt"); statErr != nil { + t.Errorf(".corrupt file missing after unrecoverable failure: %v", statErr) + } + if _, statErr := os.Stat(path); !errors.Is(statErr, fs.ErrNotExist) { + t.Errorf("original path still exists after being renamed to .corrupt: err=%v", statErr) + } + if _, statErr := os.Stat(path + ".recovering"); !errors.Is(statErr, fs.ErrNotExist) { + t.Errorf("stale .recovering file left behind: err=%v", statErr) + } + + logOutput := logBuf.String() + if !strings.Contains(logOutput, "unreadable") { + t.Errorf("no unreadable/failure log captured; log = %q", logOutput) + } +} + +func TestRunIntegrityCheck_healthy(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + problems, err := runIntegrityCheck(ctx, db) + if err != nil { + t.Fatalf("runIntegrityCheck: %v", err) + } + if len(problems) != 0 { + t.Errorf("problems = %v, want none", problems) + } +} + +func TestRecoverTable_unreadableTableReturnsZeros(t *testing.T) { + t.Parallel() + ctx := context.Background() + src := openTestDB(t) // no schema applied: every table is "unreadable" + dst := openTestDB(t) + if err := initSchema(ctx, dst); err != nil { + t.Fatalf("initSchema: %v", err) + } + + salvaged, skipped := recoverTable(ctx, src, dst, recoveryTableSpecs[0]) + if salvaged != 0 || skipped != 0 { + t.Errorf("recoverTable against an unreadable source = (%d, %d), want (0, 0)", salvaged, skipped) + } +} + +func TestRecoverSessionMeta_unreadableReturnsFalse(t *testing.T) { + t.Parallel() + ctx := context.Background() + src := openTestDB(t) // no schema: session_meta doesn't exist + dst := openTestDB(t) + if err := initSchema(ctx, dst); err != nil { + t.Fatalf("initSchema: %v", err) + } + + if recoverSessionMeta(ctx, src, dst, "any-id") { + t.Error("recoverSessionMeta against an unreadable source = true, want false") + } +} + +func TestStore_recoverSession_installFailureLeavesDataAtRecoveringPath(t *testing.T) { + t.Parallel() + st := newTestStore(t) + meta := testSessionMeta() + sess, err := st.Create(context.Background(), meta) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + path := st.sessionPath(meta.SessionID) + corruptPath := path + ".corrupt" + if err := os.Rename(path, corruptPath); err != nil { + t.Fatalf("rename to .corrupt: %v", err) + } + + // Occupy the install target with a directory so the final + // os.Rename(recoveryPath, dstPath) inside recoverSession fails after + // the recovery file is already fully built — recoverSession must not + // delete the completed .recovering file in that case. + if err := os.Mkdir(path, 0o700); err != nil { + t.Fatalf("Mkdir (blocking the install path): %v", err) + } + + _, _, err = st.recoverSession(context.Background(), corruptPath, path, meta.SessionID) + if err == nil { + t.Fatal("recoverSession (blocked install path) = nil error, want error") + } + + if _, statErr := os.Stat(path + ".recovering"); statErr != nil { + t.Errorf("recovered data at .recovering was not preserved after a failed install: %v", statErr) + } +} + +func TestStore_Open_notFoundStillReturnsErrNotFound(t *testing.T) { + // Sanity check that the integrity-check rewrite of Open didn't fold + // the plain "file doesn't exist" case into the corruption-recovery + // path — a missing file is not the same thing as a damaged one. + t.Parallel() + st := newTestStore(t) + _, err := st.Open(context.Background(), NewSessionID(time.Now())) + if !errors.Is(err, ErrNotFound) { + t.Errorf("Open (missing file) err = %v, want ErrNotFound", err) + } +} diff --git a/internal/statebackend/query.go b/internal/statebackend/query.go new file mode 100644 index 0000000..b0c9b84 --- /dev/null +++ b/internal/statebackend/query.go @@ -0,0 +1,259 @@ +package statebackend + +import ( + "context" + "database/sql" + "fmt" + "iter" + + "github.com/pluggableharness/agent/internal/telemetry" + commonv1 "github.com/pluggableharness/agent/pkg/common/proto/v1" +) + +// rowScanner is the subset of *sql.Row and *sql.Rows this package's scan +// helpers need, letting one scan function serve both a single-row +// QueryRowContext caller and a multi-row QueryContext loop. +type rowScanner interface { + Scan(dest ...any) error +} + +// Meta returns this session's session_meta row. Returns ErrClosed if Close +// has already been called. +func (s *Session) Meta(ctx context.Context) (_ SessionMeta, err error) { + if s.closed.Load() { + return SessionMeta{}, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendMetaQuery(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: querying session meta", "session_id", s.id) + + meta, metaErr := querySessionMeta(ctx, s.db, s.id) + if metaErr != nil { + err = metaErr + return SessionMeta{}, err + } + return meta, nil +} + +// Events returns every event in this session's file as a sequence-ordered +// iter.Seq2 — sequence is the sole ordering authority +// (docs/specifications/state-backend.md#ordering--concurrency, +// determinism.md), never wall-clock time. Each Event's Sequence is +// populated and Payload is byte-identical to what was appended. A decode +// or read error surfaces through the error side of the pair and stops +// iteration — the caller sees no further events after that point. If +// Close has already been called, the sequence yields exactly one +// (Event{}, ErrClosed) pair. +func (s *Session) Events(ctx context.Context) iter.Seq2[Event, error] { + return func(yield func(Event, error) bool) { + if s.closed.Load() { + yield(Event{}, ErrClosed) + return + } + + ctx, span := s.telemetry.StartStateBackendEventsQuery(ctx, s.id) + var err error + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: querying events", "session_id", s.id) + + const q = `SELECT sequence, id, timestamp, kind, producer_category, producer_name, producer_version, schema_version, payload FROM events ORDER BY sequence` + rows, queryErr := s.db.QueryContext(ctx, q) + if queryErr != nil { + err = fmt.Errorf("statebackend: query events: %w", queryErr) + yield(Event{}, err) + return + } + defer func() { _ = rows.Close() }() + + for rows.Next() { + ev, scanErr := scanEvent(rows) + if scanErr != nil { + err = fmt.Errorf("statebackend: query events: %w", scanErr) + yield(Event{}, err) + return + } + if !yield(ev, nil) { + return + } + } + if rowsErr := rows.Err(); rowsErr != nil { + err = fmt.Errorf("statebackend: query events: %w", rowsErr) + yield(Event{}, err) + } + } +} + +// scanEvent decodes one events row, translating its stored TEXT +// kind/producer_category back into their proto enum values. +func scanEvent(row rowScanner) (Event, error) { + var ( + ev Event + timestampText, kindText string + categoryText, name, version string + ) + if err := row.Scan(&ev.Sequence, &ev.ID, ×tampText, &kindText, &categoryText, &name, &version, &ev.SchemaVersion, &ev.Payload); err != nil { + return Event{}, err + } + + timestamp, err := parseTimestamp(timestampText) + if err != nil { + return Event{}, fmt.Errorf("timestamp: %w", err) + } + ev.Timestamp = timestamp + + kind, err := decodeEventKind(kindText) + if err != nil { + return Event{}, err + } + ev.Kind = kind + + category, err := decodeProducerCategory(categoryText) + if err != nil { + return Event{}, err + } + ev.Producer = &commonv1.ProducerRef{Category: category, Name: name, Version: version} + + return ev, nil +} + +// Producers returns the distinct set of producers that have written to +// this session's file (docs/specifications/state-backend.md#producers — +// the "install X to re-render this" preflight list), ordered +// deterministically by (category, name, version). Returns ErrClosed if +// Close has already been called. +func (s *Session) Producers(ctx context.Context) (_ []*commonv1.ProducerRef, err error) { + if s.closed.Load() { + return nil, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendProducersQuery(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: querying producers", "session_id", s.id) + + const q = `SELECT category, name, version FROM producers ORDER BY category, name, version` + rows, queryErr := s.db.QueryContext(ctx, q) + if queryErr != nil { + err = fmt.Errorf("statebackend: query producers: %w", queryErr) + return nil, err + } + defer func() { _ = rows.Close() }() + + var producers []*commonv1.ProducerRef + for rows.Next() { + var categoryText, name, version string + if scanErr := rows.Scan(&categoryText, &name, &version); scanErr != nil { + err = fmt.Errorf("statebackend: query producers: %w", scanErr) + return nil, err + } + category, decErr := decodeProducerCategory(categoryText) + if decErr != nil { + err = fmt.Errorf("statebackend: query producers: %w", decErr) + return nil, err + } + producers = append(producers, &commonv1.ProducerRef{Category: category, Name: name, Version: version}) + } + if rowsErr := rows.Err(); rowsErr != nil { + err = fmt.Errorf("statebackend: query producers: %w", rowsErr) + return nil, err + } + return producers, nil +} + +// TotalCostUSD returns SUM(cost_ledger.cost_usd), or 0 if the session has +// no cost_ledger rows yet. Returns ErrClosed if Close has already been +// called. +func (s *Session) TotalCostUSD(ctx context.Context) (_ float64, err error) { + if s.closed.Load() { + return 0, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendCostQuery(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: querying total cost", "session_id", s.id) + + var total sql.NullFloat64 + row := s.db.QueryRowContext(ctx, "SELECT SUM(cost_usd) FROM cost_ledger") + if scanErr := row.Scan(&total); scanErr != nil { + err = fmt.Errorf("statebackend: query total cost: %w", scanErr) + return 0, err + } + // SUM over zero rows is SQL NULL, not 0 — total.Valid is false in that + // case and total.Float64's zero value is exactly the "0 for none" + // this method documents. + return total.Float64, nil +} + +// CostLedger returns every cost_ledger row, in append order (sequence — +// the sole ordering authority). Returns ErrClosed if Close has already +// been called. +func (s *Session) CostLedger(ctx context.Context) (_ []CostEntry, err error) { + if s.closed.Load() { + return nil, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendCostLedgerQuery(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: querying cost ledger", "session_id", s.id) + + const q = `SELECT provider_name, model_id, input_tokens, output_tokens, cache_write_tokens, cache_read_tokens, cost_usd FROM cost_ledger ORDER BY sequence` + rows, queryErr := s.db.QueryContext(ctx, q) + if queryErr != nil { + err = fmt.Errorf("statebackend: query cost ledger: %w", queryErr) + return nil, err + } + defer func() { _ = rows.Close() }() + + var entries []CostEntry + for rows.Next() { + var c CostEntry + if scanErr := rows.Scan(&c.ProviderName, &c.ModelID, &c.InputTokens, &c.OutputTokens, &c.CacheWriteTokens, &c.CacheReadTokens, &c.CostUSD); scanErr != nil { + err = fmt.Errorf("statebackend: query cost ledger: %w", scanErr) + return nil, err + } + entries = append(entries, c) + } + if rowsErr := rows.Err(); rowsErr != nil { + err = fmt.Errorf("statebackend: query cost ledger: %w", rowsErr) + return nil, err + } + return entries, nil +} + +// PlanItems returns every plan_items row, in append order (sequence — the +// sole ordering authority). Returns ErrClosed if Close has already been +// called. +func (s *Session) PlanItems(ctx context.Context) (_ []PlanItem, err error) { + if s.closed.Load() { + return nil, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendPlanItemsQuery(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: querying plan items", "session_id", s.id) + + const q = `SELECT turn_id, tool_call_id, provider_name, tool_name, decision, decided_by FROM plan_items ORDER BY sequence` + rows, queryErr := s.db.QueryContext(ctx, q) + if queryErr != nil { + err = fmt.Errorf("statebackend: query plan items: %w", queryErr) + return nil, err + } + defer func() { _ = rows.Close() }() + + var items []PlanItem + for rows.Next() { + var item PlanItem + var decisionText string + if scanErr := rows.Scan(&item.TurnID, &item.ToolCallID, &item.ProviderName, &item.ToolName, &decisionText, &item.DecidedBy); scanErr != nil { + err = fmt.Errorf("statebackend: query plan items: %w", scanErr) + return nil, err + } + decision, decErr := decodePlanDecision(decisionText) + if decErr != nil { + err = fmt.Errorf("statebackend: query plan items: %w", decErr) + return nil, err + } + item.Decision = decision + items = append(items, item) + } + if rowsErr := rows.Err(); rowsErr != nil { + err = fmt.Errorf("statebackend: query plan items: %w", rowsErr) + return nil, err + } + return items, nil +} diff --git a/internal/statebackend/query_test.go b/internal/statebackend/query_test.go new file mode 100644 index 0000000..c1112ba --- /dev/null +++ b/internal/statebackend/query_test.go @@ -0,0 +1,501 @@ +package statebackend + +import ( + "bytes" + "context" + "errors" + "fmt" + "testing" + "time" + + commonv1 "github.com/pluggableharness/agent/pkg/common/proto/v1" + kernelv1 "github.com/pluggableharness/agent/pkg/kernel/proto/v1" + planv1 "github.com/pluggableharness/agent/pkg/plan/proto/v1" + sessionv1 "github.com/pluggableharness/agent/pkg/session/proto/v1" +) + +func TestSession_Meta_roundTrip(t *testing.T) { + t.Parallel() + st := newTestStore(t) + meta := testSessionMeta() + sess := createSession(t, st, meta) + + got, err := sess.Meta(context.Background()) + if err != nil { + t.Fatalf("Meta: %v", err) + } + if got.SessionID != meta.SessionID || got.Profile != meta.Profile || got.Status != meta.Status { + t.Errorf("Meta = %+v, want matching %+v", got, meta) + } +} + +func TestSession_Meta_errClosed(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := sess.Meta(context.Background()); !errors.Is(err, ErrClosed) { + t.Errorf("Meta after Close err = %v, want ErrClosed", err) + } +} + +// TestSession_Events_roundTrip is the required round-trip test: events +// with varied payloads (empty, large ~1MB, arbitrary bytes) must read back +// via Events() byte-identical, in exact sequence order. +func TestSession_Events_roundTrip(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + large := make([]byte, 1<<20) // ~1MB + for i := range large { + large[i] = byte(i % 251) + } + + payloads := [][]byte{ + {}, // empty + []byte("hello"), // ordinary text + {0x00, 0xFF, 0x01, 0x02, 0x03}, // arbitrary bytes, including NUL and high bytes + large, // ~1MB + } + + wantIDs := make([]string, len(payloads)) + for i, p := range payloads { + ev := testEvent(fmt.Sprintf("evt-%d", i)) + ev.Payload = p + if _, err := sess.AppendEvent(context.Background(), ev); err != nil { + t.Fatalf("AppendEvent[%d]: %v", i, err) + } + wantIDs[i] = ev.ID + } + + var got []Event + for ev, err := range sess.Events(context.Background()) { + if err != nil { + t.Fatalf("Events: %v", err) + } + got = append(got, ev) + } + + if len(got) != len(payloads) { + t.Fatalf("got %d events, want %d", len(got), len(payloads)) + } + for i, ev := range got { + if ev.Sequence != int64(i+1) { + t.Errorf("event[%d].Sequence = %d, want %d", i, ev.Sequence, i+1) + } + if ev.ID != wantIDs[i] { + t.Errorf("event[%d].ID = %q, want %q", i, ev.ID, wantIDs[i]) + } + if !bytes.Equal(ev.Payload, payloads[i]) { + t.Errorf("event[%d].Payload = %d bytes, want %d bytes (byte-identical)", i, len(ev.Payload), len(payloads[i])) + } + } +} + +func TestSession_Events_empty(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + var got []Event + for ev, err := range sess.Events(context.Background()) { + if err != nil { + t.Fatalf("Events: %v", err) + } + got = append(got, ev) + } + if len(got) != 0 { + t.Errorf("Events (no events appended) = %d results, want 0", len(got)) + } +} + +func TestSession_Events_errClosed(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + count := 0 + var gotErr error + for _, err := range sess.Events(context.Background()) { + count++ + gotErr = err + } + if count != 1 { + t.Fatalf("Events after Close yielded %d pairs, want exactly 1", count) + } + if !errors.Is(gotErr, ErrClosed) { + t.Errorf("Events after Close err = %v, want ErrClosed", gotErr) + } +} + +// TestSession_Events_decodeErrorSurfaces plants a row with an +// unrecognized kind directly (bypassing AppendEvent's own validation) to +// force scanEvent's decode path to fail, verifying Events() surfaces the +// error through the pair's error side and stops iterating rather than +// silently skipping the bad row. +func TestSession_Events_decodeErrorSurfaces(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + const q = `INSERT INTO events (id, timestamp, kind, producer_category, producer_name, producer_version, schema_version, payload) VALUES (?, ?, ?, ?, ?, ?, ?, ?)` + if _, err := sess.db.ExecContext(context.Background(), q, "evt-1", formatTimestamp(time.Now()), "not_a_kind", "tool", "p", "1", "v1", []byte("x")); err != nil { + t.Fatalf("insert: %v", err) + } + + count := 0 + var gotErr error + for _, err := range sess.Events(context.Background()) { + count++ + gotErr = err + } + if count != 1 { + t.Fatalf("Events yielded %d pairs, want exactly 1 (stops on first decode error)", count) + } + if !errors.Is(gotErr, ErrInvalidKind) { + t.Errorf("Events err = %v, want ErrInvalidKind", gotErr) + } +} + +func TestSession_Events_stopsOnEarlyBreak(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + for i := range 5 { + if _, err := sess.AppendEvent(context.Background(), testEvent(fmt.Sprintf("evt-%d", i))); err != nil { + t.Fatalf("AppendEvent[%d]: %v", i, err) + } + } + + count := 0 + for range sess.Events(context.Background()) { + count++ + if count == 2 { + break + } + } + if count != 2 { + t.Errorf("iteration count = %d, want 2 (stopped early)", count) + } +} + +func TestSession_Producers_dedupAndOrder(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + toolProducer := testProducer() // category=tool, name=test-tool, version=1.0.0 + + ev1 := testEvent("evt-1") + ev1.Producer = toolProducer + if _, err := sess.AppendEvent(context.Background(), ev1); err != nil { + t.Fatalf("AppendEvent[1]: %v", err) + } + + // Same producer again: must not duplicate the producers row. + ev2 := testEvent("evt-2") + ev2.Producer = toolProducer + if _, err := sess.AppendEvent(context.Background(), ev2); err != nil { + t.Fatalf("AppendEvent[2]: %v", err) + } + + ev3 := testEvent("evt-3") + ev3.Producer = &commonv1.ProducerRef{Category: commonv1.Category_CATEGORY_MEMORY, Name: "recall-plugin", Version: "2.0.0"} + if _, err := sess.AppendEvent(context.Background(), ev3); err != nil { + t.Fatalf("AppendEvent[3]: %v", err) + } + + producers, err := sess.Producers(context.Background()) + if err != nil { + t.Fatalf("Producers: %v", err) + } + if len(producers) != 2 { + t.Fatalf("Producers = %d entries, want 2 (deduped)", len(producers)) + } + // Deterministic order: category, name, version — "memory" sorts before "tool". + if producers[0].GetCategory().String() == producers[1].GetCategory().String() { + t.Errorf("Producers not distinguishable by category: %+v", producers) + } + if producers[0].GetName() != "recall-plugin" || producers[1].GetName() != "test-tool" { + t.Errorf("Producers order = [%q, %q], want [%q, %q]", producers[0].GetName(), producers[1].GetName(), "recall-plugin", "test-tool") + } +} + +// TestSession_Producers_decodeErrorSurfaces plants a producers row with an +// unrecognized category directly, verifying Producers() surfaces the +// decode error instead of returning a partially-decoded slice. +func TestSession_Producers_decodeErrorSurfaces(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + seq, err := sess.AppendEvent(context.Background(), testEvent("evt-1")) + if err != nil { + t.Fatalf("AppendEvent: %v", err) + } + + const q = `INSERT INTO producers (category, name, version, first_seen_sequence) VALUES (?, ?, ?, ?)` + if _, err := sess.db.ExecContext(context.Background(), q, "not_a_category", "p", "1", seq); err != nil { + t.Fatalf("insert producers: %v", err) + } + + if _, err := sess.Producers(context.Background()); err == nil { + t.Fatal("Producers (unrecognized category row) = nil error, want error") + } +} + +func TestSession_Producers_empty(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + producers, err := sess.Producers(context.Background()) + if err != nil { + t.Fatalf("Producers: %v", err) + } + if len(producers) != 0 { + t.Errorf("Producers (no events) = %d, want 0", len(producers)) + } +} + +func TestSession_Producers_errClosed(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := sess.Producers(context.Background()); !errors.Is(err, ErrClosed) { + t.Errorf("Producers after Close err = %v, want ErrClosed", err) + } +} + +func TestSession_TotalCostUSD_sumsAndZero(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + zero, err := sess.TotalCostUSD(context.Background()) + if err != nil { + t.Fatalf("TotalCostUSD (no rows): %v", err) + } + if zero != 0 { + t.Errorf("TotalCostUSD (no rows) = %v, want 0", zero) + } + + costs := []float64{0.01, 0.02, 0.005} + for i, c := range costs { + ev := testEvent(fmt.Sprintf("msg-%d", i)) + ev.Kind = kernelv1.EventKind_EVENT_KIND_MESSAGE + if _, err := sess.AppendMessage(context.Background(), ev, CostEntry{ProviderName: "p", ModelID: "m", CostUSD: c}); err != nil { + t.Fatalf("AppendMessage[%d]: %v", i, err) + } + } + + total, err := sess.TotalCostUSD(context.Background()) + if err != nil { + t.Fatalf("TotalCostUSD: %v", err) + } + want := 0.01 + 0.02 + 0.005 + if diff := total - want; diff > 1e-9 || diff < -1e-9 { + t.Errorf("TotalCostUSD = %v, want %v", total, want) + } +} + +func TestSession_TotalCostUSD_errClosed(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := sess.TotalCostUSD(context.Background()); !errors.Is(err, ErrClosed) { + t.Errorf("TotalCostUSD after Close err = %v, want ErrClosed", err) + } +} + +func TestSession_CostLedger_orderedContents(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + entries := []CostEntry{ + {ProviderName: "anthropic", ModelID: "m1", InputTokens: 10, OutputTokens: 5, CostUSD: 0.01}, + {ProviderName: "anthropic", ModelID: "m2", InputTokens: 20, OutputTokens: 15, CostUSD: 0.02}, + } + for i, c := range entries { + ev := testEvent(fmt.Sprintf("msg-%d", i)) + ev.Kind = kernelv1.EventKind_EVENT_KIND_MESSAGE + if _, err := sess.AppendMessage(context.Background(), ev, c); err != nil { + t.Fatalf("AppendMessage[%d]: %v", i, err) + } + } + + got, err := sess.CostLedger(context.Background()) + if err != nil { + t.Fatalf("CostLedger: %v", err) + } + if len(got) != len(entries) { + t.Fatalf("CostLedger = %d entries, want %d", len(got), len(entries)) + } + for i, want := range entries { + if got[i] != want { + t.Errorf("CostLedger[%d] = %+v, want %+v", i, got[i], want) + } + } +} + +func TestSession_CostLedger_empty(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + got, err := sess.CostLedger(context.Background()) + if err != nil { + t.Fatalf("CostLedger: %v", err) + } + if len(got) != 0 { + t.Errorf("CostLedger (no rows) = %d, want 0", len(got)) + } +} + +func TestSession_CostLedger_errClosed(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := sess.CostLedger(context.Background()); !errors.Is(err, ErrClosed) { + t.Errorf("CostLedger after Close err = %v, want ErrClosed", err) + } +} + +func TestSession_PlanItems_orderedContents(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + items := []PlanItem{ + {TurnID: "t1", ToolCallID: "c1", ProviderName: "p", ToolName: "read_file", Decision: planv1.PlanDecision_PLAN_DECISION_ALLOW, DecidedBy: "policy"}, + {TurnID: "t1", ToolCallID: "c2", ProviderName: "p", ToolName: "write_file", Decision: planv1.PlanDecision_PLAN_DECISION_DENY, DecidedBy: "operator"}, + } + ev := testEvent("plan-1") + ev.Kind = kernelv1.EventKind_EVENT_KIND_PLAN + if _, err := sess.AppendPlan(context.Background(), ev, items); err != nil { + t.Fatalf("AppendPlan: %v", err) + } + + got, err := sess.PlanItems(context.Background()) + if err != nil { + t.Fatalf("PlanItems: %v", err) + } + if len(got) != len(items) { + t.Fatalf("PlanItems = %d entries, want %d", len(got), len(items)) + } + for i, want := range items { + if got[i] != want { + t.Errorf("PlanItems[%d] = %+v, want %+v", i, got[i], want) + } + } +} + +// TestSession_PlanItems_decodeErrorSurfaces plants a plan_items row with +// an unrecognized decision directly, verifying PlanItems() surfaces the +// decode error instead of returning a partially-decoded slice. +func TestSession_PlanItems_decodeErrorSurfaces(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("plan-1") + ev.Kind = kernelv1.EventKind_EVENT_KIND_PLAN + seq, err := sess.AppendEvent(context.Background(), ev) + if err != nil { + t.Fatalf("AppendEvent: %v", err) + } + + const q = `INSERT INTO plan_items (event_sequence, turn_id, tool_call_id, provider_name, tool_name, decision, decided_by) VALUES (?, ?, ?, ?, ?, ?, ?)` + if _, err := sess.db.ExecContext(context.Background(), q, seq, "t1", "c1", "p", "tool", "not_a_decision", "policy"); err != nil { + t.Fatalf("insert plan_items: %v", err) + } + + if _, err := sess.PlanItems(context.Background()); err == nil { + t.Fatal("PlanItems (unrecognized decision row) = nil error, want error") + } +} + +func TestSession_PlanItems_empty(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + got, err := sess.PlanItems(context.Background()) + if err != nil { + t.Fatalf("PlanItems: %v", err) + } + if len(got) != 0 { + t.Errorf("PlanItems (no rows) = %d, want 0", len(got)) + } +} + +func TestSession_PlanItems_errClosed(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := sess.PlanItems(context.Background()); !errors.Is(err, ErrClosed) { + t.Errorf("PlanItems after Close err = %v, want ErrClosed", err) + } +} + +// TestSession_Meta_reflectsSetStatus is a small cross-check that Meta and +// SetStatus (session.go) agree on session_meta's contents. +func TestSession_Meta_reflectsSetStatus(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ended := time.Now().UTC().Truncate(time.Millisecond) + if err := sess.SetStatus(context.Background(), sessionv1.SessionStatus_SESSION_STATUS_FAILED, &ended); err != nil { + t.Fatalf("SetStatus: %v", err) + } + + got, err := sess.Meta(context.Background()) + if err != nil { + t.Fatalf("Meta: %v", err) + } + if got.Status != sessionv1.SessionStatus_SESSION_STATUS_FAILED { + t.Errorf("Status = %v, want FAILED", got.Status) + } + if got.EndedAt == nil || !got.EndedAt.Equal(ended) { + t.Errorf("EndedAt = %v, want %v", got.EndedAt, ended) + } +} diff --git a/internal/statebackend/schema.go b/internal/statebackend/schema.go new file mode 100644 index 0000000..574a0a3 --- /dev/null +++ b/internal/statebackend/schema.go @@ -0,0 +1,188 @@ +package statebackend + +import ( + "context" + "database/sql" + "fmt" +) + +// currentSchemaVersion is the schema revision this package implements +// (docs/specifications/state-backend.md#schema-migration). It is stamped +// into PRAGMA user_version at creation (initSchema) and checked — and +// migrated toward, if older — on every Open. +const currentSchemaVersion = 1 + +// schemaStatements is the five-table schema from +// docs/specifications/state-backend.md#schema, reproduced verbatim +// (including AUTOINCREMENT on every sequence column) as one DDL statement +// per slice element, in the document's own order, so each can be applied +// via database/sql's ExecContext — which does not guarantee multi-statement +// execution in a single call — rather than as one multi-statement string. +var schemaStatements = []string{ + `CREATE TABLE events ( + sequence INTEGER PRIMARY KEY AUTOINCREMENT, + id TEXT NOT NULL UNIQUE, + timestamp TEXT NOT NULL, + kind TEXT NOT NULL, + producer_category TEXT NOT NULL, + producer_name TEXT NOT NULL, + producer_version TEXT NOT NULL, + schema_version TEXT NOT NULL, + payload BLOB NOT NULL +)`, + `CREATE INDEX idx_events_kind ON events(kind)`, + `CREATE INDEX idx_events_producer ON events(producer_category, producer_name, producer_version)`, + `CREATE TABLE session_meta ( + session_id TEXT PRIMARY KEY, + parent_session_id TEXT, + profile TEXT NOT NULL, + status TEXT NOT NULL, + depth INTEGER NOT NULL, + started_at TEXT NOT NULL, + ended_at TEXT +)`, + `CREATE TABLE cost_ledger ( + sequence INTEGER PRIMARY KEY AUTOINCREMENT, + event_sequence INTEGER NOT NULL REFERENCES events(sequence), + provider_name TEXT NOT NULL, + model_id TEXT NOT NULL, + input_tokens INTEGER NOT NULL, + output_tokens INTEGER NOT NULL, + cache_write_tokens INTEGER NOT NULL DEFAULT 0, + cache_read_tokens INTEGER NOT NULL DEFAULT 0, + cost_usd REAL NOT NULL +)`, + `CREATE TABLE plan_items ( + sequence INTEGER PRIMARY KEY AUTOINCREMENT, + event_sequence INTEGER NOT NULL REFERENCES events(sequence), + turn_id TEXT NOT NULL, + tool_call_id TEXT NOT NULL, + provider_name TEXT NOT NULL, + tool_name TEXT NOT NULL, + decision TEXT NOT NULL, + decided_by TEXT NOT NULL +)`, + `CREATE TABLE producers ( + category TEXT NOT NULL, + name TEXT NOT NULL, + version TEXT NOT NULL, + first_seen_sequence INTEGER NOT NULL REFERENCES events(sequence), + PRIMARY KEY (category, name, version) +)`, +} + +// migrationStep is one ordered schema migration, applied transactionally by +// applyMigrations when opening a session file whose PRAGMA user_version is +// older than the target version it's being brought up to. +type migrationStep struct { + // version is the user_version this step's migrate function produces. + version int + // migrate performs the step's schema changes within tx. It MUST NOT + // touch PRAGMA user_version itself — applyMigrationStep stamps that in + // the same transaction once migrate returns successfully. + migrate func(ctx context.Context, tx *sql.Tx) error +} + +// migrations is the ordered list of schema migrations, oldest to newest, +// consulted by applyMigrations when Open finds a session file whose +// PRAGMA user_version is older than currentSchemaVersion +// (docs/specifications/state-backend.md#schema-migration). It MUST stay +// sorted ascending by version. Empty for schema version 1: there is no +// schema version older than the current baseline to migrate from yet. A +// future schema bump adds one migrationStep here — never a change to the +// baseline schemaStatements above. +var migrations = []migrationStep{} + +// initSchema creates the five-table schema and stamps +// PRAGMA user_version = currentSchemaVersion, in a single transaction, for +// a freshly created session file. Per +// docs/specifications/state-backend.md#schema-migration, every session file +// MUST carry user_version set at creation. +func initSchema(ctx context.Context, db *sql.DB) error { + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("statebackend: init schema: begin: %w", err) + } + for _, stmt := range schemaStatements { + if _, err := tx.ExecContext(ctx, stmt); err != nil { + _ = tx.Rollback() + return fmt.Errorf("statebackend: init schema: %w", err) + } + } + // PRAGMA user_version takes an integer literal, not a bind parameter; + // currentSchemaVersion is a package constant, never caller input, so + // this is not a SQL-injection surface. + if _, err := tx.ExecContext(ctx, fmt.Sprintf("PRAGMA user_version = %d", currentSchemaVersion)); err != nil { + _ = tx.Rollback() + return fmt.Errorf("statebackend: init schema: set user_version: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("statebackend: init schema: commit: %w", err) + } + return nil +} + +// readUserVersion reads PRAGMA user_version from db. Per +// docs/specifications/state-backend.md#schema-migration this MUST be +// checked before any other operation touches an opened file — mirroring +// internal/registry's lock-file-version pre-check +// (internal/registry/lockfile.go's lockFileVersion). +func readUserVersion(ctx context.Context, db *sql.DB) (int, error) { + var version int + if err := db.QueryRowContext(ctx, "PRAGMA user_version").Scan(&version); err != nil { + return 0, fmt.Errorf("statebackend: read schema version: %w", err) + } + return version, nil +} + +// applyMigrations brings db from schema version from up to target by +// running each step in steps whose version is greater than from, in +// ascending order, each inside its own transaction that also stamps +// PRAGMA user_version to that step's version +// (docs/specifications/state-backend.md#schema-migration). steps MUST be +// sorted ascending by version. Because the user_version write happens +// inside the same transaction as the step's own work, a failing step +// leaves the file at its last successfully applied version rather than +// partially migrated. +func applyMigrations(ctx context.Context, db *sql.DB, steps []migrationStep, from, target int) error { + version := from + for _, step := range steps { + if version >= target { + break + } + if step.version <= version { + continue + } + if err := applyMigrationStep(ctx, db, step); err != nil { + return err + } + version = step.version + } + if version != target { + return fmt.Errorf("statebackend: migrate: reached schema version %d, want %d: no migration path", version, target) + } + return nil +} + +// applyMigrationStep runs one migrationStep inside its own transaction, +// stamping PRAGMA user_version to step.version only after migrate succeeds, +// then commits. A failure at any point rolls the transaction back, leaving +// db's user_version unchanged. +func applyMigrationStep(ctx context.Context, db *sql.DB, step migrationStep) error { + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("statebackend: migrate to version %d: begin: %w", step.version, err) + } + if err := step.migrate(ctx, tx); err != nil { + _ = tx.Rollback() + return fmt.Errorf("statebackend: migrate to version %d: %w", step.version, err) + } + if _, err := tx.ExecContext(ctx, fmt.Sprintf("PRAGMA user_version = %d", step.version)); err != nil { + _ = tx.Rollback() + return fmt.Errorf("statebackend: migrate to version %d: set user_version: %w", step.version, err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("statebackend: migrate to version %d: commit: %w", step.version, err) + } + return nil +} diff --git a/internal/statebackend/schema_test.go b/internal/statebackend/schema_test.go new file mode 100644 index 0000000..0c109a7 --- /dev/null +++ b/internal/statebackend/schema_test.go @@ -0,0 +1,352 @@ +package statebackend + +import ( + "context" + "database/sql" + "errors" + "path/filepath" + "testing" +) + +// openTestDB opens a fresh, empty sqlite file under t.TempDir() directly +// (bypassing openDB/Store entirely), for tests exercising schema.go's +// functions in isolation. +func openTestDB(t *testing.T) *sql.DB { + t.Helper() + path := filepath.Join(t.TempDir(), "test.sqlite") + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatalf("sql.Open: %v", err) + } + t.Cleanup(func() { + if err := db.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }) + return db +} + +// tableExists reports whether name is a table in db's sqlite_master. +func tableExists(t *testing.T, db *sql.DB, name string) bool { + t.Helper() + var got string + err := db.QueryRowContext(context.Background(), "SELECT name FROM sqlite_master WHERE type = 'table' AND name = ?", name).Scan(&got) + switch { + case err == nil: + return true + case errors.Is(err, sql.ErrNoRows): + return false + default: + t.Fatalf("query sqlite_master for table %q: %v", name, err) + return false + } +} + +// indexExists reports whether name is an index in db's sqlite_master. +func indexExists(t *testing.T, db *sql.DB, name string) bool { + t.Helper() + var got string + err := db.QueryRowContext(context.Background(), "SELECT name FROM sqlite_master WHERE type = 'index' AND name = ?", name).Scan(&got) + switch { + case err == nil: + return true + case errors.Is(err, sql.ErrNoRows): + return false + default: + t.Fatalf("query sqlite_master for index %q: %v", name, err) + return false + } +} + +func TestInitSchema(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + for _, table := range []string{"events", "session_meta", "cost_ledger", "plan_items", "producers"} { + if !tableExists(t, db, table) { + t.Errorf("table %q missing after initSchema", table) + } + } + for _, idx := range []string{"idx_events_kind", "idx_events_producer"} { + if !indexExists(t, db, idx) { + t.Errorf("index %q missing after initSchema", idx) + } + } + + version, err := readUserVersion(ctx, db) + if err != nil { + t.Fatalf("readUserVersion: %v", err) + } + if version != currentSchemaVersion { + t.Errorf("user_version = %d, want %d", version, currentSchemaVersion) + } +} + +func TestInitSchema_eventsAutoincrement(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + // sqlite_sequence only exists when at least one table declares + // INTEGER PRIMARY KEY AUTOINCREMENT — its presence confirms the DDL + // kept AUTOINCREMENT verbatim rather than "optimizing" it away. + if !tableExists(t, db, "sqlite_sequence") { + t.Error("sqlite_sequence table missing — AUTOINCREMENT was not applied") + } +} + +func TestInitSchema_closedDBFailsToBegin(t *testing.T) { + t.Parallel() + db := openTestDB(t) + if err := db.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if err := initSchema(context.Background(), db); err == nil { + t.Fatal("initSchema on a closed db = nil error, want error") + } +} + +func TestInitSchema_reapplyFailsAndRollsBack(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema (first): %v", err) + } + versionAfterFirst, err := readUserVersion(ctx, db) + if err != nil { + t.Fatalf("readUserVersion: %v", err) + } + + // A second initSchema against the same, already-initialized db fails on + // its first CREATE TABLE (events already exists) and must roll back + // rather than leaving user_version half-applied. + if err := initSchema(ctx, db); err == nil { + t.Fatal("initSchema (second, against an already-initialized db) = nil error, want error") + } + + versionAfterSecond, err := readUserVersion(ctx, db) + if err != nil { + t.Fatalf("readUserVersion: %v", err) + } + if versionAfterSecond != versionAfterFirst { + t.Errorf("user_version = %d after failed reapply, want unchanged %d", versionAfterSecond, versionAfterFirst) + } +} + +func TestReadUserVersion_closedDB(t *testing.T) { + t.Parallel() + db := openTestDB(t) + if err := db.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := readUserVersion(context.Background(), db); err == nil { + t.Fatal("readUserVersion on a closed db = nil error, want error") + } +} + +func TestApplyMigrationStep_closedDBFailsToBegin(t *testing.T) { + t.Parallel() + db := openTestDB(t) + if err := db.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + step := migrationStep{version: 1, migrate: func(context.Context, *sql.Tx) error { return nil }} + if err := applyMigrationStep(context.Background(), db, step); err == nil { + t.Fatal("applyMigrationStep on a closed db = nil error, want error") + } +} + +func TestReadUserVersion_defaultsZero(t *testing.T) { + t.Parallel() + db := openTestDB(t) + version, err := readUserVersion(context.Background(), db) + if err != nil { + t.Fatalf("readUserVersion: %v", err) + } + if version != 0 { + t.Errorf("version = %d, want 0", version) + } +} + +func TestApplyMigrations(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + var applied []int + steps := []migrationStep{ + {version: 1, migrate: func(ctx context.Context, tx *sql.Tx) error { + applied = append(applied, 1) + _, err := tx.ExecContext(ctx, "CREATE TABLE migration_marker_1 (id INTEGER)") + return err + }}, + {version: 2, migrate: func(ctx context.Context, tx *sql.Tx) error { + applied = append(applied, 2) + _, err := tx.ExecContext(ctx, "CREATE TABLE migration_marker_2 (id INTEGER)") + return err + }}, + } + + if err := applyMigrations(ctx, db, steps, 0, 2); err != nil { + t.Fatalf("applyMigrations: %v", err) + } + + if len(applied) != 2 || applied[0] != 1 || applied[1] != 2 { + t.Errorf("applied steps = %v, want [1 2]", applied) + } + + version, err := readUserVersion(ctx, db) + if err != nil { + t.Fatalf("readUserVersion: %v", err) + } + if version != 2 { + t.Errorf("user_version = %d, want 2", version) + } + + for _, table := range []string{"migration_marker_1", "migration_marker_2"} { + if !tableExists(t, db, table) { + t.Errorf("table %q missing after migration", table) + } + } +} + +func TestApplyMigrations_skipsAlreadyApplied(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + var applied []int + steps := []migrationStep{ + {version: 1, migrate: func(context.Context, *sql.Tx) error { + applied = append(applied, 1) + return nil + }}, + {version: 2, migrate: func(context.Context, *sql.Tx) error { + applied = append(applied, 2) + return nil + }}, + } + + // from=1: step 1 is already applied and must be skipped, only step 2 runs. + if err := applyMigrations(ctx, db, steps, 1, 2); err != nil { + t.Fatalf("applyMigrations: %v", err) + } + if len(applied) != 1 || applied[0] != 2 { + t.Errorf("applied steps = %v, want [2] (step 1 already applied)", applied) + } +} + +func TestApplyMigrations_alreadyAtTarget(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + stepRan := false + steps := []migrationStep{ + {version: 1, migrate: func(context.Context, *sql.Tx) error { + stepRan = true + return nil + }}, + } + + if err := applyMigrations(ctx, db, steps, 1, 1); err != nil { + t.Fatalf("applyMigrations: %v", err) + } + if stepRan { + t.Error("migration step ran when db was already at target version") + } +} + +func TestApplyMigrations_noPathToTarget(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + // No steps registered at all: target 1 is unreachable from 0. + if err := applyMigrations(ctx, db, nil, 0, 1); err == nil { + t.Fatal("applyMigrations = nil error, want error (no migration path)") + } +} + +// TestApplyMigrations_failingStepLeavesVersionUnchanged verifies the +// per-step transactionality applyMigrationStep documents: a step that +// fails midway rolls back both its own schema changes and the +// PRAGMA user_version stamp together, leaving the file exactly as it was +// before the attempt. +func TestApplyMigrations_failingStepLeavesVersionUnchanged(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + boom := errors.New("boom") + steps := []migrationStep{ + {version: 1, migrate: func(ctx context.Context, tx *sql.Tx) error { + if _, err := tx.ExecContext(ctx, "CREATE TABLE should_not_persist (id INTEGER)"); err != nil { + return err + } + return boom + }}, + } + + err := applyMigrations(ctx, db, steps, 0, 1) + if !errors.Is(err, boom) { + t.Fatalf("applyMigrations err = %v, want wrapping %v", err, boom) + } + + version, verr := readUserVersion(ctx, db) + if verr != nil { + t.Fatalf("readUserVersion: %v", verr) + } + if version != 0 { + t.Errorf("user_version = %d, want 0 (unchanged after failed migration)", version) + } + + if tableExists(t, db, "should_not_persist") { + t.Error("should_not_persist table exists after a failed, rolled-back migration step") + } +} + +func TestApplyMigrations_multiStepPartialFailureStopsAtLastGood(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + + boom := errors.New("boom") + steps := []migrationStep{ + {version: 1, migrate: func(ctx context.Context, tx *sql.Tx) error { + _, err := tx.ExecContext(ctx, "CREATE TABLE step_one (id INTEGER)") + return err + }}, + {version: 2, migrate: func(context.Context, *sql.Tx) error { + return boom + }}, + } + + err := applyMigrations(ctx, db, steps, 0, 2) + if !errors.Is(err, boom) { + t.Fatalf("applyMigrations err = %v, want wrapping %v", err, boom) + } + + version, verr := readUserVersion(ctx, db) + if verr != nil { + t.Fatalf("readUserVersion: %v", verr) + } + if version != 1 { + t.Errorf("user_version = %d, want 1 (step 1 committed, step 2 rolled back)", version) + } + if !tableExists(t, db, "step_one") { + t.Error("step_one table missing — step 1's commit should have survived step 2's failure") + } +} diff --git a/internal/statebackend/session.go b/internal/statebackend/session.go new file mode 100644 index 0000000..deab3b8 --- /dev/null +++ b/internal/statebackend/session.go @@ -0,0 +1,245 @@ +package statebackend + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" + + sqlite "modernc.org/sqlite" + sqlite3 "modernc.org/sqlite/lib" + + "github.com/pluggableharness/agent/internal/telemetry" + sessionv1 "github.com/pluggableharness/agent/pkg/session/proto/v1" +) + +// AppendEvent inserts ev into the events table and upserts its producer's +// row into producers, in one transaction +// (docs/specifications/state-backend.md#events, +// docs/specifications/state-backend.md#producers), returning the sequence +// sqlite assigned. ev.Kind must not be EVENT_KIND_UNSPECIFIED +// (ErrInvalidKind); ev.ID must be unique within this session's file +// (ErrDuplicateEventID on a repeat). Returns ErrClosed if Close has already +// been called. +func (s *Session) AppendEvent(ctx context.Context, ev Event) (_ int64, err error) { + if s.closed.Load() { + return 0, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendEventAppend(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: appending event", "session_id", s.id, "event_id", ev.ID) + + seq, err := s.appendEventTx(ctx, ev, nil) + return seq, err +} + +// AppendMessage inserts ev into the events table and cost into cost_ledger +// (docs/specifications/state-backend.md#cost_ledger), plus the same +// producers upsert AppendEvent does, all in one transaction. Same +// validation and error behavior as AppendEvent. +func (s *Session) AppendMessage(ctx context.Context, ev Event, cost CostEntry) (_ int64, err error) { + if s.closed.Load() { + return 0, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendMessageAppend(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: appending message", "session_id", s.id, "event_id", ev.ID) + + seq, err := s.appendEventTx(ctx, ev, func(ctx context.Context, tx *sql.Tx, eventSeq int64) error { + return insertCostEntry(ctx, tx, eventSeq, cost) + }) + return seq, err +} + +// AppendPlan inserts ev into the events table and items into plan_items +// (docs/specifications/state-backend.md#plan_items), plus the same +// producers upsert AppendEvent does, all in one transaction. Every item's +// Decision must be ALLOW/ASK/DENY (ErrInvalidDecision otherwise). Same +// event validation and error behavior as AppendEvent. +func (s *Session) AppendPlan(ctx context.Context, ev Event, items []PlanItem) (_ int64, err error) { + if s.closed.Load() { + return 0, ErrClosed + } + ctx, span := s.telemetry.StartStateBackendPlanAppend(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: appending plan", "session_id", s.id, "event_id", ev.ID, "item_count", len(items)) + + seq, err := s.appendEventTx(ctx, ev, func(ctx context.Context, tx *sql.Tx, eventSeq int64) error { + for _, item := range items { + if err := insertPlanItem(ctx, tx, eventSeq, item); err != nil { + return err + } + } + return nil + }) + return seq, err +} + +// appendEventTx runs the one transaction shared by AppendEvent, +// AppendMessage, and AppendPlan: insert ev into events, upsert its +// producer row, then — if extra is non-nil — call it with the transaction +// and the just-assigned event sequence to insert any accompanying rows +// (cost_ledger, plan_items). extra returning an error rolls the whole +// transaction back, so the event row never exists without its +// accompanying rows. +func (s *Session) appendEventTx(ctx context.Context, ev Event, extra func(ctx context.Context, tx *sql.Tx, eventSeq int64) error) (int64, error) { + kindText, err := encodeEventKind(ev.Kind) + if err != nil { + return 0, fmt.Errorf("statebackend: append event: %w", err) + } + if ev.Producer == nil { + return 0, fmt.Errorf("statebackend: append event: producer is required") + } + categoryText, err := encodeProducerCategory(ev.Producer.GetCategory()) + if err != nil { + return 0, fmt.Errorf("statebackend: append event: %w", err) + } + + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return 0, fmt.Errorf("statebackend: append event: begin: %w", err) + } + + seq, err := insertEvent(ctx, tx, ev, kindText, categoryText) + if err != nil { + _ = tx.Rollback() + return 0, mapAppendEventError(err) + } + + if err := upsertProducer(ctx, tx, categoryText, ev.Producer.GetName(), ev.Producer.GetVersion(), seq); err != nil { + _ = tx.Rollback() + return 0, fmt.Errorf("statebackend: append event: producer: %w", err) + } + + if extra != nil { + if err := extra(ctx, tx, seq); err != nil { + _ = tx.Rollback() + return 0, fmt.Errorf("statebackend: append event: %w", err) + } + } + + if err := tx.Commit(); err != nil { + return 0, fmt.Errorf("statebackend: append event: commit: %w", err) + } + return seq, nil +} + +// mapAppendEventError translates a sqlite UNIQUE-constraint violation on +// events.id into ErrDuplicateEventID; any other error is wrapped as-is. +func mapAppendEventError(err error) error { + var sqliteErr *sqlite.Error + if errors.As(err, &sqliteErr) && sqliteErr.Code() == sqlite3.SQLITE_CONSTRAINT_UNIQUE { + return fmt.Errorf("statebackend: append event: %w", ErrDuplicateEventID) + } + return fmt.Errorf("statebackend: append event: %w", err) +} + +// insertEvent inserts ev into the events table within tx, returning the +// sequence sqlite assigned it. +func insertEvent(ctx context.Context, tx *sql.Tx, ev Event, kindText, categoryText string) (int64, error) { + const q = `INSERT INTO events (id, timestamp, kind, producer_category, producer_name, producer_version, schema_version, payload) VALUES (?, ?, ?, ?, ?, ?, ?, ?)` + result, err := tx.ExecContext(ctx, q, + ev.ID, formatTimestamp(ev.Timestamp), kindText, categoryText, + ev.Producer.GetName(), ev.Producer.GetVersion(), ev.SchemaVersion, ev.Payload, + ) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +// upsertProducer inserts a producers row for (category, name, version) +// within tx if this is the first time this session's file has seen that +// triple (docs/specifications/state-backend.md#producers) — INSERT OR +// IGNORE against the table's (category, name, version) primary key makes a +// repeat sighting a no-op rather than a constraint error. +func upsertProducer(ctx context.Context, tx *sql.Tx, category, name, version string, firstSeenSequence int64) error { + const q = `INSERT OR IGNORE INTO producers (category, name, version, first_seen_sequence) VALUES (?, ?, ?, ?)` + _, err := tx.ExecContext(ctx, q, category, name, version, firstSeenSequence) + return err +} + +// insertCostEntry inserts one cost_ledger row within tx, referencing the +// event that produced it. +func insertCostEntry(ctx context.Context, tx *sql.Tx, eventSeq int64, c CostEntry) error { + const q = `INSERT INTO cost_ledger (event_sequence, provider_name, model_id, input_tokens, output_tokens, cache_write_tokens, cache_read_tokens, cost_usd) VALUES (?, ?, ?, ?, ?, ?, ?, ?)` + _, err := tx.ExecContext(ctx, q, eventSeq, c.ProviderName, c.ModelID, c.InputTokens, c.OutputTokens, c.CacheWriteTokens, c.CacheReadTokens, c.CostUSD) + return err +} + +// insertPlanItem inserts one plan_items row within tx, referencing the +// event that produced it. +func insertPlanItem(ctx context.Context, tx *sql.Tx, eventSeq int64, item PlanItem) error { + decisionText, err := encodePlanDecision(item.Decision) + if err != nil { + return err + } + const q = `INSERT INTO plan_items (event_sequence, turn_id, tool_call_id, provider_name, tool_name, decision, decided_by) VALUES (?, ?, ?, ?, ?, ?, ?)` + _, err = tx.ExecContext(ctx, q, eventSeq, item.TurnID, item.ToolCallID, item.ProviderName, item.ToolName, decisionText, item.DecidedBy) + return err +} + +// SetStatus updates session_meta's status and ended_at in place — the one +// mutable table (docs/specifications/state-backend.md#session_meta). +// endedAt is nil while the session keeps running. Returns ErrClosed if +// Close has already been called. +func (s *Session) SetStatus(ctx context.Context, status sessionv1.SessionStatus, endedAt *time.Time) (err error) { + if s.closed.Load() { + return ErrClosed + } + ctx, span := s.telemetry.StartStateBackendStatusSet(ctx, s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: setting session status", "session_id", s.id, "status", status) + + statusText, encErr := encodeSessionStatus(status) + if encErr != nil { + err = fmt.Errorf("statebackend: set status: %w", encErr) + return err + } + + var endedAtVal any + if endedAt != nil { + endedAtVal = formatTimestamp(*endedAt) + } + + const q = `UPDATE session_meta SET status = ?, ended_at = ? WHERE session_id = ?` + result, execErr := s.db.ExecContext(ctx, q, statusText, endedAtVal, s.id) + if execErr != nil { + err = fmt.Errorf("statebackend: set status: %w", execErr) + return err + } + rows, raErr := result.RowsAffected() + if raErr != nil { + err = fmt.Errorf("statebackend: set status: %w", raErr) + return err + } + if rows == 0 { + err = fmt.Errorf("statebackend: set status: %w", ErrNotFound) + return err + } + return nil +} + +// Close checkpoints the WAL (PRAGMA wal_checkpoint(TRUNCATE), folding it +// back into the main file rather than leaving a -wal sidecar behind) and +// closes the underlying *sql.DB. Idempotent — a second Close is a no-op. +// Every write method on Session (AppendEvent, AppendMessage, AppendPlan, +// SetStatus) returns ErrClosed once Close has been called. +func (s *Session) Close() (err error) { + if !s.closed.CompareAndSwap(false, true) { + return nil + } + + ctx, span := s.telemetry.StartStateBackendSessionClose(context.Background(), s.id) + defer func() { telemetry.EndSpan(span, err) }() + s.logger.DebugContext(ctx, "statebackend: closing session", "session_id", s.id) + + if _, checkpointErr := s.db.ExecContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)"); checkpointErr != nil { + err = fmt.Errorf("statebackend: close %s: checkpoint: %w", s.id, checkpointErr) + } + if closeErr := s.db.Close(); closeErr != nil { + err = errors.Join(err, fmt.Errorf("statebackend: close %s: %w", s.id, closeErr)) + } + return err +} diff --git a/internal/statebackend/session_test.go b/internal/statebackend/session_test.go new file mode 100644 index 0000000..1bf14bd --- /dev/null +++ b/internal/statebackend/session_test.go @@ -0,0 +1,508 @@ +package statebackend + +import ( + "context" + "database/sql" + "errors" + "fmt" + "sort" + "sync" + "testing" + "time" + + commonv1 "github.com/pluggableharness/agent/pkg/common/proto/v1" + kernelv1 "github.com/pluggableharness/agent/pkg/kernel/proto/v1" + planv1 "github.com/pluggableharness/agent/pkg/plan/proto/v1" + sessionv1 "github.com/pluggableharness/agent/pkg/session/proto/v1" +) + +// testSessionMeta returns a minimal, valid SessionMeta with a fresh +// session ID, for tests that only care about a session existing. +func testSessionMeta() SessionMeta { + return SessionMeta{ + SessionID: NewSessionID(time.Now()), + Profile: "default", + Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, + StartedAt: time.Now(), + } +} + +// testProducer returns a minimal, valid producer reference. +func testProducer() *commonv1.ProducerRef { + return &commonv1.ProducerRef{ + Category: commonv1.Category_CATEGORY_TOOL, + Name: "test-tool", + Version: "1.0.0", + } +} + +// testEvent returns a minimal, valid Event with the given ID. +func testEvent(id string) Event { + return Event{ + ID: id, + Timestamp: time.Now(), + Kind: kernelv1.EventKind_EVENT_KIND_TOOL_CALL, + Producer: testProducer(), + SchemaVersion: "v1", + Payload: []byte(`{"ok":true}`), + } +} + +func TestSession_AppendEvent(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("evt-1") + seq, err := sess.AppendEvent(context.Background(), ev) + if err != nil { + t.Fatalf("AppendEvent: %v", err) + } + if seq != 1 { + t.Fatalf("sequence = %d, want 1", seq) + } + + var ( + id, kindText, category, name, version, schemaVersion string + payload []byte + ) + row := sess.db.QueryRowContext(context.Background(), + "SELECT id, kind, producer_category, producer_name, producer_version, schema_version, payload FROM events WHERE sequence = ?", seq) + if err := row.Scan(&id, &kindText, &category, &name, &version, &schemaVersion, &payload); err != nil { + t.Fatalf("query events: %v", err) + } + if id != ev.ID { + t.Errorf("id = %q, want %q", id, ev.ID) + } + if kindText != "tool_call" { + t.Errorf("kind = %q, want %q", kindText, "tool_call") + } + if category != "tool" || name != ev.Producer.Name || version != ev.Producer.Version { + t.Errorf("producer = (%q, %q, %q), want (%q, %q, %q)", category, name, version, "tool", ev.Producer.Name, ev.Producer.Version) + } + if schemaVersion != ev.SchemaVersion { + t.Errorf("schema_version = %q, want %q", schemaVersion, ev.SchemaVersion) + } + if string(payload) != string(ev.Payload) { + t.Errorf("payload = %q, want %q", payload, ev.Payload) + } + + // The producers table must have gained exactly one row for this triple. + var producerCount int + if err := sess.db.QueryRowContext(context.Background(), "SELECT COUNT(*) FROM producers WHERE category = ? AND name = ? AND version = ?", category, name, version).Scan(&producerCount); err != nil { + t.Fatalf("query producers: %v", err) + } + if producerCount != 1 { + t.Errorf("producers rows for (%q,%q,%q) = %d, want 1", category, name, version, producerCount) + } +} + +func TestSession_AppendEvent_sequential(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + for i := 1; i <= 5; i++ { + seq, err := sess.AppendEvent(context.Background(), testEvent(fmt.Sprintf("evt-%d", i))) + if err != nil { + t.Fatalf("AppendEvent[%d]: %v", i, err) + } + if seq != int64(i) { + t.Fatalf("sequence[%d] = %d, want %d", i, seq, i) + } + } +} + +// TestSession_AppendEvent_concurrentSequencesAreExactlyOneToN is the +// required race test: N goroutines append concurrently on the same +// Session, and the set of returned sequences must be exactly {1..N} with +// no gaps or duplicates — proving the sole-writer connection (SetMaxOpenConns(1)) +// serializes concurrent transactions correctly. Run under -race. +func TestSession_AppendEvent_concurrentSequencesAreExactlyOneToN(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + const n = 50 + var wg sync.WaitGroup + seqs := make([]int64, n) + errs := make([]error, n) + wg.Add(n) + for i := range n { + go func(idx int) { + defer wg.Done() + seq, err := sess.AppendEvent(context.Background(), testEvent(fmt.Sprintf("evt-%d", idx))) + seqs[idx] = seq + errs[idx] = err + }(i) + } + wg.Wait() + + for i, err := range errs { + if err != nil { + t.Fatalf("AppendEvent[%d]: %v", i, err) + } + } + + sort.Slice(seqs, func(i, j int) bool { return seqs[i] < seqs[j] }) + for i, seq := range seqs { + want := int64(i + 1) + if seq != want { + t.Fatalf("sequences = %v, want exactly 1..%d with no gaps/dupes (sorted index %d = %d, want %d)", seqs, n, i, seq, want) + } + } +} + +// TestSession_concurrentReaderDuringWrites exercises WAL mode: a second, +// independent connection to the same file must be able to read while the +// Session's sole-writer connection is actively appending — never blocked +// by, or blocking, the writer (docs/specifications/state-backend.md#ordering--concurrency). +func TestSession_concurrentReaderDuringWrites(t *testing.T) { + t.Parallel() + st := newTestStore(t) + meta := testSessionMeta() + sess := createSession(t, st, meta) + path := st.sessionPath(meta.SessionID) + + const n = 30 + writeErrs := make(chan error, 1) + go func() { + defer close(writeErrs) + for i := range n { + if _, err := sess.AppendEvent(context.Background(), testEvent(fmt.Sprintf("evt-%d", i))); err != nil { + writeErrs <- err + return + } + } + }() + + readerDB, err := sql.Open("sqlite", path) + if err != nil { + t.Fatalf("sql.Open (reader): %v", err) + } + t.Cleanup(func() { _ = readerDB.Close() }) + + for range n { + var count int + if err := readerDB.QueryRowContext(context.Background(), "SELECT COUNT(*) FROM events").Scan(&count); err != nil { + t.Fatalf("concurrent read during writes: %v", err) + } + } + + if err := <-writeErrs; err != nil { + t.Fatalf("AppendEvent: %v", err) + } +} + +func TestSession_AppendEvent_duplicateID(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("dup-event") + seq1, err := sess.AppendEvent(context.Background(), ev) + if err != nil { + t.Fatalf("AppendEvent (first): %v", err) + } + if seq1 != 1 { + t.Fatalf("first sequence = %d, want 1", seq1) + } + + if _, err := sess.AppendEvent(context.Background(), ev); !errors.Is(err, ErrDuplicateEventID) { + t.Fatalf("AppendEvent (duplicate id) err = %v, want ErrDuplicateEventID", err) + } + + // The failed duplicate insert's transaction was rolled back entirely, + // so it must not have consumed a sequence value — the next append + // gets 2, not 3. + seq2, err := sess.AppendEvent(context.Background(), testEvent("not-a-dup")) + if err != nil { + t.Fatalf("AppendEvent (second): %v", err) + } + if seq2 != 2 { + t.Errorf("sequence after a failed duplicate append = %d, want 2 (no gap)", seq2) + } +} + +func TestSession_AppendEvent_invalidKind(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("evt-1") + ev.Kind = kernelv1.EventKind_EVENT_KIND_UNSPECIFIED + if _, err := sess.AppendEvent(context.Background(), ev); !errors.Is(err, ErrInvalidKind) { + t.Fatalf("AppendEvent (unspecified kind) err = %v, want ErrInvalidKind", err) + } +} + +func TestSession_AppendEvent_missingProducer(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("evt-1") + ev.Producer = nil + if _, err := sess.AppendEvent(context.Background(), ev); err == nil { + t.Fatal("AppendEvent (nil producer) = nil error, want error") + } +} + +// TestSession_appendEventTx_extraFailureRollsBackEvent tests the +// same-tx-atomicity mechanism AppendMessage and AppendPlan both build on +// directly: if the "extra" step (cost_ledger or plan_items insert) fails, +// the event row itself must not exist either — the whole append is one +// transaction, not "insert the event, then try to insert the rest." +func TestSession_appendEventTx_extraFailureRollsBackEvent(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("rolled-back") + boom := errors.New("boom") + _, err := sess.appendEventTx(context.Background(), ev, func(context.Context, *sql.Tx, int64) error { + return boom + }) + if !errors.Is(err, boom) { + t.Fatalf("appendEventTx err = %v, want wrapping %v", err, boom) + } + + var count int + if scanErr := sess.db.QueryRowContext(context.Background(), "SELECT COUNT(*) FROM events WHERE id = ?", ev.ID).Scan(&count); scanErr != nil { + t.Fatalf("query events: %v", scanErr) + } + if count != 0 { + t.Error("events row exists after a failed same-tx append, want it rolled back") + } +} + +func TestSession_AppendMessage(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("msg-event") + ev.Kind = kernelv1.EventKind_EVENT_KIND_MESSAGE + cost := CostEntry{ + ProviderName: "anthropic", + ModelID: "claude-x", + InputTokens: 100, + OutputTokens: 50, + CacheWriteTokens: 5, + CacheReadTokens: 10, + CostUSD: 0.0123, + } + + seq, err := sess.AppendMessage(context.Background(), ev, cost) + if err != nil { + t.Fatalf("AppendMessage: %v", err) + } + + var ( + gotEventSeq int64 + gotProvider, gotModel string + gotInput, gotOutput, gotCacheWrite, gotCacheRead int64 + gotCost float64 + ) + row := sess.db.QueryRowContext(context.Background(), + "SELECT event_sequence, provider_name, model_id, input_tokens, output_tokens, cache_write_tokens, cache_read_tokens, cost_usd FROM cost_ledger WHERE event_sequence = ?", seq) + if err := row.Scan(&gotEventSeq, &gotProvider, &gotModel, &gotInput, &gotOutput, &gotCacheWrite, &gotCacheRead, &gotCost); err != nil { + t.Fatalf("query cost_ledger: %v", err) + } + if gotEventSeq != seq { + t.Errorf("event_sequence = %d, want %d", gotEventSeq, seq) + } + if gotProvider != cost.ProviderName || gotModel != cost.ModelID { + t.Errorf("provider/model = (%q, %q), want (%q, %q)", gotProvider, gotModel, cost.ProviderName, cost.ModelID) + } + if gotInput != cost.InputTokens || gotOutput != cost.OutputTokens || gotCacheWrite != cost.CacheWriteTokens || gotCacheRead != cost.CacheReadTokens { + t.Errorf("token counts = (%d,%d,%d,%d), want (%d,%d,%d,%d)", gotInput, gotOutput, gotCacheWrite, gotCacheRead, cost.InputTokens, cost.OutputTokens, cost.CacheWriteTokens, cost.CacheReadTokens) + } + if gotCost != cost.CostUSD { + t.Errorf("cost_usd = %v, want %v", gotCost, cost.CostUSD) + } +} + +func TestSession_AppendPlan(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("plan-event") + ev.Kind = kernelv1.EventKind_EVENT_KIND_PLAN + items := []PlanItem{ + {TurnID: "t1", ToolCallID: "c1", ProviderName: "p", ToolName: "read_file", Decision: planv1.PlanDecision_PLAN_DECISION_ALLOW, DecidedBy: "policy"}, + {TurnID: "t1", ToolCallID: "c2", ProviderName: "p", ToolName: "write_file", Decision: planv1.PlanDecision_PLAN_DECISION_ASK, DecidedBy: "operator"}, + } + + seq, err := sess.AppendPlan(context.Background(), ev, items) + if err != nil { + t.Fatalf("AppendPlan: %v", err) + } + + rows, err := sess.db.QueryContext(context.Background(), "SELECT tool_name, decision, decided_by FROM plan_items WHERE event_sequence = ? ORDER BY sequence", seq) + if err != nil { + t.Fatalf("query plan_items: %v", err) + } + defer rows.Close() + + var got []PlanItem + for rows.Next() { + var toolName, decisionText, decidedBy string + if err := rows.Scan(&toolName, &decisionText, &decidedBy); err != nil { + t.Fatalf("scan: %v", err) + } + decision, err := decodePlanDecision(decisionText) + if err != nil { + t.Fatalf("decodePlanDecision(%q): %v", decisionText, err) + } + got = append(got, PlanItem{ToolName: toolName, Decision: decision, DecidedBy: decidedBy}) + } + if err := rows.Err(); err != nil { + t.Fatalf("rows: %v", err) + } + + if len(got) != len(items) { + t.Fatalf("plan_items rows = %d, want %d", len(got), len(items)) + } + for i, item := range items { + if got[i].ToolName != item.ToolName || got[i].Decision != item.Decision || got[i].DecidedBy != item.DecidedBy { + t.Errorf("plan_items[%d] = %+v, want %+v", i, got[i], item) + } + } +} + +func TestSession_AppendPlan_invalidDecisionRollsBackEvent(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + ev := testEvent("plan-event") + ev.Kind = kernelv1.EventKind_EVENT_KIND_PLAN + items := []PlanItem{ + {TurnID: "t1", ToolCallID: "c1", ProviderName: "p", ToolName: "read_file", Decision: planv1.PlanDecision_PLAN_DECISION_UNSPECIFIED, DecidedBy: "policy"}, + } + + if _, err := sess.AppendPlan(context.Background(), ev, items); !errors.Is(err, ErrInvalidDecision) { + t.Fatalf("AppendPlan (invalid decision) err = %v, want ErrInvalidDecision", err) + } + + var count int + if scanErr := sess.db.QueryRowContext(context.Background(), "SELECT COUNT(*) FROM events WHERE id = ?", ev.ID).Scan(&count); scanErr != nil { + t.Fatalf("query events: %v", scanErr) + } + if count != 0 { + t.Error("events row exists after AppendPlan's plan_items insert failed, want it rolled back") + } +} + +func TestSession_SetStatus_roundTrip(t *testing.T) { + t.Parallel() + st := newTestStore(t) + meta := testSessionMeta() + sess := createSession(t, st, meta) + + // Truncated to millisecond precision: formatTimestamp's on-disk layout + // is millisecond precision, so a nanosecond-precision time.Now() would + // never compare equal after the round trip. + ended := time.Now().UTC().Truncate(time.Millisecond) + if err := sess.SetStatus(context.Background(), sessionv1.SessionStatus_SESSION_STATUS_COMPLETED, &ended); err != nil { + t.Fatalf("SetStatus: %v", err) + } + + got, err := querySessionMeta(context.Background(), sess.db, meta.SessionID) + if err != nil { + t.Fatalf("querySessionMeta: %v", err) + } + if got.Status != sessionv1.SessionStatus_SESSION_STATUS_COMPLETED { + t.Errorf("Status = %v, want COMPLETED", got.Status) + } + if got.EndedAt == nil || !got.EndedAt.Equal(ended) { + t.Errorf("EndedAt = %v, want %v", got.EndedAt, ended) + } +} + +func TestSession_SetStatus_stillRunningLeavesEndedAtNil(t *testing.T) { + t.Parallel() + st := newTestStore(t) + meta := testSessionMeta() + sess := createSession(t, st, meta) + + if err := sess.SetStatus(context.Background(), sessionv1.SessionStatus_SESSION_STATUS_RUNNING, nil); err != nil { + t.Fatalf("SetStatus: %v", err) + } + + got, err := querySessionMeta(context.Background(), sess.db, meta.SessionID) + if err != nil { + t.Fatalf("querySessionMeta: %v", err) + } + if got.EndedAt != nil { + t.Errorf("EndedAt = %v, want nil", got.EndedAt) + } +} + +func TestSession_SetStatus_invalidStatus(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess := createSession(t, st, testSessionMeta()) + + if err := sess.SetStatus(context.Background(), sessionv1.SessionStatus_SESSION_STATUS_UNSPECIFIED, nil); err == nil { + t.Fatal("SetStatus (unspecified) = nil error, want error") + } +} + +func TestSession_SetStatus_notFound(t *testing.T) { + t.Parallel() + st := newTestStore(t) + meta := testSessionMeta() + sess := createSession(t, st, meta) + + if _, err := sess.db.ExecContext(context.Background(), "DELETE FROM session_meta WHERE session_id = ?", meta.SessionID); err != nil { + t.Fatalf("delete session_meta: %v", err) + } + + if err := sess.SetStatus(context.Background(), sessionv1.SessionStatus_SESSION_STATUS_COMPLETED, nil); !errors.Is(err, ErrNotFound) { + t.Fatalf("SetStatus (no row) err = %v, want ErrNotFound", err) + } +} + +func TestSession_Close_idempotent(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + + if err := sess.Close(); err != nil { + t.Fatalf("Close (first): %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close (second) = %v, want nil (idempotent)", err) + } +} + +func TestSession_errClosedAfterClose(t *testing.T) { + t.Parallel() + st := newTestStore(t) + sess, err := st.Create(context.Background(), testSessionMeta()) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := sess.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + if _, err := sess.AppendEvent(context.Background(), testEvent("e1")); !errors.Is(err, ErrClosed) { + t.Errorf("AppendEvent after Close err = %v, want ErrClosed", err) + } + if _, err := sess.AppendMessage(context.Background(), testEvent("e2"), CostEntry{}); !errors.Is(err, ErrClosed) { + t.Errorf("AppendMessage after Close err = %v, want ErrClosed", err) + } + if _, err := sess.AppendPlan(context.Background(), testEvent("e3"), nil); !errors.Is(err, ErrClosed) { + t.Errorf("AppendPlan after Close err = %v, want ErrClosed", err) + } + if err := sess.SetStatus(context.Background(), sessionv1.SessionStatus_SESSION_STATUS_COMPLETED, nil); !errors.Is(err, ErrClosed) { + t.Errorf("SetStatus after Close err = %v, want ErrClosed", err) + } +} diff --git a/internal/statebackend/sessionid.go b/internal/statebackend/sessionid.go new file mode 100644 index 0000000..c003b6b --- /dev/null +++ b/internal/statebackend/sessionid.go @@ -0,0 +1,43 @@ +package statebackend + +import ( + "crypto/rand" + "fmt" + "sync" + "time" + + "github.com/oklog/ulid/v2" +) + +var ( + // mu guards monotonic to ensure concurrent ULID generation produces unique IDs. + mu sync.Mutex + // monotonic is a ULID entropy source that produces monotonically increasing ULIDs + // when called from the same millisecond, ensuring uniqueness across concurrent calls. + monotonic = ulid.Monotonic(rand.Reader, 0) +) + +// NewSessionID generates a new session ID as a ULID with the given timestamp. +// The ULID is in canonical Crockford base32 (uppercase, 26 characters), +// making session IDs sortable chronologically by filename alone. +func NewSessionID(t time.Time) string { + mu.Lock() + ms := ulid.Timestamp(t) + id, _ := ulid.New(ms, monotonic) + mu.Unlock() + return id.String() +} + +// ValidateSessionID returns an error if id is not a strictly valid canonical ULID. +// It rejects lowercase ULIDs, invalid characters, and values outside the ULID range. +func ValidateSessionID(id string) error { + parsed, err := ulid.ParseStrict(id) + if err != nil { + return fmt.Errorf("statebackend: %w", err) + } + // ParseStrict validates the format; ensure round-trip matches (catches lowercase). + if parsed.String() != id { + return fmt.Errorf("statebackend: session id must be canonical uppercase") + } + return nil +} diff --git a/internal/statebackend/sessionid_fuzz_test.go b/internal/statebackend/sessionid_fuzz_test.go new file mode 100644 index 0000000..88581a3 --- /dev/null +++ b/internal/statebackend/sessionid_fuzz_test.go @@ -0,0 +1,39 @@ +package statebackend + +import ( + "testing" + "time" + + "github.com/oklog/ulid/v2" +) + +// FuzzValidateSessionID exercises ValidateSessionID against arbitrary strings, +// asserting: +// 1. Any generated session ID must always validate with a nil error. +// 2. No panic on arbitrary attacker-controlled input. +// 3. If ValidateSessionID returns nil, ulid.ParseStrict must also succeed +// and String() must round-trip exactly to the input. +func FuzzValidateSessionID(f *testing.F) { + // Add seed examples. + f.Add(NewSessionID(time.Now())) + f.Add(NewSessionID(time.Unix(0, 0))) + f.Add("") + f.Add("invalid") + f.Add("01ARZ3NDEKTSV4RRFFQ69G5FAV") + + f.Fuzz(func(t *testing.T, input string) { + // Property 2: never panic on arbitrary input. + err := ValidateSessionID(input) + + // Property 3: if validation succeeds, ensure round-trip consistency. + if err == nil { + parsed, parseErr := ulid.ParseStrict(input) + if parseErr != nil { + t.Fatalf("ValidateSessionID(%q) returned nil but ParseStrict failed: %v", input, parseErr) + } + if parsed.String() != input { + t.Fatalf("ValidateSessionID(%q) succeeded but String() round-trip mismatch: %q", input, parsed.String()) + } + } + }) +} diff --git a/internal/statebackend/sessionid_test.go b/internal/statebackend/sessionid_test.go new file mode 100644 index 0000000..4e7b3ab --- /dev/null +++ b/internal/statebackend/sessionid_test.go @@ -0,0 +1,141 @@ +package statebackend + +import ( + "regexp" + "sort" + "strings" + "sync" + "testing" + "time" + + "github.com/oklog/ulid/v2" +) + +func TestNewSessionID(t *testing.T) { + t.Parallel() + now := time.Now() + + tests := []struct { + name string + t time.Time + }{ + {"zero time", time.Time{}}, + {"unix epoch", time.Unix(0, 0)}, + {"now", now}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + id := NewSessionID(tt.t) + + // Must be exactly 26 characters. + if len(id) != 26 { + t.Errorf("NewSessionID length = %d, want 26", len(id)) + } + + // Must be uppercase Crockford base32 only. + if !regexp.MustCompile(`^[0-7][0-9A-Z]{25}$`).MatchString(id) { + t.Errorf("NewSessionID format invalid: %q", id) + } + }) + } +} + +func TestNewSessionIDChronologicalOrder(t *testing.T) { + t.Parallel() + + t1 := time.Date(2024, 1, 1, 10, 0, 0, 0, time.UTC) + t2 := t1.Add(1 * time.Second) + t3 := t2.Add(1 * time.Second) + + id1 := NewSessionID(t1) + id2 := NewSessionID(t2) + id3 := NewSessionID(t3) + + ids := []string{id3, id1, id2} + sort.Strings(ids) + + if ids[0] != id1 || ids[1] != id2 || ids[2] != id3 { + t.Errorf("Chronological sort failed: %v", ids) + } +} + +func TestNewSessionIDConcurrentUniqueness(t *testing.T) { + t.Parallel() + + const goroutines = 100 + ids := make([]string, goroutines) + var wg sync.WaitGroup + wg.Add(goroutines) + + now := time.Now() + + for i := range goroutines { + go func(idx int) { + defer wg.Done() + ids[idx] = NewSessionID(now) + }(i) + } + wg.Wait() + + // All IDs must be unique. + seen := make(map[string]bool) + for _, id := range ids { + if seen[id] { + t.Errorf("Duplicate ID generated: %q", id) + break + } + seen[id] = true + } +} + +func TestValidateSessionID(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + id string + wantErr bool + }{ + {"generated id is valid", NewSessionID(time.Now()), false}, + {"empty string", "", true}, + {"too short", "0123456789", true}, + {"too long", "0123456789ABCDEFGHIJKLMNOPQRST", true}, + {"lowercase rejected", strings.ToLower(NewSessionID(time.Now())), true}, + {"invalid chars", "01234567890ABCDEFGHIJKLMNO!!!", true}, + {"overflow (8 prefix)", "8123456789ABCDEFGHIJKLMNOPQR", true}, + {"all zeros", "00000000000000000000000000", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + err := ValidateSessionID(tt.id) + if (err != nil) != tt.wantErr { + t.Errorf("ValidateSessionID(%q) error = %v, wantErr %v", tt.id, err, tt.wantErr) + } + }) + } +} + +func TestValidateSessionIDRoundTrip(t *testing.T) { + t.Parallel() + + // Any generated ID must pass validation and round-trip. + for range 10 { + id := NewSessionID(time.Now()) + if err := ValidateSessionID(id); err != nil { + t.Errorf("Generated ID failed validation: %q, %v", id, err) + } + + // Ensure it round-trips through ulid.ParseStrict. + parsed, err := ulid.ParseStrict(id) + if err != nil { + t.Errorf("Generated ID failed ParseStrict: %q, %v", id, err) + } + if parsed.String() != id { + t.Errorf("Round-trip mismatch: %q -> %q", id, parsed.String()) + } + } +} diff --git a/internal/statebackend/statebackend.go b/internal/statebackend/statebackend.go new file mode 100644 index 0000000..7770d9f --- /dev/null +++ b/internal/statebackend/statebackend.go @@ -0,0 +1,536 @@ +package statebackend + +import ( + "context" + "database/sql" + "errors" + "fmt" + "io/fs" + "log/slog" + "os" + "path/filepath" + "sort" + "strings" + "sync/atomic" + "time" + + _ "modernc.org/sqlite" // registers the "sqlite" database/sql driver + + "github.com/pluggableharness/agent/internal/telemetry" + "github.com/pluggableharness/agent/internal/telemetry/drivers/noop" + sessionv1 "github.com/pluggableharness/agent/pkg/session/proto/v1" +) + +// timestampLayout is the RFC 3339, UTC, millisecond-precision layout every +// TEXT timestamp column in the schema is stored and parsed with +// (docs/specifications/state-backend.md#ordering--concurrency: timestamps +// are display-only, never ordering-authoritative, but still need a single +// fixed format for round-tripping). +const timestampLayout = "2006-01-02T15:04:05.000Z07:00" + +// formatTimestamp renders t as UTC RFC 3339 with millisecond precision, the +// canonical on-disk form for every TEXT timestamp column in the schema. +func formatTimestamp(t time.Time) string { + return t.UTC().Format(timestampLayout) +} + +// parseTimestamp parses a TEXT timestamp column back into a time.Time, +// inverse of formatTimestamp. +func parseTimestamp(s string) (time.Time, error) { + return time.Parse(timestampLayout, s) +} + +// sessionStatusText maps SessionStatus to the exact lowercase snake_case +// text docs/specifications/state-backend.md#session_meta's status column +// documents ("running | completed | error_max_turns | error_max_budget_usd +// | error_max_wall_clock | cancelled | failed") — the wire enum's own +// SCREAMING_SNAKE_CASE String() is not what gets stored. +var sessionStatusText = map[sessionv1.SessionStatus]string{ + sessionv1.SessionStatus_SESSION_STATUS_RUNNING: "running", + sessionv1.SessionStatus_SESSION_STATUS_COMPLETED: "completed", + sessionv1.SessionStatus_SESSION_STATUS_ERROR_MAX_TURNS: "error_max_turns", + sessionv1.SessionStatus_SESSION_STATUS_ERROR_MAX_BUDGET_USD: "error_max_budget_usd", + sessionv1.SessionStatus_SESSION_STATUS_ERROR_MAX_WALL_CLOCK: "error_max_wall_clock", + sessionv1.SessionStatus_SESSION_STATUS_CANCELLED: "cancelled", + sessionv1.SessionStatus_SESSION_STATUS_FAILED: "failed", +} + +// sessionTextStatus is sessionStatusText inverted, built once at package +// init time from sessionStatusText itself so the two can never drift. +var sessionTextStatus = func() map[string]sessionv1.SessionStatus { + m := make(map[string]sessionv1.SessionStatus, len(sessionStatusText)) + for status, text := range sessionStatusText { + m[text] = status + } + return m +}() + +// encodeSessionStatus renders status as its stored TEXT representation. +// SESSION_STATUS_UNSPECIFIED and any unrecognized value are rejected — like +// EventKind's zero value (docs/specifications/state-backend.md#the-kind-enum), +// SessionStatus's zero value MUST NOT ever be persisted. +func encodeSessionStatus(status sessionv1.SessionStatus) (string, error) { + text, ok := sessionStatusText[status] + if !ok { + return "", fmt.Errorf("statebackend: session status %v has no stored representation", status) + } + return text, nil +} + +// decodeSessionStatus is the inverse of encodeSessionStatus, used when +// reading a session_meta row back. +func decodeSessionStatus(text string) (sessionv1.SessionStatus, error) { + status, ok := sessionTextStatus[text] + if !ok { + return sessionv1.SessionStatus_SESSION_STATUS_UNSPECIFIED, fmt.Errorf("statebackend: unrecognized session status %q", text) + } + return status, nil +} + +// SessionMeta mirrors the session_meta table's columns +// (docs/specifications/state-backend.md#session_meta) — the one table that +// isn't append-only, updated in place as a session progresses. +// ParentSessionID is empty for a root session (the column's SQL NULL); +// EndedAt is nil while the session is still running (the column's SQL +// NULL). +type SessionMeta struct { + // SessionID matches the filename stem: a canonical ULID (NewSessionID). + SessionID string + // ParentSessionID is empty for a root session. + ParentSessionID string + // Profile is the agent profile this session was started with. + Profile string + // Status is the session's current lifecycle state. + Status sessionv1.SessionStatus + // Depth is the cached depth-budget value + // (docs/specifications/agent-loop/subagents.md#depth-limits), avoiding + // a parent-chain walk when scanning. + Depth int + // StartedAt is when the session began. + StartedAt time.Time + // EndedAt is nil while the session is still running. + EndedAt *time.Time +} + +// Session is an open handle to one session's sqlite file: the *sql.DB +// backing it, its identifying ids, and the logger/telemetry provider it +// instruments through (inherited from the Store that opened it). Append and +// query methods live in session.go (event/cost/plan-item writes) and +// query.go (Stage 3's replay reads). +type Session struct { + id string + db *sql.DB + path string + logger *slog.Logger + telemetry *telemetry.Provider + + // closed is set by Close so every write method can reject a call made + // after it with ErrClosed instead of surfacing a raw database/sql + // error — see session.go. + closed atomic.Bool +} + +// ID returns the session's ULID. +func (s *Session) ID() string { + return s.id +} + +// Store manages the directory of per-session sqlite files described by +// docs/specifications/state-backend.md#file-layout +// ($XDG_STATE_HOME/agent/sessions/.sqlite). It is not itself a +// database handle — each Session opened through it owns its own *sql.DB. +type Store struct { + dir string + clock func() time.Time + logger *slog.Logger + telemetry *telemetry.Provider +} + +// Option configures a Store constructed by NewStore. +type Option func(*Store) + +// WithClock overrides the Store's source of the current time, used to +// default SessionMeta.StartedAt when Create is called without one. Tests +// use this for a deterministic clock; production code leaves it unset, +// defaulting to time.Now. +func WithClock(clock func() time.Time) Option { + return func(s *Store) { + if clock != nil { + s.clock = clock + } + } +} + +// WithLogger sets the *slog.Logger the Store logs through. A nil logger (or +// omitting this option) leaves the default of slog.Default(). +func WithLogger(logger *slog.Logger) Option { + return func(s *Store) { + if logger != nil { + s.logger = logger + } + } +} + +// WithTelemetry sets the *telemetry.Provider the Store and every Session it +// opens instrument through (internal/telemetry/span.go's +// StartStateBackend* helpers). Omitting this option (or passing nil) +// leaves the default: a Provider with every signal disabled, so New wires +// OTel's own no-op tracer/meter/logger providers directly +// (internal/telemetry.New's documented behavior for a disabled signal) — +// the instrumentation code path still runs on every call, at effectively +// zero cost, rather than being conditionally skipped. +func WithTelemetry(prov *telemetry.Provider) Option { + return func(s *Store) { + if prov != nil { + s.telemetry = prov + } + } +} + +// defaultTelemetryProvider builds the Provider a Store falls back to when +// WithTelemetry isn't supplied. Every signal is disabled, so +// telemetry.New never calls into the noop.Backend passed here at all — it +// exists only to satisfy New's non-nil Backend requirement. Constructing +// this at NewStore time (rather than propagating a caller context, which +// NewStore's signature has none of) is this package's one ingress-style +// use of context.Background(), the same carve-out go-style.md gives +// "main, an HTTP handler boundary, a message-consume loop": NewStore is +// this package's construction entry point. +func defaultTelemetryProvider() (*telemetry.Provider, error) { + return telemetry.New(context.Background(), telemetry.Config{}, noop.New(), nil) +} + +// NewStore returns a Store rooted at dir, creating it (mode 0700) if it +// does not already exist. dir is expected to be +// $XDG_STATE_HOME/agent/sessions per +// docs/specifications/state-backend.md#file-layout, but NewStore itself +// takes the resolved path rather than reading XDG env vars — that +// resolution belongs to the caller doing kernel-wide path setup. +func NewStore(dir string, opts ...Option) (*Store, error) { + if dir == "" { + return nil, fmt.Errorf("statebackend: new store: dir is required") + } + if err := os.MkdirAll(dir, 0o700); err != nil { + return nil, fmt.Errorf("statebackend: new store: %w", err) + } + + st := &Store{dir: dir, clock: time.Now, logger: slog.Default()} + for _, opt := range opts { + opt(st) + } + if st.telemetry == nil { + prov, err := defaultTelemetryProvider() + if err != nil { + return nil, fmt.Errorf("statebackend: new store: %w", err) + } + st.telemetry = prov + } + return st, nil +} + +// sessionPath returns the on-disk path for sessionID's file, per +// docs/specifications/state-backend.md#file-layout. +func (st *Store) sessionPath(sessionID string) string { + return filepath.Join(st.dir, sessionID+".sqlite") +} + +// openDB opens the sqlite file at path for exclusive, sole-writer access +// (docs/specifications/state-backend.md#ordering--concurrency: the kernel +// is the only writer to any given session's file). The returned *sql.DB is +// capped at one open connection, and WAL mode plus foreign key enforcement +// are set immediately after connecting. +// +// foreign_keys is requested via the DSN's _pragma parameter +// (modernc.org/sqlite-specific) rather than a one-time PRAGMA exec: it's a +// per-connection setting, and a PRAGMA exec only applies to the connection +// live at the time it runs. If database/sql ever replaces the pooled +// connection (e.g. after a driver-level ErrBadConn), a same-process PRAGMA +// wouldn't follow it to the replacement, silently reverting to +// foreign_keys=OFF — the DSN parameter is applied by the driver on every +// new connection it opens, including that replacement. journal_mode stays +// a PRAGMA exec deliberately: WAL is a durable property of the database +// file itself (recorded in the file header), not a per-connection +// setting, so a replacement connection sees it automatically. +func openDB(ctx context.Context, path string) (*sql.DB, error) { + db, err := sql.Open("sqlite", path+"?_pragma=foreign_keys(1)") + if err != nil { + return nil, fmt.Errorf("statebackend: open %s: %w", path, err) + } + db.SetMaxOpenConns(1) + + if _, err := db.ExecContext(ctx, "PRAGMA journal_mode = WAL"); err != nil { + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: set journal_mode: %w", path, err) + } + return db, nil +} + +// Create creates a new session file for meta.SessionID +// (docs/specifications/state-backend.md#file-layout), applies the +// five-table schema, stamps PRAGMA user_version, and inserts meta as the +// initial session_meta row. meta.SessionID MUST already be a valid, +// caller-generated ULID (see NewSessionID) — Create does not generate one +// itself. If meta.StartedAt is zero, it defaults to the Store's clock. Any +// failure after the file is created (a bad schema, an invalid +// meta.Status, ...) removes the partial file rather than leaving it behind +// to block a retry with the same session ID. +func (st *Store) Create(ctx context.Context, meta SessionMeta) (_ *Session, err error) { + ctx, span := st.telemetry.StartStateBackendSessionCreate(ctx, meta.SessionID) + defer func() { telemetry.EndSpan(span, err) }() + + if err = ValidateSessionID(meta.SessionID); err != nil { + err = fmt.Errorf("statebackend: create: %w", err) + return nil, err + } + if meta.StartedAt.IsZero() { + meta.StartedAt = st.clock() + } + + path := st.sessionPath(meta.SessionID) + st.logger.DebugContext(ctx, "statebackend: creating session", "session_id", meta.SessionID) + + // #nosec G304 -- path is built from meta.SessionID, which ValidateSessionID above already rejected unless it's a canonical ULID; it is never attacker-controlled arbitrary input. + f, openErr := os.OpenFile(path, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0o600) + if openErr != nil { + if errors.Is(openErr, fs.ErrExist) { + err = fmt.Errorf("statebackend: create %s: session file already exists", meta.SessionID) + return nil, err + } + err = fmt.Errorf("statebackend: create %s: %w", meta.SessionID, openErr) + return nil, err + } + if closeErr := f.Close(); closeErr != nil { + _ = os.Remove(path) + err = fmt.Errorf("statebackend: create %s: %w", meta.SessionID, closeErr) + return nil, err + } + + sess, popErr := st.populateCreatedFile(ctx, path, meta) + if popErr != nil { + _ = os.Remove(path) + err = fmt.Errorf("statebackend: create %s: %w", meta.SessionID, popErr) + return nil, err + } + return sess, nil +} + +// populateCreatedFile opens the just-created, still-empty file at path, +// applies the schema, and inserts meta's session_meta row. Split out of +// Create so every failure path shares one os.Remove(path) cleanup call. +func (st *Store) populateCreatedFile(ctx context.Context, path string, meta SessionMeta) (*Session, error) { + db, err := openDB(ctx, path) + if err != nil { + return nil, err + } + if err := initSchema(ctx, db); err != nil { + _ = db.Close() + return nil, err + } + if err := insertSessionMeta(ctx, db, meta); err != nil { + _ = db.Close() + return nil, err + } + return &Session{id: meta.SessionID, db: db, path: path, logger: st.logger, telemetry: st.telemetry}, nil +} + +// Open opens an existing session file for sessionID. It validates +// sessionID, then runs the corruption-recovery check +// (docs/specifications/state-backend.md#corruption-recovery — see +// integrity.go's checkIntegrity), then checks PRAGMA user_version before +// any other operation touches the file +// (docs/specifications/state-backend.md#schema-migration): newer than +// currentSchemaVersion returns ErrSchemaTooNew; older applies the ordered +// migrations slice. sessionID not found returns ErrNotFound. +func (st *Store) Open(ctx context.Context, sessionID string) (_ *Session, err error) { + ctx, span := st.telemetry.StartStateBackendSessionOpen(ctx, sessionID) + defer func() { telemetry.EndSpan(span, err) }() + + if err = ValidateSessionID(sessionID); err != nil { + err = fmt.Errorf("statebackend: open: %w", err) + return nil, err + } + st.logger.DebugContext(ctx, "statebackend: opening session", "session_id", sessionID) + + path := st.sessionPath(sessionID) + if _, statErr := os.Stat(path); statErr != nil { + if errors.Is(statErr, fs.ErrNotExist) { + err = fmt.Errorf("statebackend: open %s: %w", sessionID, ErrNotFound) + return nil, err + } + err = fmt.Errorf("statebackend: open %s: %w", sessionID, statErr) + return nil, err + } + + db, icErr := st.checkIntegrity(ctx, path, sessionID) + if icErr != nil { + err = fmt.Errorf("statebackend: open %s: %w", sessionID, icErr) + return nil, err + } + + version, verErr := readUserVersion(ctx, db) + if verErr != nil { + _ = db.Close() + err = fmt.Errorf("statebackend: open %s: %w", sessionID, verErr) + return nil, err + } + switch { + case version > currentSchemaVersion: + _ = db.Close() + err = fmt.Errorf("statebackend: open %s: schema version %d: %w", sessionID, version, ErrSchemaTooNew) + return nil, err + case version < currentSchemaVersion: + if migErr := applyMigrations(ctx, db, migrations, version, currentSchemaVersion); migErr != nil { + _ = db.Close() + err = fmt.Errorf("statebackend: open %s: %w", sessionID, migErr) + return nil, err + } + } + + return &Session{id: sessionID, db: db, path: path, logger: st.logger, telemetry: st.telemetry}, nil +} + +// List returns every session's session_meta row, ordered by session_id +// (docs/specifications/state-backend.md#cross-session-queries: there is no +// separate index, so this scans the store directory). Since session IDs are +// canonical ULIDs, sorting by session_id is equivalent to chronological +// order. +func (st *Store) List(ctx context.Context) ([]SessionMeta, error) { + return st.scan(ctx, func(SessionMeta) bool { return true }) +} + +// Children returns every session whose session_meta.parent_session_id is +// sessionID, ordered by session_id +// (docs/specifications/state-backend.md#live-vs-post-hoc-tree-walking: a +// post-hoc query, reconstructing the tree by scanning files). +func (st *Store) Children(ctx context.Context, sessionID string) ([]SessionMeta, error) { + if err := ValidateSessionID(sessionID); err != nil { + return nil, fmt.Errorf("statebackend: children: %w", err) + } + return st.scan(ctx, func(m SessionMeta) bool { return m.ParentSessionID == sessionID }) +} + +// scan reads every *.sqlite file's session_meta row directly (bypassing +// Open's schema-version-check-and-migrate path — a metadata scan MUST NOT +// have the side effect of migrating every file it merely lists), filters by +// keep, and returns the result ordered by session_id. A file whose name +// isn't a valid session ID is logged at WARN and skipped rather than +// failing the whole scan, since the sessions directory is not guaranteed to +// contain only session files (e.g. future *.sqlite.corrupt entries from +// Stage 3's recovery flow). +func (st *Store) scan(ctx context.Context, keep func(SessionMeta) bool) ([]SessionMeta, error) { + entries, err := os.ReadDir(st.dir) + if err != nil { + return nil, fmt.Errorf("statebackend: scan sessions: %w", err) + } + + metas := make([]SessionMeta, 0, len(entries)) + for _, entry := range entries { + if entry.IsDir() || filepath.Ext(entry.Name()) != ".sqlite" { + continue + } + sessionID := strings.TrimSuffix(entry.Name(), ".sqlite") + if err := ValidateSessionID(sessionID); err != nil { + st.logger.WarnContext(ctx, "statebackend: skipping non-session file", "file", entry.Name(), "err", err) + continue + } + + meta, err := st.readSessionMeta(ctx, sessionID) + if err != nil { + return nil, fmt.Errorf("statebackend: scan sessions: %s: %w", sessionID, err) + } + if keep(meta) { + metas = append(metas, meta) + } + } + + sort.Slice(metas, func(i, j int) bool { return metas[i].SessionID < metas[j].SessionID }) + return metas, nil +} + +// readSessionMeta opens sessionID's file directly (not through Open — see +// scan's comment) purely to read its single session_meta row. +func (st *Store) readSessionMeta(ctx context.Context, sessionID string) (SessionMeta, error) { + db, err := openDB(ctx, st.sessionPath(sessionID)) + if err != nil { + return SessionMeta{}, err + } + defer func() { _ = db.Close() }() + + return querySessionMeta(ctx, db, sessionID) +} + +// insertSessionMeta inserts meta as session_meta's single row for a +// freshly created session file. +func insertSessionMeta(ctx context.Context, db *sql.DB, meta SessionMeta) error { + statusText, err := encodeSessionStatus(meta.Status) + if err != nil { + return fmt.Errorf("statebackend: insert session_meta: %w", err) + } + + var parentSessionID any + if meta.ParentSessionID != "" { + parentSessionID = meta.ParentSessionID + } + var endedAt any + if meta.EndedAt != nil { + endedAt = formatTimestamp(*meta.EndedAt) + } + + const q = `INSERT INTO session_meta (session_id, parent_session_id, profile, status, depth, started_at, ended_at) VALUES (?, ?, ?, ?, ?, ?, ?)` + if _, err := db.ExecContext(ctx, q, + meta.SessionID, parentSessionID, meta.Profile, statusText, meta.Depth, + formatTimestamp(meta.StartedAt), endedAt, + ); err != nil { + return fmt.Errorf("statebackend: insert session_meta: %w", err) + } + return nil +} + +// querySessionMeta reads sessionID's session_meta row from db. +func querySessionMeta(ctx context.Context, db *sql.DB, sessionID string) (SessionMeta, error) { + const q = `SELECT session_id, parent_session_id, profile, status, depth, started_at, ended_at FROM session_meta WHERE session_id = ?` + row := db.QueryRowContext(ctx, q, sessionID) + return scanSessionMeta(row) +} + +// scanSessionMeta decodes one session_meta row, translating its stored TEXT +// status back to sessionv1.SessionStatus and its TEXT timestamps back to +// time.Time. +func scanSessionMeta(row *sql.Row) (SessionMeta, error) { + var ( + meta SessionMeta + parentSessionID sql.NullString + statusText string + startedAtText string + endedAtText sql.NullString + ) + if err := row.Scan(&meta.SessionID, &parentSessionID, &meta.Profile, &statusText, &meta.Depth, &startedAtText, &endedAtText); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return SessionMeta{}, fmt.Errorf("statebackend: session_meta: %w", ErrNotFound) + } + return SessionMeta{}, fmt.Errorf("statebackend: session_meta: %w", err) + } + meta.ParentSessionID = parentSessionID.String + + status, err := decodeSessionStatus(statusText) + if err != nil { + return SessionMeta{}, fmt.Errorf("statebackend: session_meta: %w", err) + } + meta.Status = status + + startedAt, err := parseTimestamp(startedAtText) + if err != nil { + return SessionMeta{}, fmt.Errorf("statebackend: session_meta: started_at: %w", err) + } + meta.StartedAt = startedAt + + if endedAtText.Valid { + endedAt, err := parseTimestamp(endedAtText.String) + if err != nil { + return SessionMeta{}, fmt.Errorf("statebackend: session_meta: ended_at: %w", err) + } + meta.EndedAt = &endedAt + } + + return meta, nil +} diff --git a/internal/statebackend/statebackend_test.go b/internal/statebackend/statebackend_test.go new file mode 100644 index 0000000..db9f39d --- /dev/null +++ b/internal/statebackend/statebackend_test.go @@ -0,0 +1,712 @@ +package statebackend + +import ( + "bytes" + "context" + "database/sql" + "errors" + "io/fs" + "log/slog" + "os" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + sessionv1 "github.com/pluggableharness/agent/pkg/session/proto/v1" +) + +// fixedClock returns a Store clock (WithClock) that always reports t. +func fixedClock(t time.Time) func() time.Time { + return func() time.Time { return t } +} + +// newTestStore returns a Store rooted at a fresh t.TempDir(). +func newTestStore(t *testing.T, opts ...Option) *Store { + t.Helper() + st, err := NewStore(t.TempDir(), opts...) + if err != nil { + t.Fatalf("NewStore: %v", err) + } + return st +} + +// createSession creates a session in st and registers its Close via +// t.Cleanup, failing the test immediately on error. +func createSession(t *testing.T, st *Store, meta SessionMeta) *Session { + t.Helper() + sess, err := st.Create(context.Background(), meta) + if err != nil { + t.Fatalf("Create %s: %v", meta.SessionID, err) + } + t.Cleanup(func() { _ = sess.Close() }) + return sess +} + +func TestNewStore(t *testing.T) { + t.Parallel() + + t.Run("creates directory with 0700", func(t *testing.T) { + t.Parallel() + dir := filepath.Join(t.TempDir(), "sessions") + if _, err := NewStore(dir); err != nil { + t.Fatalf("NewStore: %v", err) + } + info, err := os.Stat(dir) + if err != nil { + t.Fatalf("Stat: %v", err) + } + if !info.IsDir() { + t.Fatal("dir is not a directory") + } + // Unix permission bits aren't meaningful on Windows (ACL-based + // instead) — os.Mkdir's mode argument is effectively ignored + // there, so Stat reads back something like 0777 regardless of + // what NewStore requested. + if runtime.GOOS != "windows" { + if perm := info.Mode().Perm(); perm != 0o700 { + t.Errorf("dir perm = %o, want %o", perm, 0o700) + } + } + }) + + t.Run("empty dir rejected", func(t *testing.T) { + t.Parallel() + if _, err := NewStore(""); err == nil { + t.Fatal("NewStore(\"\") = nil error, want error") + } + }) +} + +func TestStore_Create(t *testing.T) { + t.Parallel() + + fixedNow := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) + st := newTestStore(t, WithClock(fixedClock(fixedNow))) + ctx := context.Background() + + id := NewSessionID(fixedNow) + meta := SessionMeta{ + SessionID: id, + Profile: "default", + Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, + Depth: 2, + // StartedAt deliberately left zero to exercise the Store-clock default. + } + + sess := createSession(t, st, meta) + if sess.ID() != id { + t.Errorf("ID() = %q, want %q", sess.ID(), id) + } + + path := st.sessionPath(id) + info, err := os.Stat(path) + if err != nil { + t.Fatalf("Stat: %v", err) + } + // Unix permission bits aren't meaningful on Windows — see TestNewStore's + // equivalent guard. + if runtime.GOOS != "windows" { + if perm := info.Mode().Perm(); perm != 0o600 { + t.Errorf("file perm = %o, want %o", perm, 0o600) + } + } + + version, err := readUserVersion(ctx, sess.db) + if err != nil { + t.Fatalf("readUserVersion: %v", err) + } + if version != currentSchemaVersion { + t.Errorf("user_version = %d, want %d", version, currentSchemaVersion) + } + + got, err := querySessionMeta(ctx, sess.db, id) + if err != nil { + t.Fatalf("querySessionMeta: %v", err) + } + if got.SessionID != id { + t.Errorf("SessionID = %q, want %q", got.SessionID, id) + } + if got.ParentSessionID != "" { + t.Errorf("ParentSessionID = %q, want empty", got.ParentSessionID) + } + if got.Profile != meta.Profile { + t.Errorf("Profile = %q, want %q", got.Profile, meta.Profile) + } + if got.Status != meta.Status { + t.Errorf("Status = %v, want %v", got.Status, meta.Status) + } + if got.Depth != meta.Depth { + t.Errorf("Depth = %d, want %d", got.Depth, meta.Depth) + } + if !got.StartedAt.Equal(fixedNow) { + t.Errorf("StartedAt = %v, want %v (Store clock default)", got.StartedAt, fixedNow) + } + if got.EndedAt != nil { + t.Errorf("EndedAt = %v, want nil", got.EndedAt) + } +} + +func TestStore_Create_endedAtRoundTrips(t *testing.T) { + t.Parallel() + st := newTestStore(t) + ctx := context.Background() + + started := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + ended := started.Add(5 * time.Minute) + id := NewSessionID(started) + meta := SessionMeta{ + SessionID: id, + Profile: "default", + Status: sessionv1.SessionStatus_SESSION_STATUS_COMPLETED, + StartedAt: started, + EndedAt: &ended, + } + + sess := createSession(t, st, meta) + got, err := querySessionMeta(ctx, sess.db, id) + if err != nil { + t.Fatalf("querySessionMeta: %v", err) + } + if got.EndedAt == nil { + t.Fatal("EndedAt = nil, want non-nil") + } + if !got.EndedAt.Equal(ended) { + t.Errorf("EndedAt = %v, want %v", *got.EndedAt, ended) + } +} + +func TestStore_Create_duplicate(t *testing.T) { + t.Parallel() + st := newTestStore(t) + id := NewSessionID(time.Now()) + meta := SessionMeta{SessionID: id, Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()} + + createSession(t, st, meta) + + if _, err := st.Create(context.Background(), meta); err == nil { + t.Fatal("Create (duplicate session id) = nil error, want error") + } +} + +func TestStore_Create_invalidSessionID(t *testing.T) { + t.Parallel() + st := newTestStore(t) + if _, err := st.Create(context.Background(), SessionMeta{SessionID: "not-a-ulid"}); err == nil { + t.Fatal("Create with invalid session id = nil error, want error") + } +} + +func TestStore_Create_invalidStatus(t *testing.T) { + t.Parallel() + st := newTestStore(t) + id := NewSessionID(time.Now()) + meta := SessionMeta{SessionID: id, Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_UNSPECIFIED, StartedAt: time.Now()} + if _, err := st.Create(context.Background(), meta); err == nil { + t.Fatal("Create with SESSION_STATUS_UNSPECIFIED = nil error, want error") + } + + // The partially-created file must not survive a failed Create — it + // would otherwise permanently block a retry with the same session ID + // (Create uses O_EXCL and treats an existing file as a real conflict). + if _, err := os.Stat(st.sessionPath(id)); !errors.Is(err, fs.ErrNotExist) { + t.Errorf("session file after failed Create: err = %v, want fs.ErrNotExist", err) + } +} + +func TestStore_Open(t *testing.T) { + t.Parallel() + st := newTestStore(t) + ctx := context.Background() + now := time.Now() + id := NewSessionID(now) + meta := SessionMeta{SessionID: id, Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: now} + + created, err := st.Create(ctx, meta) + if err != nil { + t.Fatalf("Create: %v", err) + } + if err := created.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + sess, err := st.Open(ctx, id) + if err != nil { + t.Fatalf("Open: %v", err) + } + t.Cleanup(func() { _ = sess.Close() }) + + if sess.ID() != id { + t.Errorf("ID() = %q, want %q", sess.ID(), id) + } + + got, err := querySessionMeta(ctx, sess.db, id) + if err != nil { + t.Fatalf("querySessionMeta: %v", err) + } + if got.Profile != meta.Profile { + t.Errorf("Profile = %q, want %q", got.Profile, meta.Profile) + } +} + +func TestStore_Open_notFound(t *testing.T) { + t.Parallel() + st := newTestStore(t) + _, err := st.Open(context.Background(), NewSessionID(time.Now())) + if !errors.Is(err, ErrNotFound) { + t.Errorf("Open (missing session) err = %v, want ErrNotFound", err) + } +} + +func TestStore_Open_invalidSessionID(t *testing.T) { + t.Parallel() + st := newTestStore(t) + if _, err := st.Open(context.Background(), "not-a-ulid"); err == nil { + t.Fatal("Open with invalid session id = nil error, want error") + } +} + +func TestStore_Open_schemaTooNew(t *testing.T) { + t.Parallel() + st := newTestStore(t) + id := NewSessionID(time.Now()) + path := st.sessionPath(id) + + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatalf("sql.Open: %v", err) + } + if _, err := db.ExecContext(context.Background(), "PRAGMA user_version = 999"); err != nil { + t.Fatalf("set user_version: %v", err) + } + if err := db.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + _, err = st.Open(context.Background(), id) + if !errors.Is(err, ErrSchemaTooNew) { + t.Errorf("Open (newer schema) err = %v, want ErrSchemaTooNew", err) + } +} + +func TestStore_Open_olderSchemaWithNoMigrationPath(t *testing.T) { + t.Parallel() + st := newTestStore(t) + id := NewSessionID(time.Now()) + path := st.sessionPath(id) + + // A bare file at user_version 0 (the pre-schema default): older than + // currentSchemaVersion, so Open must take the migration branch — which, + // since the real migrations slice is empty for schema version 1, has no + // path from 0 to 1 and must fail rather than silently proceed. + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatalf("sql.Open: %v", err) + } + // sql.Open is lazy — force the file into existence before Close. + if err := db.PingContext(context.Background()); err != nil { + t.Fatalf("Ping: %v", err) + } + if err := db.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + if _, err := st.Open(context.Background(), id); err == nil { + t.Fatal("Open (schema older than current, no migration path) = nil error, want error") + } +} + +func TestStore_ListAndChildren(t *testing.T) { + t.Parallel() + st := newTestStore(t) + ctx := context.Background() + + base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + + root := NewSessionID(base) + createSession(t, st, SessionMeta{SessionID: root, Profile: "root", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: base}) + + child1 := NewSessionID(base.Add(time.Millisecond)) + createSession(t, st, SessionMeta{SessionID: child1, ParentSessionID: root, Profile: "child", Status: sessionv1.SessionStatus_SESSION_STATUS_COMPLETED, StartedAt: base.Add(time.Millisecond)}) + + child2 := NewSessionID(base.Add(2 * time.Millisecond)) + createSession(t, st, SessionMeta{SessionID: child2, ParentSessionID: root, Profile: "child", Status: sessionv1.SessionStatus_SESSION_STATUS_FAILED, StartedAt: base.Add(2 * time.Millisecond)}) + + other := NewSessionID(base.Add(3 * time.Millisecond)) + createSession(t, st, SessionMeta{SessionID: other, Profile: "unrelated", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: base.Add(3 * time.Millisecond)}) + + all, err := st.List(ctx) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(all) != 4 { + t.Fatalf("List returned %d sessions, want 4", len(all)) + } + for i := 1; i < len(all); i++ { + if all[i-1].SessionID >= all[i].SessionID { + t.Errorf("List not ordered by session_id: %q before %q", all[i-1].SessionID, all[i].SessionID) + } + } + + children, err := st.Children(ctx, root) + if err != nil { + t.Fatalf("Children: %v", err) + } + if len(children) != 2 { + t.Fatalf("Children returned %d sessions, want 2", len(children)) + } + if children[0].SessionID != child1 || children[1].SessionID != child2 { + t.Errorf("Children = [%q, %q], want [%q, %q]", children[0].SessionID, children[1].SessionID, child1, child2) + } + for _, c := range children { + if c.ParentSessionID != root { + t.Errorf("child %q ParentSessionID = %q, want %q", c.SessionID, c.ParentSessionID, root) + } + } + + noChildren, err := st.Children(ctx, other) + if err != nil { + t.Fatalf("Children(other): %v", err) + } + if len(noChildren) != 0 { + t.Errorf("Children(other) = %d sessions, want 0", len(noChildren)) + } +} + +func TestStore_List_empty(t *testing.T) { + t.Parallel() + st := newTestStore(t) + all, err := st.List(context.Background()) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(all) != 0 { + t.Errorf("List = %d sessions, want 0", len(all)) + } +} + +func TestStore_Children_invalidSessionID(t *testing.T) { + t.Parallel() + st := newTestStore(t) + if _, err := st.Children(context.Background(), "not-a-ulid"); err == nil { + t.Fatal("Children with invalid session id = nil error, want error") + } +} + +func TestStore_scan_skipsNonSessionFiles(t *testing.T) { + t.Parallel() + st := newTestStore(t) + ctx := context.Background() + + id := NewSessionID(time.Now()) + createSession(t, st, SessionMeta{SessionID: id, Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()}) + + // A file that isn't a valid session ID's .sqlite (e.g. a future + // Stage 3 *.sqlite.corrupt sidecar) must be skipped, not fail the scan. + if err := os.WriteFile(filepath.Join(st.dir, "not-a-session.sqlite"), []byte("garbage"), 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if err := os.WriteFile(filepath.Join(st.dir, "README.txt"), []byte("hello"), 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + all, err := st.List(ctx) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(all) != 1 || all[0].SessionID != id { + t.Errorf("List = %v, want exactly [%q]", all, id) + } +} + +func TestSessionStatus_roundTrip(t *testing.T) { + t.Parallel() + + for status := range sessionStatusText { + text, err := encodeSessionStatus(status) + if err != nil { + t.Fatalf("encodeSessionStatus(%v): %v", status, err) + } + got, err := decodeSessionStatus(text) + if err != nil { + t.Fatalf("decodeSessionStatus(%q): %v", text, err) + } + if got != status { + t.Errorf("round trip %v -> %q -> %v, want %v", status, text, got, status) + } + } +} + +func TestEncodeSessionStatus_unspecifiedRejected(t *testing.T) { + t.Parallel() + if _, err := encodeSessionStatus(sessionv1.SessionStatus_SESSION_STATUS_UNSPECIFIED); err == nil { + t.Fatal("encodeSessionStatus(UNSPECIFIED) = nil error, want error") + } +} + +func TestDecodeSessionStatus_unrecognized(t *testing.T) { + t.Parallel() + if _, err := decodeSessionStatus("not_a_status"); err == nil { + t.Fatal("decodeSessionStatus(garbage) = nil error, want error") + } +} + +func TestNewStore_mkdirFails(t *testing.T) { + t.Parallel() + base := t.TempDir() + blocker := filepath.Join(base, "blocker") + if err := os.WriteFile(blocker, []byte("x"), 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + // blocker is a regular file, not a directory, so MkdirAll(blocker/sessions) + // must fail with ENOTDIR. + if _, err := NewStore(filepath.Join(blocker, "sessions")); err == nil { + t.Fatal("NewStore under a non-directory parent = nil error, want error") + } +} + +func TestStore_Create_openFileError(t *testing.T) { + t.Parallel() + if runtime.GOOS == "windows" { + // A chmod'd-read-only directory doesn't stop file creation on + // Windows (ACL-based permissions, not Unix mode bits) — the + // premise this test relies on doesn't hold there. + t.Skip("read-only-directory permission enforcement is Unix-specific") + } + + st := newTestStore(t) + if err := os.Chmod(st.dir, 0o500); err != nil { + t.Fatalf("Chmod: %v", err) + } + defer func() { _ = os.Chmod(st.dir, 0o700) }() + + meta := SessionMeta{SessionID: NewSessionID(time.Now()), Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()} + sess, err := st.Create(context.Background(), meta) + if err == nil { + // Belt-and-suspenders: if the read-only directory somehow didn't + // block creation (a platform surprise beyond the Windows case + // already skipped above), close the session immediately rather + // than leaving it open — an unclosed *sql.DB on a file under + // st.dir would otherwise wedge t.TempDir()'s own cleanup. + _ = sess.Close() + t.Fatal("Create in a read-only directory = nil error, want error") + } +} + +func TestOpenDB_directoryPathFails(t *testing.T) { + t.Parallel() + if _, err := openDB(context.Background(), t.TempDir()); err == nil { + t.Fatal("openDB(directory) = nil error, want error") + } +} + +// TestOpenDB_foreignKeysEnabled locks in the security-review fix: since +// foreign_keys is a per-connection setting, it's requested via openDB's +// DSN (_pragma=foreign_keys(1)) rather than a one-time PRAGMA exec, so a +// replacement pooled connection still gets it. This asserts the +// observable effect — PRAGMA foreign_keys reads back 1 — rather than +// inspecting the DSN string itself. +func TestOpenDB_foreignKeysEnabled(t *testing.T) { + t.Parallel() + path := filepath.Join(t.TempDir(), "test.sqlite") + db, err := openDB(context.Background(), path) + if err != nil { + t.Fatalf("openDB: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + + var enabled int + if err := db.QueryRowContext(context.Background(), "PRAGMA foreign_keys").Scan(&enabled); err != nil { + t.Fatalf("PRAGMA foreign_keys: %v", err) + } + if enabled != 1 { + t.Errorf("PRAGMA foreign_keys = %d, want 1", enabled) + } +} + +func TestPopulateCreatedFile_openDBFailure(t *testing.T) { + t.Parallel() + st := newTestStore(t) + dirAsFile := filepath.Join(st.dir, "not-a-file") + if err := os.Mkdir(dirAsFile, 0o700); err != nil { + t.Fatalf("Mkdir: %v", err) + } + + meta := SessionMeta{SessionID: NewSessionID(time.Now()), Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()} + if _, err := st.populateCreatedFile(context.Background(), dirAsFile, meta); err == nil { + t.Fatal("populateCreatedFile against a directory = nil error, want error") + } +} + +func TestPopulateCreatedFile_initSchemaFailure(t *testing.T) { + t.Parallel() + st := newTestStore(t) + path := filepath.Join(st.dir, "existing.sqlite") + + db, err := openDB(context.Background(), path) + if err != nil { + t.Fatalf("openDB: %v", err) + } + if err := initSchema(context.Background(), db); err != nil { + t.Fatalf("initSchema: %v", err) + } + if err := db.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + // path already carries the full schema; populateCreatedFile's own + // initSchema call must fail (tables already exist) rather than silently + // succeed against a file it didn't create. + meta := SessionMeta{SessionID: NewSessionID(time.Now()), Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()} + if _, err := st.populateCreatedFile(context.Background(), path, meta); err == nil { + t.Fatal("populateCreatedFile against an already-initialized file = nil error, want error") + } +} + +func TestStore_List_readDirError(t *testing.T) { + t.Parallel() + dir := t.TempDir() + st, err := NewStore(dir) + if err != nil { + t.Fatalf("NewStore: %v", err) + } + if err := os.RemoveAll(dir); err != nil { + t.Fatalf("RemoveAll: %v", err) + } + + if _, err := st.List(context.Background()); err == nil { + t.Fatal("List after the store directory was removed = nil error, want error") + } +} + +func TestStore_scan_corruptFileReturnsError(t *testing.T) { + t.Parallel() + st := newTestStore(t) + id := NewSessionID(time.Now()) + path := st.sessionPath(id) + + // A validly-named session file (per ValidateSessionID) whose content + // isn't a sqlite database at all — distinct from + // TestStore_scan_skipsNonSessionFiles, where the *filename* itself + // isn't a valid session ID and is skipped outright. This one must + // surface as an error, not be silently skipped. + if err := os.WriteFile(path, []byte("this is not a valid sqlite file"), 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + if _, err := st.List(context.Background()); err == nil { + t.Fatal("List with a corrupt session file = nil error, want error") + } +} + +func TestInsertSessionMeta_duplicateSessionIDFails(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + meta := SessionMeta{SessionID: "01ARZ3NDEKTSV4RRFFQ69G5FAV", Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()} + if err := insertSessionMeta(ctx, db, meta); err != nil { + t.Fatalf("insertSessionMeta: %v", err) + } + if err := insertSessionMeta(ctx, db, meta); err == nil { + t.Fatal("insertSessionMeta (duplicate session_id) = nil error, want error") + } +} + +func TestQuerySessionMeta_invalidStatus(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + const q = `INSERT INTO session_meta (session_id, parent_session_id, profile, status, depth, started_at, ended_at) VALUES (?, NULL, ?, ?, ?, ?, NULL)` + if _, err := db.ExecContext(ctx, q, "01ARZ3NDEKTSV4RRFFQ69G5FAV", "default", "not_a_status", 0, formatTimestamp(time.Now())); err != nil { + t.Fatalf("insert: %v", err) + } + + if _, err := querySessionMeta(ctx, db, "01ARZ3NDEKTSV4RRFFQ69G5FAV"); err == nil { + t.Fatal("querySessionMeta with invalid status = nil error, want error") + } +} + +func TestQuerySessionMeta_invalidStartedAt(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + const q = `INSERT INTO session_meta (session_id, parent_session_id, profile, status, depth, started_at, ended_at) VALUES (?, NULL, ?, ?, ?, ?, NULL)` + if _, err := db.ExecContext(ctx, q, "01ARZ3NDEKTSV4RRFFQ69G5FAV", "default", "running", 0, "not-a-timestamp"); err != nil { + t.Fatalf("insert: %v", err) + } + + if _, err := querySessionMeta(ctx, db, "01ARZ3NDEKTSV4RRFFQ69G5FAV"); err == nil { + t.Fatal("querySessionMeta with invalid started_at = nil error, want error") + } +} + +func TestQuerySessionMeta_invalidEndedAt(t *testing.T) { + t.Parallel() + db := openTestDB(t) + ctx := context.Background() + if err := initSchema(ctx, db); err != nil { + t.Fatalf("initSchema: %v", err) + } + + const q = `INSERT INTO session_meta (session_id, parent_session_id, profile, status, depth, started_at, ended_at) VALUES (?, NULL, ?, ?, ?, ?, ?)` + if _, err := db.ExecContext(ctx, q, "01ARZ3NDEKTSV4RRFFQ69G5FAV", "default", "running", 0, formatTimestamp(time.Now()), "not-a-timestamp"); err != nil { + t.Fatalf("insert: %v", err) + } + + if _, err := querySessionMeta(ctx, db, "01ARZ3NDEKTSV4RRFFQ69G5FAV"); err == nil { + t.Fatal("querySessionMeta with invalid ended_at = nil error, want error") + } +} + +func TestWithLogger(t *testing.T) { + t.Parallel() + var buf bytes.Buffer + logger := slog.New(slog.NewTextHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug})) + st := newTestStore(t, WithLogger(logger)) + + createSession(t, st, SessionMeta{SessionID: NewSessionID(time.Now()), Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()}) + + if !strings.Contains(buf.String(), "creating session") { + t.Errorf("logger output = %q, want it to contain %q", buf.String(), "creating session") + } +} + +func TestWithLogger_nilIgnored(t *testing.T) { + t.Parallel() + st := newTestStore(t, WithLogger(nil)) + if st.logger == nil { + t.Error("logger = nil, want the default of slog.Default()") + } +} + +func TestTimestamp_roundTrip(t *testing.T) { + t.Parallel() + + in := time.Date(2026, 3, 4, 5, 6, 7, 123_000_000, time.FixedZone("EST", -5*60*60)) + text := formatTimestamp(in) + + got, err := parseTimestamp(text) + if err != nil { + t.Fatalf("parseTimestamp(%q): %v", text, err) + } + if !got.Equal(in) { + t.Errorf("round trip = %v, want %v", got, in) + } + if got.Location() != time.UTC { + t.Errorf("parsed location = %v, want UTC", got.Location()) + } +} diff --git a/internal/telemetry/span.go b/internal/telemetry/span.go index e8cfd5f..ac62c6c 100644 --- a/internal/telemetry/span.go +++ b/internal/telemetry/span.go @@ -26,6 +26,21 @@ const ( spanNameLockFileLoad = "registry.lockfile.load" spanNameChecksumVerify = "registry.checksum.verify" spanNamePluginLaunch = "plugin.launch" + + spanNameStateBackendSessionCreate = "statebackend.session.create" + spanNameStateBackendSessionOpen = "statebackend.session.open" + spanNameStateBackendEventAppend = "statebackend.event.append" + spanNameStateBackendMessageAppend = "statebackend.message.append" + spanNameStateBackendPlanAppend = "statebackend.plan.append" + spanNameStateBackendStatusSet = "statebackend.session.set_status" + spanNameStateBackendSessionClose = "statebackend.session.close" + spanNameStateBackendIntegrityCheck = "statebackend.session.integrity_check" + spanNameStateBackendMetaQuery = "statebackend.query.meta" + spanNameStateBackendEventsQuery = "statebackend.query.events" + spanNameStateBackendProducersQuery = "statebackend.query.producers" + spanNameStateBackendCostQuery = "statebackend.query.total_cost" + spanNameStateBackendCostLedgerQuery = "statebackend.query.cost_ledger" + spanNameStateBackendPlanItemsQuery = "statebackend.query.plan_items" ) // SessionSpan describes the session a StartSession call is opening @@ -151,6 +166,115 @@ func (p *Provider) StartPluginLaunch(ctx context.Context, category, name, versio )) } +// StartStateBackendSessionCreate opens the span covering one session file's +// creation — file create, schema apply, PRAGMA user_version stamp, and the +// initial session_meta insert (docs/specifications/state-backend.md#file-layout, +// docs/specifications/state-backend.md#schema-migration) — for use by +// statebackend.Store.Create. +func (p *Provider) StartStateBackendSessionCreate(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendSessionCreate, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendSessionOpen opens the span covering one session file's +// open — the PRAGMA user_version check before any other operation touches +// the file, plus any migration it triggers +// (docs/specifications/state-backend.md#schema-migration) — for use by +// statebackend.Store.Open. +func (p *Provider) StartStateBackendSessionOpen(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendSessionOpen, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendEventAppend opens the span covering one events-table +// append plus its same-transaction producers upsert +// (docs/specifications/state-backend.md#events, +// docs/specifications/state-backend.md#producers), for use by +// statebackend.Session.AppendEvent. +func (p *Provider) StartStateBackendEventAppend(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendEventAppend, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendMessageAppend opens the span covering one events-table +// append plus its same-transaction cost_ledger insert +// (docs/specifications/state-backend.md#cost_ledger), for use by +// statebackend.Session.AppendMessage. +func (p *Provider) StartStateBackendMessageAppend(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendMessageAppend, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendPlanAppend opens the span covering one events-table +// append plus its same-transaction plan_items inserts +// (docs/specifications/state-backend.md#plan_items), for use by +// statebackend.Session.AppendPlan. +func (p *Provider) StartStateBackendPlanAppend(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendPlanAppend, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendStatusSet opens the span covering one in-place +// session_meta update (docs/specifications/state-backend.md#session_meta — +// the one mutable table), for use by statebackend.Session.SetStatus. +func (p *Provider) StartStateBackendStatusSet(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendStatusSet, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendSessionClose opens the span covering one session file's +// close: PRAGMA wal_checkpoint(TRUNCATE) followed by closing the +// underlying *sql.DB, for use by statebackend.Session.Close. +func (p *Provider) StartStateBackendSessionClose(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendSessionClose, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendIntegrityCheck opens the span covering one session +// file's PRAGMA integrity_check on open, plus any salvage recovery it +// triggers (docs/specifications/state-backend.md#corruption-recovery), for +// use by statebackend.Store's Open-time integrity check. +func (p *Provider) StartStateBackendIntegrityCheck(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendIntegrityCheck, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendMetaQuery opens the span covering one session_meta read +// (docs/specifications/state-backend.md#session_meta), for use by +// statebackend.Session.Meta. +func (p *Provider) StartStateBackendMetaQuery(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendMetaQuery, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendEventsQuery opens the span covering one full, +// sequence-ordered events replay read +// (docs/specifications/state-backend.md#events), for use by +// statebackend.Session.Events. The span stays open for the whole +// iter.Seq2 iteration, not just the initial query. +func (p *Provider) StartStateBackendEventsQuery(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendEventsQuery, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendProducersQuery opens the span covering one distinct +// producers read (docs/specifications/state-backend.md#producers), for use +// by statebackend.Session.Producers. +func (p *Provider) StartStateBackendProducersQuery(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendProducersQuery, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendCostQuery opens the span covering one cost_ledger +// SUM(cost_usd) rollup read (docs/specifications/state-backend.md#cost_ledger), +// for use by statebackend.Session.TotalCostUSD. +func (p *Provider) StartStateBackendCostQuery(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendCostQuery, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendCostLedgerQuery opens the span covering one full +// cost_ledger read (docs/specifications/state-backend.md#cost_ledger), for +// use by statebackend.Session.CostLedger. +func (p *Provider) StartStateBackendCostLedgerQuery(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendCostLedgerQuery, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + +// StartStateBackendPlanItemsQuery opens the span covering one full +// plan_items read (docs/specifications/state-backend.md#plan_items), for +// use by statebackend.Session.PlanItems. +func (p *Provider) StartStateBackendPlanItemsQuery(ctx context.Context, sessionID string) (context.Context, trace.Span) { + return p.tracer.Start(ctx, spanNameStateBackendPlanItemsQuery, trace.WithAttributes(SessionIDKey.String(sessionID))) +} + // EndSpan ends span, recording err onto it first if non-nil (RecordError // plus a codes.Error status) so a failed hook/tool/model call is visibly // distinguishable from a successful one in any trace viewer. Every Start*