From b5b488a0d46606ab791d7a6bfecb58c99c1bbe3f Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 11:14:03 -0400 Subject: [PATCH 1/7] Scaffold internal/statebackend package Add package doc, sentinel errors, and ULID session-id helpers with unit and fuzz coverage, per docs/specifications/state-backend.md. Adds github.com/oklog/ulid/v2. --- go.mod | 3 +- go.sum | 9 +- internal/statebackend/doc.go | 3 + internal/statebackend/errors.go | 21 +++ internal/statebackend/sessionid.go | 43 ++++++ internal/statebackend/sessionid_fuzz_test.go | 39 +++++ internal/statebackend/sessionid_test.go | 141 +++++++++++++++++++ 7 files changed, 255 insertions(+), 4 deletions(-) create mode 100644 internal/statebackend/doc.go create mode 100644 internal/statebackend/errors.go create mode 100644 internal/statebackend/sessionid.go create mode 100644 internal/statebackend/sessionid_fuzz_test.go create mode 100644 internal/statebackend/sessionid_test.go diff --git a/go.mod b/go.mod index 9a99009..1e06735 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 @@ -43,7 +44,7 @@ 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/oklog/run v1.1.0 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect diff --git a/go.sum b/go.sum index dc1f869..7e73006 100644 --- a/go.sum +++ b/go.sum @@ -45,12 +45,15 @@ 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/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/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= @@ -120,7 +123,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= diff --git a/internal/statebackend/doc.go b/internal/statebackend/doc.go new file mode 100644 index 0000000..29859e9 --- /dev/null +++ b/internal/statebackend/doc.go @@ -0,0 +1,3 @@ +// Package statebackend implements the kernel state backend specified in docs/specifications/state-backend.md. +// It provides sqlite-per-session persistence: an append-only event log, session metadata, cost ledger, and plan audit trail. +package statebackend diff --git a/internal/statebackend/errors.go b/internal/statebackend/errors.go new file mode 100644 index 0000000..5f2f379 --- /dev/null +++ b/internal/statebackend/errors.go @@ -0,0 +1,21 @@ +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") + +// 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/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..80a6b40 --- /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 := 0; i < goroutines; i++ { + 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 i := 0; i < 10; i++ { + 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()) + } + } +} From 603918ae49feddc6a48bdcff6ca94492e2008e2a Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 11:29:24 -0400 Subject: [PATCH 2/7] Add statebackend store core with schema and migrations Implement the five-table DDL verbatim from state-backend.md, PRAGMA user_version handling with per-step transactional migrations, and the Store lifecycle: Create, Open (version check before any table access, ErrSchemaTooNew on newer files), List and Children via metadata-only directory scans that never migrate as a side effect. Adds modernc.org/sqlite as a direct dependency. --- go.mod | 7 + go.sum | 38 ++ internal/statebackend/schema.go | 188 ++++++ internal/statebackend/schema_test.go | 352 +++++++++++ internal/statebackend/statebackend.go | 481 +++++++++++++++ internal/statebackend/statebackend_test.go | 663 +++++++++++++++++++++ 6 files changed, 1729 insertions(+) create mode 100644 internal/statebackend/schema.go create mode 100644 internal/statebackend/schema_test.go create mode 100644 internal/statebackend/statebackend.go create mode 100644 internal/statebackend/statebackend_test.go diff --git a/go.mod b/go.mod index 1e06735..4262821 100644 --- a/go.mod +++ b/go.mod @@ -28,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 ( @@ -36,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 @@ -46,7 +48,9 @@ require ( github.com/mattn/go-colorable v0.1.12 // 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 @@ -59,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 7e73006..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= @@ -49,6 +55,8 @@ github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWE 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= @@ -56,6 +64,8 @@ github.com/oklog/ulid/v2 v2.1.2/go.mod h1:rcEKHmBBKfef9DhnvX7y1HZBYxjXb0cP5ExxNs 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= @@ -145,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/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..4e55ffe --- /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, ctx context.Context, db *sql.DB, name string) bool { + t.Helper() + var got string + err := db.QueryRowContext(ctx, "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, ctx context.Context, db *sql.DB, name string) bool { + t.Helper() + var got string + err := db.QueryRowContext(ctx, "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, ctx, db, table) { + t.Errorf("table %q missing after initSchema", table) + } + } + for _, idx := range []string{"idx_events_kind", "idx_events_producer"} { + if !indexExists(t, ctx, 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, ctx, 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, ctx, 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, ctx, 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, ctx, 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/statebackend.go b/internal/statebackend/statebackend.go new file mode 100644 index 0000000..9638271 --- /dev/null +++ b/internal/statebackend/statebackend.go @@ -0,0 +1,481 @@ +package statebackend + +import ( + "context" + "database/sql" + "errors" + "fmt" + "io/fs" + "log/slog" + "os" + "path/filepath" + "sort" + "strings" + "time" + + _ "modernc.org/sqlite" // registers the "sqlite" database/sql driver + + 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 plus its identifying ids. Append and query methods land in +// Stages 2-3 (event/cost/plan-item writes, replay reads) — Stage 1 only +// needs enough surface for Create and Open to return a working handle. +type Session struct { + id string + db *sql.DB + path string +} + +// ID returns the session's ULID. +func (s *Session) ID() string { + return s.id +} + +// Close closes the session's underlying *sql.DB. +func (s *Session) Close() error { + if err := s.db.Close(); err != nil { + return fmt.Errorf("statebackend: close %s: %w", s.id, err) + } + return nil +} + +// 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 +} + +// 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 + } + } +} + +// 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) + } + 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. +func openDB(path string) (*sql.DB, error) { + db, err := sql.Open("sqlite", path) + if err != nil { + return nil, fmt.Errorf("statebackend: open %s: %w", path, err) + } + db.SetMaxOpenConns(1) + + if _, err := db.Exec("PRAGMA journal_mode = WAL"); err != nil { + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: set journal_mode: %w", path, err) + } + if _, err := db.Exec("PRAGMA foreign_keys = ON"); err != nil { + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: set foreign_keys: %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, error) { + if err := ValidateSessionID(meta.SessionID); err != nil { + return nil, fmt.Errorf("statebackend: create: %w", 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) + + f, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0o600) + if err != nil { + if errors.Is(err, fs.ErrExist) { + return nil, fmt.Errorf("statebackend: create %s: session file already exists", meta.SessionID) + } + return nil, fmt.Errorf("statebackend: create %s: %w", meta.SessionID, err) + } + if err := f.Close(); err != nil { + _ = os.Remove(path) + return nil, fmt.Errorf("statebackend: create %s: %w", meta.SessionID, err) + } + + sess, err := st.populateCreatedFile(ctx, path, meta) + if err != nil { + _ = os.Remove(path) + return nil, fmt.Errorf("statebackend: create %s: %w", meta.SessionID, 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(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}, nil +} + +// checkIntegrity is a seam for Stage 3's +// docs/specifications/state-backend.md#corruption-recovery flow +// (PRAGMA integrity_check plus salvage-and-rename on failure). It is a +// no-op in Stage 1/2 so Open's structure does not need to change when +// recovery is added. +func (st *Store) checkIntegrity(_ context.Context, _ string, _ *sql.DB) error { + return nil +} + +// Open opens an existing session file for sessionID. It validates +// sessionID, runs the (currently no-op) corruption-recovery seam, 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, error) { + if err := ValidateSessionID(sessionID); err != nil { + return nil, fmt.Errorf("statebackend: open: %w", err) + } + st.logger.DebugContext(ctx, "statebackend: opening session", "session_id", sessionID) + + path := st.sessionPath(sessionID) + if _, err := os.Stat(path); err != nil { + if errors.Is(err, fs.ErrNotExist) { + return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, ErrNotFound) + } + return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + } + + db, err := openDB(path) + if err != nil { + return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + } + + if err := st.checkIntegrity(ctx, path, db); err != nil { + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + } + + version, err := readUserVersion(ctx, db) + if err != nil { + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + } + switch { + case version > currentSchemaVersion: + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: schema version %d: %w", sessionID, version, ErrSchemaTooNew) + case version < currentSchemaVersion: + if err := applyMigrations(ctx, db, migrations, version, currentSchemaVersion); err != nil { + _ = db.Close() + return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + } + } + + return &Session{id: sessionID, db: db, path: path}, 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(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..a503de7 --- /dev/null +++ b/internal/statebackend/statebackend_test.go @@ -0,0 +1,663 @@ +package statebackend + +import ( + "bytes" + "context" + "database/sql" + "errors" + "io/fs" + "log/slog" + "os" + "path/filepath" + "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") + } + 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) + } + 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.Exec("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() + 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()} + if _, err := st.Create(context.Background(), meta); err == nil { + t.Fatal("Create in a read-only directory = nil error, want error") + } +} + +func TestOpenDB_directoryPathFails(t *testing.T) { + t.Parallel() + if _, err := openDB(t.TempDir()); err == nil { + t.Fatal("openDB(directory) = nil error, want error") + } +} + +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(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()) + } +} From f353aad42221cacbdf4678adc166307fefbfb016 Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 11:48:57 -0400 Subject: [PATCH 3/7] Add statebackend append paths and telemetry spans Implement sole-writer append transactions: AppendEvent with same-tx producer upsert, AppendMessage with cost_ledger row, AppendPlan with plan_items rows, SetStatus, and idempotent Close with WAL checkpoint. Event kinds, session statuses, plan decisions, and producer categories encode to the spec's lowercase text vocabulary; unspecified enum values are rejected. Add StartStateBackend* span helpers to internal/telemetry and instrument Create/Open. Race tests assert dense 1..N sequences under concurrent appenders. --- internal/statebackend/errors.go | 3 + internal/statebackend/event.go | 214 +++++++++ internal/statebackend/event_test.go | 117 +++++ internal/statebackend/schema_test.go | 20 +- internal/statebackend/session.go | 245 ++++++++++ internal/statebackend/session_test.go | 508 +++++++++++++++++++++ internal/statebackend/sessionid_test.go | 4 +- internal/statebackend/statebackend.go | 170 ++++--- internal/statebackend/statebackend_test.go | 6 +- internal/telemetry/span.go | 65 +++ 10 files changed, 1282 insertions(+), 70 deletions(-) create mode 100644 internal/statebackend/event.go create mode 100644 internal/statebackend/event_test.go create mode 100644 internal/statebackend/session.go create mode 100644 internal/statebackend/session_test.go diff --git a/internal/statebackend/errors.go b/internal/statebackend/errors.go index 5f2f379..1367ea0 100644 --- a/internal/statebackend/errors.go +++ b/internal/statebackend/errors.go @@ -14,6 +14,9 @@ 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") 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_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/schema_test.go b/internal/statebackend/schema_test.go index 4e55ffe..0c109a7 100644 --- a/internal/statebackend/schema_test.go +++ b/internal/statebackend/schema_test.go @@ -27,10 +27,10 @@ func openTestDB(t *testing.T) *sql.DB { } // tableExists reports whether name is a table in db's sqlite_master. -func tableExists(t *testing.T, ctx context.Context, db *sql.DB, name string) bool { +func tableExists(t *testing.T, db *sql.DB, name string) bool { t.Helper() var got string - err := db.QueryRowContext(ctx, "SELECT name FROM sqlite_master WHERE type = 'table' AND name = ?", name).Scan(&got) + err := db.QueryRowContext(context.Background(), "SELECT name FROM sqlite_master WHERE type = 'table' AND name = ?", name).Scan(&got) switch { case err == nil: return true @@ -43,10 +43,10 @@ func tableExists(t *testing.T, ctx context.Context, db *sql.DB, name string) boo } // indexExists reports whether name is an index in db's sqlite_master. -func indexExists(t *testing.T, ctx context.Context, db *sql.DB, name string) bool { +func indexExists(t *testing.T, db *sql.DB, name string) bool { t.Helper() var got string - err := db.QueryRowContext(ctx, "SELECT name FROM sqlite_master WHERE type = 'index' AND name = ?", name).Scan(&got) + err := db.QueryRowContext(context.Background(), "SELECT name FROM sqlite_master WHERE type = 'index' AND name = ?", name).Scan(&got) switch { case err == nil: return true @@ -68,12 +68,12 @@ func TestInitSchema(t *testing.T) { } for _, table := range []string{"events", "session_meta", "cost_ledger", "plan_items", "producers"} { - if !tableExists(t, ctx, db, table) { + 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, ctx, db, idx) { + if !indexExists(t, db, idx) { t.Errorf("index %q missing after initSchema", idx) } } @@ -99,7 +99,7 @@ func TestInitSchema_eventsAutoincrement(t *testing.T) { // 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, ctx, db, "sqlite_sequence") { + if !tableExists(t, db, "sqlite_sequence") { t.Error("sqlite_sequence table missing — AUTOINCREMENT was not applied") } } @@ -216,7 +216,7 @@ func TestApplyMigrations(t *testing.T) { } for _, table := range []string{"migration_marker_1", "migration_marker_2"} { - if !tableExists(t, ctx, db, table) { + if !tableExists(t, db, table) { t.Errorf("table %q missing after migration", table) } } @@ -313,7 +313,7 @@ func TestApplyMigrations_failingStepLeavesVersionUnchanged(t *testing.T) { t.Errorf("user_version = %d, want 0 (unchanged after failed migration)", version) } - if tableExists(t, ctx, db, "should_not_persist") { + if tableExists(t, db, "should_not_persist") { t.Error("should_not_persist table exists after a failed, rolled-back migration step") } } @@ -346,7 +346,7 @@ func TestApplyMigrations_multiStepPartialFailureStopsAtLastGood(t *testing.T) { if version != 1 { t.Errorf("user_version = %d, want 1 (step 1 committed, step 2 rolled back)", version) } - if !tableExists(t, ctx, db, "step_one") { + 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_test.go b/internal/statebackend/sessionid_test.go index 80a6b40..4e7b3ab 100644 --- a/internal/statebackend/sessionid_test.go +++ b/internal/statebackend/sessionid_test.go @@ -71,7 +71,7 @@ func TestNewSessionIDConcurrentUniqueness(t *testing.T) { now := time.Now() - for i := 0; i < goroutines; i++ { + for i := range goroutines { go func(idx int) { defer wg.Done() ids[idx] = NewSessionID(now) @@ -123,7 +123,7 @@ func TestValidateSessionIDRoundTrip(t *testing.T) { t.Parallel() // Any generated ID must pass validation and round-trip. - for i := 0; i < 10; i++ { + for range 10 { id := NewSessionID(time.Now()) if err := ValidateSessionID(id); err != nil { t.Errorf("Generated ID failed validation: %q, %v", id, err) diff --git a/internal/statebackend/statebackend.go b/internal/statebackend/statebackend.go index 9638271..f0a50c2 100644 --- a/internal/statebackend/statebackend.go +++ b/internal/statebackend/statebackend.go @@ -11,10 +11,13 @@ import ( "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" ) @@ -110,13 +113,21 @@ type SessionMeta struct { } // Session is an open handle to one session's sqlite file: the *sql.DB -// backing it plus its identifying ids. Append and query methods land in -// Stages 2-3 (event/cost/plan-item writes, replay reads) — Stage 1 only -// needs enough surface for Create and Open to return a working handle. +// 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 + 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. @@ -124,22 +135,15 @@ func (s *Session) ID() string { return s.id } -// Close closes the session's underlying *sql.DB. -func (s *Session) Close() error { - if err := s.db.Close(); err != nil { - return fmt.Errorf("statebackend: close %s: %w", s.id, err) - } - return nil -} - // 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 + dir string + clock func() time.Time + logger *slog.Logger + telemetry *telemetry.Provider } // Option configures a Store constructed by NewStore. @@ -167,6 +171,35 @@ func WithLogger(logger *slog.Logger) Option { } } +// 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 @@ -185,6 +218,13 @@ func NewStore(dir string, opts ...Option) (*Store, error) { 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 } @@ -199,18 +239,18 @@ func (st *Store) sessionPath(sessionID string) string { // 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. -func openDB(path string) (*sql.DB, error) { +func openDB(ctx context.Context, path string) (*sql.DB, error) { db, err := sql.Open("sqlite", path) if err != nil { return nil, fmt.Errorf("statebackend: open %s: %w", path, err) } db.SetMaxOpenConns(1) - if _, err := db.Exec("PRAGMA journal_mode = WAL"); err != nil { + 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) } - if _, err := db.Exec("PRAGMA foreign_keys = ON"); err != nil { + if _, err := db.ExecContext(ctx, "PRAGMA foreign_keys = ON"); err != nil { _ = db.Close() return nil, fmt.Errorf("statebackend: open %s: set foreign_keys: %w", path, err) } @@ -226,9 +266,13 @@ func openDB(path string) (*sql.DB, error) { // 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, error) { - if err := ValidateSessionID(meta.SessionID); err != nil { - return nil, fmt.Errorf("statebackend: create: %w", err) +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() @@ -237,22 +281,27 @@ func (st *Store) Create(ctx context.Context, meta SessionMeta) (*Session, error) path := st.sessionPath(meta.SessionID) st.logger.DebugContext(ctx, "statebackend: creating session", "session_id", meta.SessionID) - f, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0o600) - if err != nil { - if errors.Is(err, fs.ErrExist) { - return nil, fmt.Errorf("statebackend: create %s: session file already exists", 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 } - return nil, fmt.Errorf("statebackend: create %s: %w", meta.SessionID, err) + err = fmt.Errorf("statebackend: create %s: %w", meta.SessionID, openErr) + return nil, err } - if err := f.Close(); err != nil { + if closeErr := f.Close(); closeErr != nil { _ = os.Remove(path) - return nil, fmt.Errorf("statebackend: create %s: %w", meta.SessionID, err) + err = fmt.Errorf("statebackend: create %s: %w", meta.SessionID, closeErr) + return nil, err } - sess, err := st.populateCreatedFile(ctx, path, meta) - if err != nil { + sess, popErr := st.populateCreatedFile(ctx, path, meta) + if popErr != nil { _ = os.Remove(path) - return nil, fmt.Errorf("statebackend: create %s: %w", meta.SessionID, err) + err = fmt.Errorf("statebackend: create %s: %w", meta.SessionID, popErr) + return nil, err } return sess, nil } @@ -261,7 +310,7 @@ func (st *Store) Create(ctx context.Context, meta SessionMeta) (*Session, error) // 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(path) + db, err := openDB(ctx, path) if err != nil { return nil, err } @@ -273,7 +322,7 @@ func (st *Store) populateCreatedFile(ctx context.Context, path string, meta Sess _ = db.Close() return nil, err } - return &Session{id: meta.SessionID, db: db, path: path}, nil + return &Session{id: meta.SessionID, db: db, path: path, logger: st.logger, telemetry: st.telemetry}, nil } // checkIntegrity is a seam for Stage 3's @@ -291,47 +340,58 @@ func (st *Store) checkIntegrity(_ context.Context, _ string, _ *sql.DB) error { // (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, error) { - if err := ValidateSessionID(sessionID); err != nil { - return nil, fmt.Errorf("statebackend: open: %w", err) +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 _, err := os.Stat(path); err != nil { - if errors.Is(err, fs.ErrNotExist) { - return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, ErrNotFound) + 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 } - return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + err = fmt.Errorf("statebackend: open %s: %w", sessionID, statErr) + return nil, err } - db, err := openDB(path) - if err != nil { - return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + db, openErr := openDB(ctx, path) + if openErr != nil { + err = fmt.Errorf("statebackend: open %s: %w", sessionID, openErr) + return nil, err } - if err := st.checkIntegrity(ctx, path, db); err != nil { + if icErr := st.checkIntegrity(ctx, path, db); icErr != nil { _ = db.Close() - return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + err = fmt.Errorf("statebackend: open %s: %w", sessionID, icErr) + return nil, err } - version, err := readUserVersion(ctx, db) - if err != nil { + version, verErr := readUserVersion(ctx, db) + if verErr != nil { _ = db.Close() - return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + err = fmt.Errorf("statebackend: open %s: %w", sessionID, verErr) + return nil, err } switch { case version > currentSchemaVersion: _ = db.Close() - return nil, fmt.Errorf("statebackend: open %s: schema version %d: %w", sessionID, version, ErrSchemaTooNew) + err = fmt.Errorf("statebackend: open %s: schema version %d: %w", sessionID, version, ErrSchemaTooNew) + return nil, err case version < currentSchemaVersion: - if err := applyMigrations(ctx, db, migrations, version, currentSchemaVersion); err != nil { + if migErr := applyMigrations(ctx, db, migrations, version, currentSchemaVersion); migErr != nil { _ = db.Close() - return nil, fmt.Errorf("statebackend: open %s: %w", sessionID, err) + err = fmt.Errorf("statebackend: open %s: %w", sessionID, migErr) + return nil, err } } - return &Session{id: sessionID, db: db, path: path}, nil + 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 @@ -395,7 +455,7 @@ func (st *Store) scan(ctx context.Context, keep func(SessionMeta) bool) ([]Sessi // 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(st.sessionPath(sessionID)) + db, err := openDB(ctx, st.sessionPath(sessionID)) if err != nil { return SessionMeta{}, err } diff --git a/internal/statebackend/statebackend_test.go b/internal/statebackend/statebackend_test.go index a503de7..68cbc3d 100644 --- a/internal/statebackend/statebackend_test.go +++ b/internal/statebackend/statebackend_test.go @@ -266,7 +266,7 @@ func TestStore_Open_schemaTooNew(t *testing.T) { if err != nil { t.Fatalf("sql.Open: %v", err) } - if _, err := db.Exec("PRAGMA user_version = 999"); err != nil { + 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 { @@ -472,7 +472,7 @@ func TestStore_Create_openFileError(t *testing.T) { func TestOpenDB_directoryPathFails(t *testing.T) { t.Parallel() - if _, err := openDB(t.TempDir()); err == nil { + if _, err := openDB(context.Background(), t.TempDir()); err == nil { t.Fatal("openDB(directory) = nil error, want error") } } @@ -496,7 +496,7 @@ func TestPopulateCreatedFile_initSchemaFailure(t *testing.T) { st := newTestStore(t) path := filepath.Join(st.dir, "existing.sqlite") - db, err := openDB(path) + db, err := openDB(context.Background(), path) if err != nil { t.Fatalf("openDB: %v", err) } diff --git a/internal/telemetry/span.go b/internal/telemetry/span.go index e8cfd5f..ac75d9c 100644 --- a/internal/telemetry/span.go +++ b/internal/telemetry/span.go @@ -26,6 +26,14 @@ 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" ) // SessionSpan describes the session a StartSession call is opening @@ -151,6 +159,63 @@ 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))) +} + // 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* From 31c9206e027341e2630fc6844afd51f834175798 Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 12:12:33 -0400 Subject: [PATCH 4/7] Add statebackend queries and corruption recovery Implement the read-side API: sequence-ordered Events replay iterator, deduplicated Producers set, TotalCostUSD, CostLedger, PlanItems, and Meta. Wire Open through a real integrity check: corrupt or unopenable files are renamed to .corrupt (never deleted), salvaged table-by-table into a fresh schema-correct file preserving sequence values, with slog WARN reporting; unsalvageable files return ErrUnrecoverable. Add an event round-trip fuzz target and query/recovery span helpers. --- internal/statebackend/event_fuzz_test.go | 109 +++++ internal/statebackend/integrity.go | 281 +++++++++++++ internal/statebackend/integrity_test.go | 303 ++++++++++++++ internal/statebackend/query.go | 259 ++++++++++++ internal/statebackend/query_test.go | 501 +++++++++++++++++++++++ internal/statebackend/statebackend.go | 25 +- internal/telemetry/span.go | 73 +++- 7 files changed, 1525 insertions(+), 26 deletions(-) create mode 100644 internal/statebackend/event_fuzz_test.go create mode 100644 internal/statebackend/integrity.go create mode 100644 internal/statebackend/integrity_test.go create mode 100644 internal/statebackend/query.go create mode 100644 internal/statebackend/query_test.go 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/integrity.go b/internal/statebackend/integrity.go new file mode 100644 index 0000000..4b1fbf1 --- /dev/null +++ b/internal/statebackend/integrity.go @@ -0,0 +1,281 @@ +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) + err = fmt.Errorf("statebackend: recover %s: %w", sessionID, ErrUnrecoverable) + 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 + + recDB, err := openDB(ctx, recoveryPath) + if err != nil { + 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..4043f88 --- /dev/null +++ b/internal/statebackend/integrity_test.go @@ -0,0 +1,303 @@ +package statebackend + +import ( + "bytes" + "context" + "errors" + "fmt" + "io/fs" + "log/slog" + "os" + "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) + } + + 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/statebackend.go b/internal/statebackend/statebackend.go index f0a50c2..ab97aa8 100644 --- a/internal/statebackend/statebackend.go +++ b/internal/statebackend/statebackend.go @@ -325,18 +325,11 @@ func (st *Store) populateCreatedFile(ctx context.Context, path string, meta Sess return &Session{id: meta.SessionID, db: db, path: path, logger: st.logger, telemetry: st.telemetry}, nil } -// checkIntegrity is a seam for Stage 3's -// docs/specifications/state-backend.md#corruption-recovery flow -// (PRAGMA integrity_check plus salvage-and-rename on failure). It is a -// no-op in Stage 1/2 so Open's structure does not need to change when -// recovery is added. -func (st *Store) checkIntegrity(_ context.Context, _ string, _ *sql.DB) error { - return nil -} - // Open opens an existing session file for sessionID. It validates -// sessionID, runs the (currently no-op) corruption-recovery seam, then -// checks PRAGMA user_version before any other operation touches the file +// 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. @@ -360,14 +353,8 @@ func (st *Store) Open(ctx context.Context, sessionID string) (_ *Session, err er return nil, err } - db, openErr := openDB(ctx, path) - if openErr != nil { - err = fmt.Errorf("statebackend: open %s: %w", sessionID, openErr) - return nil, err - } - - if icErr := st.checkIntegrity(ctx, path, db); icErr != nil { - _ = db.Close() + db, icErr := st.checkIntegrity(ctx, path, sessionID) + if icErr != nil { err = fmt.Errorf("statebackend: open %s: %w", sessionID, icErr) return nil, err } diff --git a/internal/telemetry/span.go b/internal/telemetry/span.go index ac75d9c..ac62c6c 100644 --- a/internal/telemetry/span.go +++ b/internal/telemetry/span.go @@ -27,13 +27,20 @@ const ( 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" + 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 @@ -216,6 +223,58 @@ func (p *Provider) StartStateBackendSessionClose(ctx context.Context, sessionID 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* From 8787d247b1210ca77a9e6d4bb851d93358b9c3f7 Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 12:20:08 -0400 Subject: [PATCH 5/7] Document statebackend and align session-tree spec wording Expand the package godoc with the design decisions the spec left open, add README and package CLAUDE.md, enumerate producer_category's text vocabulary in state-backend.md, and fix forward the stale parent-to- children index claims in subagents.md, conformance.md, and turn-algorithm.md: parent linkage lives in session_meta, child lookup is a session_meta scan, and cost-rollup, depth threading, and cancellation are live in-memory mechanisms, not state-backend reads. --- docs/specifications/agent-loop/conformance.md | 2 +- docs/specifications/agent-loop/subagents.md | 6 +- .../agent-loop/turn-algorithm.md | 2 +- docs/specifications/state-backend.md | 2 +- internal/statebackend/CLAUDE.md | 8 +++ internal/statebackend/README.md | 43 ++++++++++++ internal/statebackend/doc.go | 65 ++++++++++++++++++- 7 files changed, 120 insertions(+), 8 deletions(-) create mode 100644 internal/statebackend/CLAUDE.md create mode 100644 internal/statebackend/README.md 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/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 index 29859e9..d712324 100644 --- a/internal/statebackend/doc.go +++ b/internal/statebackend/doc.go @@ -1,3 +1,64 @@ -// Package statebackend implements the kernel state backend specified in docs/specifications/state-backend.md. -// It provides sqlite-per-session persistence: an append-only event log, session metadata, cost ledger, and plan audit trail. +// 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. +// - 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 From 0538a93e3ea10b58b1225b3e401f0350141c95ba Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 12:22:50 -0400 Subject: [PATCH 6/7] Harden statebackend recovery and connection pragmas Pre-create the recovery scratch file at 0600 so recovered session files keep the documented mode instead of SQLite's 0644 default, move foreign_keys enforcement into the DSN so every pooled connection gets it rather than only the first, and chain the salvage failure cause through ErrUnrecoverable so callers keep both errors.Is matching and the underlying reason. --- internal/statebackend/integrity.go | 23 ++++++++++++++++++++- internal/statebackend/integrity_test.go | 13 ++++++++++++ internal/statebackend/statebackend.go | 18 +++++++++++----- internal/statebackend/statebackend_test.go | 24 ++++++++++++++++++++++ 4 files changed, 72 insertions(+), 6 deletions(-) diff --git a/internal/statebackend/integrity.go b/internal/statebackend/integrity.go index 4b1fbf1..687982d 100644 --- a/internal/statebackend/integrity.go +++ b/internal/statebackend/integrity.go @@ -102,7 +102,10 @@ func (st *Store) checkIntegrity(ctx context.Context, path, sessionID string) (_ 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) - err = fmt.Errorf("statebackend: recover %s: %w", sessionID, ErrUnrecoverable) + // 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 } @@ -161,8 +164,26 @@ func (st *Store) recoverSession(ctx context.Context, srcPath, dstPath, sessionID 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 diff --git a/internal/statebackend/integrity_test.go b/internal/statebackend/integrity_test.go index 4043f88..6ed64f9 100644 --- a/internal/statebackend/integrity_test.go +++ b/internal/statebackend/integrity_test.go @@ -146,6 +146,19 @@ func TestStore_Open_corruptionTriggersRecovery(t *testing.T) { 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). + 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) diff --git a/internal/statebackend/statebackend.go b/internal/statebackend/statebackend.go index ab97aa8..7770d9f 100644 --- a/internal/statebackend/statebackend.go +++ b/internal/statebackend/statebackend.go @@ -239,8 +239,20 @@ func (st *Store) sessionPath(sessionID string) string { // 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) + db, err := sql.Open("sqlite", path+"?_pragma=foreign_keys(1)") if err != nil { return nil, fmt.Errorf("statebackend: open %s: %w", path, err) } @@ -250,10 +262,6 @@ func openDB(ctx context.Context, path string) (*sql.DB, error) { _ = db.Close() return nil, fmt.Errorf("statebackend: open %s: set journal_mode: %w", path, err) } - if _, err := db.ExecContext(ctx, "PRAGMA foreign_keys = ON"); err != nil { - _ = db.Close() - return nil, fmt.Errorf("statebackend: open %s: set foreign_keys: %w", path, err) - } return db, nil } diff --git a/internal/statebackend/statebackend_test.go b/internal/statebackend/statebackend_test.go index 68cbc3d..bfae7f4 100644 --- a/internal/statebackend/statebackend_test.go +++ b/internal/statebackend/statebackend_test.go @@ -477,6 +477,30 @@ func TestOpenDB_directoryPathFails(t *testing.T) { } } +// 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) From 50ccaf903eeea7e2a94b830c6d9b9e5aa8ea93fe Mon Sep 17 00:00:00 2001 From: Steven Crothers Date: Fri, 24 Jul 2026 12:32:39 -0400 Subject: [PATCH 7/7] Gate Unix permission assertions off on Windows Unix mode bits aren't meaningful on Windows: created files report 666, directories 777, and read-only directory modes don't block file creation. Guard the permission assertions per-OS, skip the read-only directory test on Windows, close the session on an unexpected Create success so TempDir cleanup can't wedge on an open handle, and note in the package doc that the 0700/0600 guarantee is Unix-only. --- internal/statebackend/doc.go | 3 +- internal/statebackend/integrity_test.go | 18 ++++++----- internal/statebackend/statebackend_test.go | 35 ++++++++++++++++++---- 3 files changed, 43 insertions(+), 13 deletions(-) diff --git a/internal/statebackend/doc.go b/internal/statebackend/doc.go index d712324..da1f03d 100644 --- a/internal/statebackend/doc.go +++ b/internal/statebackend/doc.go @@ -47,7 +47,8 @@ // 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. +// 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 diff --git a/internal/statebackend/integrity_test.go b/internal/statebackend/integrity_test.go index 6ed64f9..f75ffe4 100644 --- a/internal/statebackend/integrity_test.go +++ b/internal/statebackend/integrity_test.go @@ -8,6 +8,7 @@ import ( "io/fs" "log/slog" "os" + "runtime" "strings" "testing" "time" @@ -150,13 +151,16 @@ func TestStore_Open_corruptionTriggersRecovery(t *testing.T) { // 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). - 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) + // 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()) diff --git a/internal/statebackend/statebackend_test.go b/internal/statebackend/statebackend_test.go index bfae7f4..db9f39d 100644 --- a/internal/statebackend/statebackend_test.go +++ b/internal/statebackend/statebackend_test.go @@ -9,6 +9,7 @@ import ( "log/slog" "os" "path/filepath" + "runtime" "strings" "testing" "time" @@ -59,8 +60,14 @@ func TestNewStore(t *testing.T) { if !info.IsDir() { t.Fatal("dir is not a directory") } - if perm := info.Mode().Perm(); perm != 0o700 { - t.Errorf("dir perm = %o, want %o", perm, 0o700) + // 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) + } } }) @@ -98,8 +105,12 @@ func TestStore_Create(t *testing.T) { if err != nil { t.Fatalf("Stat: %v", err) } - if perm := info.Mode().Perm(); perm != 0o600 { - t.Errorf("file perm = %o, want %o", perm, 0o600) + // 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) @@ -458,6 +469,13 @@ func TestNewStore_mkdirFails(t *testing.T) { 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) @@ -465,7 +483,14 @@ func TestStore_Create_openFileError(t *testing.T) { defer func() { _ = os.Chmod(st.dir, 0o700) }() meta := SessionMeta{SessionID: NewSessionID(time.Now()), Profile: "default", Status: sessionv1.SessionStatus_SESSION_STATUS_RUNNING, StartedAt: time.Now()} - if _, err := st.Create(context.Background(), meta); err == nil { + 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") } }