iop/packages/go/agentstate/store_test.go

516 lines
16 KiB
Go

package agentstate
import (
"context"
"encoding/json"
"errors"
"os"
"path/filepath"
"reflect"
"sync"
"testing"
"time"
"iop/packages/go/agentpolicy"
"iop/packages/go/agentprovider/cli/status"
"iop/packages/go/agentruntime"
"iop/packages/go/agenttask"
)
func TestStoreRoundTripAndStaleCAS(t *testing.T) {
path := filepath.Join(t.TempDir(), "state", "manager.json")
store, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
state, revision, err := store.Load(context.Background())
if err != nil {
t.Fatalf("initial Load: %v", err)
}
if revision != "0" || state.SchemaVersion != agenttask.StateSchemaVersion {
t.Fatalf("initial revision/schema = %q/%d", revision, state.SchemaVersion)
}
state.NextOrdinal = 7
committed, err := store.CompareAndSwap(context.Background(), revision, state)
if err != nil {
t.Fatalf("CompareAndSwap: %v", err)
}
if committed != "1" {
t.Fatalf("committed revision = %q, want 1", committed)
}
reopened, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore reopen: %v", err)
}
got, gotRevision, err := reopened.Load(context.Background())
if err != nil {
t.Fatalf("reopened Load: %v", err)
}
if gotRevision != "1" || got.NextOrdinal != 7 {
t.Fatalf("reopened revision/state = %q/%d", gotRevision, got.NextOrdinal)
}
if _, err := reopened.CompareAndSwap(context.Background(), "0", got); !errors.Is(err, agenttask.ErrRevisionConflict) {
t.Fatalf("stale CompareAndSwap error = %v", err)
}
}
func TestStoreRejectsCorruptionWithoutOverwrite(t *testing.T) {
path := filepath.Join(t.TempDir(), "manager.json")
store, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
state, revision, err := store.Load(context.Background())
if err != nil {
t.Fatalf("Load: %v", err)
}
if _, err := store.CompareAndSwap(context.Background(), revision, state); err != nil {
t.Fatalf("CompareAndSwap: %v", err)
}
payload, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile: %v", err)
}
var envelope map[string]any
if err := json.Unmarshal(payload, &envelope); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
envelope["checksum"] = "tampered"
corrupt, err := json.Marshal(envelope)
if err != nil {
t.Fatalf("Marshal: %v", err)
}
if err := os.WriteFile(path, corrupt, 0o600); err != nil {
t.Fatalf("WriteFile: %v", err)
}
before, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile before rejected CAS: %v", err)
}
if _, _, err := store.Load(context.Background()); !errors.Is(err, ErrCorruptState) {
t.Fatalf("corrupt Load error = %v", err)
}
if _, err := store.CompareAndSwap(context.Background(), "1", state); !errors.Is(err, ErrCorruptState) {
t.Fatalf("corrupt CompareAndSwap error = %v", err)
}
after, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile after rejected CAS: %v", err)
}
if string(after) != string(before) {
t.Fatal("rejected CAS overwrote corrupt checkpoint evidence")
}
}
func TestStoreConcurrentCASIsSerialized(t *testing.T) {
path := filepath.Join(t.TempDir(), "manager.json")
const writers = 12
var wait sync.WaitGroup
errs := make(chan error, writers)
for range writers {
wait.Add(1)
go func() {
defer wait.Done()
store, err := NewStore(path)
if err != nil {
errs <- err
return
}
for {
state, revision, err := store.Load(context.Background())
if err != nil {
errs <- err
return
}
state.NextOrdinal++
if _, err := store.CompareAndSwap(context.Background(), revision, state); errors.Is(err, agenttask.ErrRevisionConflict) {
continue
} else if err != nil {
errs <- err
}
return
}
}()
}
wait.Wait()
close(errs)
for err := range errs {
t.Errorf("concurrent writer: %v", err)
}
store, _ := NewStore(path)
state, revision, err := store.Load(context.Background())
if err != nil {
t.Fatalf("final Load: %v", err)
}
if state.NextOrdinal != writers || revision != "12" {
t.Fatalf("final ordinal/revision = %d/%s, want %d/12", state.NextOrdinal, revision, writers)
}
}
func TestStoreIntegrationRecordSnapshot(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "state", "manager.json")
store, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
payloads := map[string][]byte{
"projectlog:one": []byte(`{"value":"one"}`),
"projectlog:two": []byte(`{"value":"two"}`),
"other:three": []byte(`{"value":"three"}`),
}
for key, payload := range payloads {
if _, err := store.CompareAndSwapIntegrationRecord(ctx, key, "", payload); err != nil {
t.Fatalf("seed %s: %v", key, err)
}
}
snapshot, err := store.LoadIntegrationRecords(ctx, "projectlog:")
if err != nil {
t.Fatalf("LoadIntegrationRecords: %v", err)
}
if len(snapshot) != 2 {
t.Fatalf("snapshot size = %d, want 2", len(snapshot))
}
for _, key := range []string{"projectlog:one", "projectlog:two"} {
entry, ok := snapshot[key]
if !ok {
t.Fatalf("snapshot missing %q", key)
}
if string(entry.Payload) != string(payloads[key]) ||
entry.Revision != integrationRecordRevision(payloads[key]) {
t.Fatalf("snapshot %q = %+v", key, entry)
}
}
snapshot["projectlog:one"].Payload[0] = '['
reloaded, err := store.LoadIntegrationRecords(ctx, "projectlog:")
if err != nil {
t.Fatalf("reload snapshot: %v", err)
}
if string(reloaded["projectlog:one"].Payload) != string(payloads["projectlog:one"]) {
t.Fatal("caller mutation changed retained integration record")
}
}
func TestStoreIntegrationRecordBatchCAS(t *testing.T) {
ctx := context.Background()
t.Run("two key commit preserves sibling and survives reopen", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "manager.json")
store, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
sibling := []byte(`{"value":"sibling"}`)
if _, err := store.CompareAndSwapIntegrationRecord(ctx, "other:sibling", "", sibling); err != nil {
t.Fatalf("seed sibling: %v", err)
}
updates := []IntegrationRecordUpdate{
{Key: "projectlog:index", Payload: []byte(`{"value":"index"}`)},
{Key: "projectlog:journal", Payload: []byte(`{"value":"journal"}`)},
}
revisions, err := store.CompareAndSwapIntegrationRecords(ctx, updates)
if err != nil {
t.Fatalf("CompareAndSwapIntegrationRecords: %v", err)
}
for _, update := range updates {
if revisions[update.Key] != integrationRecordRevision(update.Payload) {
t.Fatalf("revision %q = %q", update.Key, revisions[update.Key])
}
}
reopened, err := NewStore(path)
if err != nil {
t.Fatalf("reopen NewStore: %v", err)
}
for _, update := range append(updates, IntegrationRecordUpdate{
Key: "other:sibling", Payload: sibling,
}) {
payload, _, found, err := reopened.LoadIntegrationRecord(ctx, update.Key)
if err != nil || !found || string(payload) != string(update.Payload) {
t.Fatalf("reopened %q = %q, found=%t, err=%v", update.Key, payload, found, err)
}
}
})
t.Run("one stale revision rejects every update", func(t *testing.T) {
store, err := NewStore(filepath.Join(t.TempDir(), "manager.json"))
if err != nil {
t.Fatalf("NewStore: %v", err)
}
first := []byte(`{"version":1}`)
firstRevision, err := store.CompareAndSwapIntegrationRecord(ctx, "batch:first", "", first)
if err != nil {
t.Fatalf("seed first: %v", err)
}
second := []byte(`{"version":1}`)
secondRevision, err := store.CompareAndSwapIntegrationRecord(ctx, "batch:second", "", second)
if err != nil {
t.Fatalf("seed second: %v", err)
}
advanced := []byte(`{"version":2}`)
if _, err := store.CompareAndSwapIntegrationRecord(
ctx,
"batch:first",
firstRevision,
advanced,
); err != nil {
t.Fatalf("advance first: %v", err)
}
_, err = store.CompareAndSwapIntegrationRecords(ctx, []IntegrationRecordUpdate{
{Key: "batch:first", Expected: firstRevision, Payload: []byte(`{"version":3}`)},
{Key: "batch:second", Expected: secondRevision, Payload: []byte(`{"version":2}`)},
})
if !errors.Is(err, agenttask.ErrRevisionConflict) {
t.Fatalf("stale batch error = %v", err)
}
firstPayload, _, _, _ := store.LoadIntegrationRecord(ctx, "batch:first")
secondPayload, _, _, _ := store.LoadIntegrationRecord(ctx, "batch:second")
if string(firstPayload) != string(advanced) || string(secondPayload) != string(second) {
t.Fatalf("stale batch partially mutated records: first=%s second=%s", firstPayload, secondPayload)
}
})
t.Run("invalid requests are rejected before mutation", func(t *testing.T) {
store, err := NewStore(filepath.Join(t.TempDir(), "manager.json"))
if err != nil {
t.Fatalf("NewStore: %v", err)
}
tests := map[string][]IntegrationRecordUpdate{
"empty": nil,
"duplicate key": {
{Key: "duplicate:key", Payload: []byte(`{"value":1}`)},
{Key: "duplicate:key", Payload: []byte(`{"value":2}`)},
},
"invalid key": {
{Key: " invalid", Payload: []byte(`{"value":1}`)},
},
"invalid json": {
{Key: "invalid:json", Payload: []byte(`{"value":`)},
},
}
for name, updates := range tests {
t.Run(name, func(t *testing.T) {
if _, err := store.CompareAndSwapIntegrationRecords(ctx, updates); err == nil {
t.Fatal("invalid batch succeeded")
}
})
}
snapshot, err := store.LoadIntegrationRecords(ctx, "duplicate:")
if err != nil {
t.Fatalf("LoadIntegrationRecords: %v", err)
}
if len(snapshot) != 0 {
t.Fatalf("invalid request mutated records: %+v", snapshot)
}
})
t.Run("competing shared revision has one winner", func(t *testing.T) {
path := filepath.Join(t.TempDir(), "manager.json")
seed, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
indexRevision, err := seed.CompareAndSwapIntegrationRecord(
ctx,
"replay:index",
"",
[]byte(`{"owner":"none"}`),
)
if err != nil {
t.Fatalf("seed index: %v", err)
}
start := make(chan struct{})
results := make(chan error, 2)
for _, owner := range []string{"one", "two"} {
owner := owner
go func() {
store, openErr := NewStore(path)
if openErr != nil {
results <- openErr
return
}
<-start
_, updateErr := store.CompareAndSwapIntegrationRecords(ctx, []IntegrationRecordUpdate{
{
Key: "replay:index",
Expected: indexRevision,
Payload: []byte(`{"owner":"` + owner + `"}`),
},
{
Key: "replay:journal:" + owner,
Payload: []byte(`{"owner":"` + owner + `"}`),
},
})
results <- updateErr
}()
}
close(start)
var successes, conflicts int
for range 2 {
err := <-results
switch {
case err == nil:
successes++
case errors.Is(err, agenttask.ErrRevisionConflict):
conflicts++
default:
t.Fatalf("competing writer: %v", err)
}
}
if successes != 1 || conflicts != 1 {
t.Fatalf("successes/conflicts = %d/%d, want 1/1", successes, conflicts)
}
journals, err := seed.LoadIntegrationRecords(ctx, "replay:journal:")
if err != nil {
t.Fatalf("load journals: %v", err)
}
if len(journals) != 1 {
t.Fatalf("committed journals = %d, want 1", len(journals))
}
})
}
func TestStorePersistsSealedQuotaObservation(t *testing.T) {
path := filepath.Join(t.TempDir(), "manager.json")
store, err := NewStore(path)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
now := time.Date(2026, 7, 29, 1, 2, 3, 0, time.UTC)
attempt := agenttask.AttemptID("attempt-v1/4:work/1:1")
locator := agenttask.LocatorRecord{
Kind: agenttask.LocatorProcess, Opaque: "pid:42:start:9001", Revision: "process-r1",
ProjectID: "project", WorkspaceID: "workspace",
WorkUnitID: "work", AttemptID: attempt,
}
sealedObservation := agentpolicy.NormalizeAttemptObservation(
status.NormalizeQuotaSnapshot(
"provider",
"profile",
[]string{"overall"},
now,
&status.UsageStatus{DailyLimit: "50%"},
nil,
),
&agentruntime.Failure{Code: agentruntime.FailureCodeUnavailable, Retryable: true},
now,
time.Minute,
)
state := agenttask.ManagerState{
SchemaVersion: agenttask.StateSchemaVersion,
DeviceLease: &agenttask.LeaseRecord{
OwnerID: "daemon", Token: "device-token", ExpiresAt: now.Add(time.Minute),
},
WorkspaceLeases: map[agenttask.WorkspaceID]agenttask.LeaseRecord{
"workspace": {
OwnerID: "daemon", Token: "workspace-token", ExpiresAt: now.Add(time.Minute),
},
},
Projects: map[agenttask.ProjectID]agenttask.ProjectRecord{
"project": {
ProjectID: "project", WorkspaceID: "workspace",
Status: agenttask.ProjectStatusRunning,
Works: map[agenttask.WorkUnitID]agenttask.WorkRecord{
"work": {
Unit: agenttask.WorkUnit{ID: "work", MilestoneID: "milestone"},
State: agenttask.WorkStateDispatching, Attempt: 1, AttemptID: attempt,
Locators: map[agenttask.LocatorKind]agenttask.LocatorRecord{
agenttask.LocatorProcess: locator,
},
FailureBudgets: map[agenttask.FailureStage]agenttask.FailureBudgetRecord{
agenttask.FailureStageDispatch: {
Stage: agenttask.FailureStageDispatch, Consecutive: 2, Limit: 10,
LastCode: agenttask.BlockerInvocationFailed,
AttemptID: attempt, UpdatedAt: now,
},
},
AttemptObservations: []agenttask.AttemptObservationRecord{{
AttemptID: attempt,
Target: agenttask.ExecutionTarget{
ProviderID: "provider", ModelID: "model", ProfileID: "profile",
ProfileRevision: "profile-r1", ConfigRevision: "config-r1", Capacity: 1,
},
Observation: sealedObservation,
}},
},
},
},
},
}
_, revision, err := store.Load(context.Background())
if err != nil {
t.Fatalf("Load: %v", err)
}
if _, err := store.CompareAndSwap(context.Background(), revision, state); err != nil {
t.Fatalf("CompareAndSwap: %v", err)
}
got, _, err := store.Load(context.Background())
if err != nil {
t.Fatalf("Load committed state: %v", err)
}
if !reflect.DeepEqual(got, state) {
t.Fatalf("recovery state changed across disk round trip\ngot: %#v\nwant: %#v", got, state)
}
payload, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile: %v", err)
}
var envelope diskEnvelope
if err := decodeOne(payload, &envelope); err != nil {
t.Fatalf("decode envelope: %v", err)
}
var durableState map[string]any
if err := json.Unmarshal(envelope.State, &durableState); err != nil {
t.Fatalf("decode state object: %v", err)
}
quota := durableState["Projects"].(map[string]any)["project"].(map[string]any)["Works"].(map[string]any)["work"].(map[string]any)["AttemptObservations"].([]any)[0].(map[string]any)["Observation"].(map[string]any)["Quota"].(map[string]any)
quota["state"] = "exhausted"
envelope.State, err = json.Marshal(durableState)
if err != nil {
t.Fatalf("encode tampered state: %v", err)
}
envelope.Checksum, err = stateChecksum(envelope.SchemaVersion, envelope.Revision, envelope.State)
if err != nil {
t.Fatalf("stateChecksum: %v", err)
}
payload, err = json.Marshal(envelope)
if err != nil {
t.Fatalf("encode tampered envelope: %v", err)
}
if err := os.WriteFile(path, payload, 0o600); err != nil {
t.Fatalf("WriteFile tampered state: %v", err)
}
if _, _, err := store.Load(context.Background()); !errors.Is(err, ErrCorruptState) {
t.Fatalf("Load checksum-valid seal drift error = %v, want corrupt state", err)
}
}
func TestStoreRejectsUnsupportedEnvelopeSchema(t *testing.T) {
path := filepath.Join(t.TempDir(), "manager.json")
state, err := json.Marshal(agenttask.ManagerState{SchemaVersion: agenttask.StateSchemaVersion})
if err != nil {
t.Fatalf("Marshal state: %v", err)
}
checksum, err := stateChecksum(99, 1, state)
if err != nil {
t.Fatalf("stateChecksum: %v", err)
}
payload, err := json.Marshal(diskEnvelope{
SchemaVersion: 99, Revision: 1, Checksum: checksum, State: state,
})
if err != nil {
t.Fatalf("Marshal envelope: %v", err)
}
if err := os.WriteFile(path, payload, 0o600); err != nil {
t.Fatalf("WriteFile: %v", err)
}
store, _ := NewStore(path)
if _, _, err := store.Load(context.Background()); !errors.Is(err, ErrUnsupportedSchema) {
t.Fatalf("Load error = %v, want unsupported schema", err)
}
}