Files
notarius/internal/framework/chunkplan/store_test.go

509 lines
16 KiB
Go

package chunkplan
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"reflect"
"strings"
"sync"
"testing"
"time"
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
)
const testSourceDigest = "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"
func TestFilesystemStoreRoundTripAndExactPath(t *testing.T) {
root := filepath.Join(t.TempDir(), "plans")
store := newStore(t, root)
record := testRecord(t, 1)
if err := store.Save(record); err != nil {
t.Fatalf("Save() error = %v", err)
}
target := filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:"), "plan.json")
data, err := os.ReadFile(target)
if err != nil {
t.Fatalf("read exact plan path: %v", err)
}
if bytes.Contains(data, []byte("RAW REFERENCE CONTENT")) {
t.Fatal("stored record contains raw reference content")
}
var envelope struct {
Producer struct {
References []map[string]any `json:"references"`
} `json:"producer"`
}
if err := json.Unmarshal(data, &envelope); err != nil {
t.Fatal(err)
}
if len(envelope.Producer.References) != 1 {
t.Fatalf("stored references = %#v", envelope.Producer.References)
}
if _, exists := envelope.Producer.References[0]["content"]; exists {
t.Fatalf("stored reference contains content field: %#v", envelope.Producer.References[0])
}
got, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanHit {
t.Fatalf("Load() decision=%#v error=%v", decision, err)
}
if !reflect.DeepEqual(got, record) {
t.Fatalf("round trip record = %#v, want %#v", got, record)
}
}
func TestFilesystemStorePermissions(t *testing.T) {
root := filepath.Join(t.TempDir(), "plans")
store := newStore(t, root)
if err := store.Save(testRecord(t, 1)); err != nil {
t.Fatal(err)
}
digestDir := filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:"))
for path, want := range map[string]os.FileMode{
root: 0o700,
digestDir: 0o700,
filepath.Join(digestDir, "plan.json"): 0o600,
} {
info, err := os.Stat(path)
if err != nil {
t.Fatal(err)
}
if got := info.Mode().Perm(); got != want {
t.Fatalf("%s mode = %04o, want %04o", path, got, want)
}
}
}
func TestFilesystemStoreMissingAndOperationalErrors(t *testing.T) {
root := filepath.Join(t.TempDir(), "plans")
store := newStore(t, root)
_, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanMissing {
t.Fatalf("missing decision=%#v error=%v", decision, err)
}
if _, err := os.Stat(root); !os.IsNotExist(err) {
t.Fatalf("Load() created missing root: %v", err)
}
target := filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:"), "plan.json")
if err := os.MkdirAll(target, 0o700); err != nil {
t.Fatal(err)
}
_, decision, err = store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanInvalid {
t.Fatalf("Load(plan.json directory) decision=%#v error=%v", decision, err)
}
}
func TestFilesystemStoreRejectsSymlinkedEntries(t *testing.T) {
tests := []struct {
name string
setup func(t *testing.T, root, outside string)
}{
{
name: "digest directory",
setup: func(t *testing.T, root, outside string) {
t.Helper()
if err := os.MkdirAll(outside, 0o700); err != nil {
t.Fatal(err)
}
writeFile(t, filepath.Join(outside, "plan.json"), []byte("outside plan"), 0o640)
createSymlink(t, outside, filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:")))
},
},
{
name: "plan file",
setup: func(t *testing.T, root, outside string) {
t.Helper()
digestDir := filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:"))
if err := os.MkdirAll(digestDir, 0o700); err != nil {
t.Fatal(err)
}
writeFile(t, outside, []byte("outside plan"), 0o640)
createSymlink(t, outside, filepath.Join(digestDir, "plan.json"))
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
root := filepath.Join(t.TempDir(), "plans")
if err := os.MkdirAll(root, 0o700); err != nil {
t.Fatal(err)
}
outside := filepath.Join(t.TempDir(), "outside")
tc.setup(t, root, outside)
before := readFile(t, outsidePlanPath(tc.name, outside))
beforeMode := fileMode(t, outsidePlanPath(tc.name, outside))
store := newStore(t, root)
_, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanInvalid {
t.Fatalf("Load() decision=%#v error=%v", decision, err)
}
if err := store.Save(testRecord(t, 1)); err == nil {
t.Fatal("Save() error = nil")
}
outsidePlan := outsidePlanPath(tc.name, outside)
if got := readFile(t, outsidePlan); !bytes.Equal(got, before) {
t.Fatalf("outside content = %q, want %q", got, before)
}
if got := fileMode(t, outsidePlan); got != beforeMode {
t.Fatalf("outside mode = %04o, want %04o", got, beforeMode)
}
})
}
}
func TestFilesystemStoreRejectsUnexpectedEntryTypes(t *testing.T) {
tests := []struct {
name string
setup func(t *testing.T, root string)
}{
{
name: "digest file",
setup: func(t *testing.T, root string) {
writeFile(t, filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:")), []byte("not a directory"), 0o600)
},
},
{
name: "plan directory",
setup: func(t *testing.T, root string) {
if err := os.MkdirAll(filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:"), "plan.json"), 0o700); err != nil {
t.Fatal(err)
}
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
root := t.TempDir()
tc.setup(t, root)
store := newStore(t, root)
_, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanInvalid {
t.Fatalf("Load() decision=%#v error=%v", decision, err)
}
if err := store.Save(testRecord(t, 1)); err == nil {
t.Fatal("Save() error = nil")
}
})
}
}
func TestFilesystemStoreRejectsMalformedSourceDigests(t *testing.T) {
store := newStore(t, t.TempDir())
for _, digest := range []string{"", "sha1:" + strings.Repeat("a", 64), "sha256:../escape", "sha256:" + strings.Repeat("A", 64), "sha256:" + strings.Repeat("a", 63)} {
t.Run(digest, func(t *testing.T) {
if _, _, err := store.Load(digest); err == nil {
t.Fatal("Load() error = nil")
}
record := testRecord(t, 1)
record.SourceDigest = digest
record.Plan.SourceDigest = digest
record.PlanDigest, _ = source.DigestChunkPlan(record.Plan)
if err := store.Save(record); err == nil {
t.Fatal("Save() error = nil")
}
})
}
}
func TestFilesystemStoreReportsInvalidRecordsAsRecoverable(t *testing.T) {
tests := []struct {
name string
mutate func([]byte) []byte
}{
{name: "unknown field", mutate: func(data []byte) []byte {
return bytes.Replace(data, []byte(`{"schema_version"`), []byte(`{"SENTINEL_UNKNOWN_FIELD":true,"schema_version"`), 1)
}},
{name: "truncated JSON", mutate: func(data []byte) []byte { return data[:len(data)/2] }},
{name: "schema mismatch", mutate: replaceJSON(`notarius.chunk-plan.v1`, `SENTINEL_SCHEMA_VALUE`)},
{name: "source mismatch", mutate: replaceJSON(testSourceDigest, "sha256:"+strings.Repeat("b", 64))},
{name: "plan digest mismatch", mutate: func(data []byte) []byte {
prefix := []byte(`"plan_digest":"sha256:`)
index := bytes.Index(data, prefix)
if index >= 0 {
data[index+len(prefix)] = '0'
}
return data
}},
{name: "noncanonical annotation", mutate: func(data []byte) []byte {
return bytes.Replace(data, []byte(`"test/value"`), []byte(`"SENTINEL_ANNOTATION_NAMESPACE"`), 1)
}},
{name: "bad boundary", mutate: func(data []byte) []byte {
return bytes.Replace(data, []byte(`"start_unit_id":1`), []byte(`"start_unit_id":0`), 1)
}},
{name: "timestamp", mutate: replaceJSON(`2026-07-18T12:00:00Z`, `SENTINEL_TIMESTAMP`)},
{name: "trailing JSON", mutate: func(data []byte) []byte { return append(data, []byte(` {"SENTINEL_TRAILING":true}`)...) }},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
root := t.TempDir()
store := newStore(t, root)
record := testRecord(t, 1)
if err := store.Save(record); err != nil {
t.Fatal(err)
}
path := filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:"), "plan.json")
data, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, tc.mutate(data), 0o600); err != nil {
t.Fatal(err)
}
got, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanInvalid || !reflect.DeepEqual(got, pipeline.ChunkPlanRecord{}) || decision.Reason != "stored chunk plan is invalid" {
t.Fatalf("record=%#v decision=%#v error=%v", got, decision, err)
}
for _, sentinel := range []string{"SENTINEL_UNKNOWN_FIELD", "SENTINEL_SCHEMA_VALUE", "SENTINEL_ANNOTATION_NAMESPACE", "SENTINEL_TIMESTAMP", "SENTINEL_TRAILING"} {
if strings.Contains(decision.Reason, sentinel) {
t.Fatalf("decision leaked %q: %#v", sentinel, decision)
}
}
})
}
}
func TestFilesystemStoreAtomicallyReplacesAndPreservesValidRecordOnFailure(t *testing.T) {
root := t.TempDir()
store := newStore(t, root)
first := testRecord(t, 1)
second := testRecord(t, 2)
if err := store.Save(first); err != nil {
t.Fatal(err)
}
invalid := second
invalid.PlanDigest = "sha256:" + strings.Repeat("0", 64)
if err := store.Save(invalid); err == nil {
t.Fatal("Save(invalid) error = nil")
}
got, _, _ := store.Load(testSourceDigest)
if !reflect.DeepEqual(got, first) {
t.Fatalf("record after failed replacement = %#v", got)
}
if err := store.Save(second); err != nil {
t.Fatal(err)
}
got, _, _ = store.Load(testSourceDigest)
if !reflect.DeepEqual(got, second) {
t.Fatalf("record after replacement = %#v", got)
}
assertNoTemps(t, filepath.Join(root, strings.TrimPrefix(testSourceDigest, "sha256:")))
}
func TestFilesystemStoreConcurrentWritersExposeCompleteRecord(t *testing.T) {
store := newStore(t, t.TempDir())
const writers = 24
records := make([]pipeline.ChunkPlanRecord, writers)
for i := range records {
records[i] = testRecord(t, i+1)
}
var wg sync.WaitGroup
errs := make(chan error, writers)
for i := range records {
wg.Add(1)
go func(record pipeline.ChunkPlanRecord) {
defer wg.Done()
errs <- store.Save(record)
}(records[i])
}
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatalf("concurrent Save() error = %v", err)
}
}
got, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanHit {
t.Fatalf("Load() decision=%#v error=%v", decision, err)
}
var annotation struct {
Value int `json:"value"`
}
if err := json.Unmarshal(got.Plan.Annotations["test/value"], &annotation); err != nil || annotation.Value < 1 || annotation.Value > writers {
t.Fatalf("final annotation=%#v error=%v", annotation, err)
}
}
func TestFilesystemStoreReadersObserveOnlyCompleteRecordsDuringWrites(t *testing.T) {
store := newStore(t, t.TempDir())
if err := store.Save(testRecord(t, 1)); err != nil {
t.Fatal(err)
}
const writers = 12
const readers = 12
errs := make(chan error, writers+readers)
start := make(chan struct{})
var writersDone sync.WaitGroup
for i := 0; i < writers; i++ {
writersDone.Add(1)
go func(value int) {
defer writersDone.Done()
<-start
errs <- store.Save(testRecord(t, value+2))
}(i)
}
for i := 0; i < readers; i++ {
go func() {
<-start
for attempt := 0; attempt < 50; attempt++ {
record, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanHit {
errs <- fmt.Errorf("Load() decision=%#v error=%v", decision, err)
return
}
if err := validateRecord(record, testSourceDigest); err != nil {
errs <- fmt.Errorf("reader observed invalid record: %w", err)
return
}
}
errs <- nil
}()
}
close(start)
writersDone.Wait()
for i := 0; i < readers+writers; i++ {
if err := <-errs; err != nil {
t.Fatal(err)
}
}
}
func TestFilesystemStoreInterruptedWritesPreservePreviousRecord(t *testing.T) {
store := newStore(t, t.TempDir()).(*filesystemStore)
first := testRecord(t, 1)
if err := store.Save(first); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
hooks atomicWriteHooks
}{
{name: "before temporary file", hooks: atomicWriteHooks{BeforeCreateTemp: func() error { return errors.New("interrupted before temporary file") }}},
{name: "before rename", hooks: atomicWriteHooks{BeforeRename: func() error { return errors.New("interrupted before rename") }}},
} {
t.Run(tc.name, func(t *testing.T) {
store.write = func(root *os.Root, target string, data []byte) error {
return writeAtomicWithHooks(root, target, data, tc.hooks)
}
if err := store.Save(testRecord(t, 2)); err == nil {
t.Fatal("Save() error = nil")
}
store.write = writeAtomic
got, decision, err := store.Load(testSourceDigest)
if err != nil || decision.Status != pipeline.ChunkPlanHit || !reflect.DeepEqual(got, first) {
t.Fatalf("record after interruption=%#v decision=%#v error=%v", got, decision, err)
}
})
}
}
func newStore(t *testing.T, root string) pipeline.ChunkPlanStore {
t.Helper()
store, err := NewFilesystemStore(root)
if err != nil {
t.Fatalf("NewFilesystemStore() error = %v", err)
}
return store
}
func testRecord(t *testing.T, value int) pipeline.ChunkPlanRecord {
t.Helper()
annotation, err := json.Marshal(map[string]int{"value": value})
if err != nil {
t.Fatal(err)
}
plan := source.ChunkPlan{
SourceDigest: testSourceDigest,
Ranges: []source.ChunkRange{{StartUnitID: 1, EndUnitID: 2, Annotations: source.ChunkAnnotations{"test/range": json.RawMessage(`{"range":true}`)}}},
Annotations: source.ChunkAnnotations{"test/value": annotation},
}
planDigest, err := source.DigestChunkPlan(plan)
if err != nil {
t.Fatal(err)
}
return pipeline.ChunkPlanRecord{
SchemaVersion: SchemaVersion,
SourceDigest: testSourceDigest,
PlanDigest: planDigest,
Plan: plan,
Producer: pipeline.ChunkPlanProducer{
InputModule: "input/test", ChunkModule: "chunk/test", LLMProfile: "profile/test",
References: []artifacts.ReferenceProvenance{{Stage: "chunk", SlotName: "guide", OriginType: "file", OriginURI: "file:///guide.txt", Digest: "sha256:reference"}},
Metadata: map[string]any{"prompt_id": "test/prompt", "enabled": true},
},
Warnings: []contracts.Warning{{Scope: "chunk/test", ReasonCode: "observed", Message: "warning"}},
CreatedAt: time.Date(2026, 7, 18, 12, 0, 0, 0, time.UTC),
}
}
func replaceJSON(old, replacement string) func([]byte) []byte {
return func(data []byte) []byte { return bytes.Replace(data, []byte(old), []byte(replacement), 1) }
}
func assertNoTemps(t *testing.T, dir string) {
t.Helper()
entries, err := os.ReadDir(dir)
if err != nil {
t.Fatal(err)
}
for _, entry := range entries {
if strings.Contains(entry.Name(), ".tmp-") {
t.Fatalf("temporary file remains: %s", entry.Name())
}
}
}
func createSymlink(t *testing.T, target, link string) {
t.Helper()
if err := os.Symlink(target, link); err != nil {
t.Skipf("create symlink: %v", err)
}
}
func writeFile(t *testing.T, path string, data []byte, mode os.FileMode) {
t.Helper()
if err := os.WriteFile(path, data, mode); err != nil {
t.Fatal(err)
}
if err := os.Chmod(path, mode); err != nil {
t.Fatal(err)
}
}
func readFile(t *testing.T, path string) []byte {
t.Helper()
data, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
return data
}
func fileMode(t *testing.T, path string) os.FileMode {
t.Helper()
info, err := os.Stat(path)
if err != nil {
t.Fatal(err)
}
return info.Mode().Perm()
}
func outsidePlanPath(name, outside string) string {
if name == "digest directory" {
return filepath.Join(outside, "plan.json")
}
return outside
}