| 1 | package control |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 6 | "sync/atomic" |
| 7 | "testing" |
| 8 | "time" |
| 9 | |
| 10 | "reasonix/internal/agent" |
| 11 | "reasonix/internal/event" |
| 12 | "reasonix/internal/provider" |
| 13 | "reasonix/internal/tool" |
| 14 | ) |
| 15 | |
| 16 | type checkpointEventRunner struct { |
| 17 | session *agent.Session |
| 18 | err error |
| 19 | started chan struct{} |
| 20 | wait bool |
| 21 | skipUser bool |
| 22 | localOnly bool |
| 23 | } |
| 24 | |
| 25 | func (r *checkpointEventRunner) Run(ctx context.Context, input string) error { |
| 26 | if !r.skipUser { |
| 27 | r.session.Add(provider.Message{ |
| 28 | Role: provider.RoleUser, Content: input, LocalOnly: r.localOnly, |
| 29 | CreatedAt: time.Now().UnixMilli(), |
| 30 | }) |
| 31 | } |
| 32 | if r.started != nil { |
| 33 | close(r.started) |
| 34 | } |
| 35 | if r.wait { |
| 36 | <-ctx.Done() |
| 37 | return ctx.Err() |
| 38 | } |
| 39 | return r.err |
| 40 | } |
| 41 | |
| 42 | func newCheckpointEventController(t *testing.T, runner *checkpointEventRunner) (*Controller, <-chan event.Event) { |
| 43 | t.Helper() |
| 44 | events := make(chan event.Event, 8) |
| 45 | executor := agent.New(nil, tool.NewRegistry(), runner.session, agent.Options{}, event.Discard) |
| 46 | dir := t.TempDir() |
| 47 | controller := newOwnedTestController(t, Options{ |
| 48 | Runner: runner, Executor: executor, |
| 49 | SessionDir: dir, SessionPath: dir + "/session.jsonl", |
| 50 | Sink: event.FuncSink(func(e event.Event) { |
| 51 | if e.Kind == event.TurnDone { |
| 52 | events <- e |
| 53 | } |
| 54 | }), |
| 55 | }) |
| 56 | return controller, events |
| 57 | } |
| 58 | |
| 59 | func receiveCheckpointTurnDone(t *testing.T, events <-chan event.Event) event.Event { |
| 60 | t.Helper() |
| 61 | select { |
| 62 | case e := <-events: |
| 63 | return e |
| 64 | case <-time.After(5 * time.Second): |
| 65 | t.Fatal("timed out waiting for TurnDone") |
| 66 | return event.Event{} |
| 67 | } |
| 68 | } |
| 69 | |
| 70 | func requireCheckpointTurn(t *testing.T, e event.Event, want int) { |
| 71 | t.Helper() |
| 72 | if e.CheckpointTurn == nil || *e.CheckpointTurn != want { |
| 73 | t.Fatalf("TurnDone checkpoint = %v, want %d", e.CheckpointTurn, want) |
| 74 | } |
| 75 | } |
| 76 | |
| 77 | func TestTurnDoneCarriesValidatedCheckpointAcrossSuccessAndError(t *testing.T) { |
| 78 | session := agent.NewSession("system") |
| 79 | runner := &checkpointEventRunner{session: session} |
| 80 | controller, events := newCheckpointEventController(t, runner) |
| 81 | defer controller.Close() |
| 82 | |
| 83 | controller.Send("first prompt") |
| 84 | first := receiveCheckpointTurnDone(t, events) |
| 85 | if first.Err != nil { |
| 86 | t.Fatalf("successful TurnDone error = %v", first.Err) |
| 87 | } |
| 88 | requireCheckpointTurn(t, first, 0) |
| 89 | |
| 90 | runner.err = errors.New("provider failed") |
| 91 | controller.Send("second prompt") |
| 92 | second := receiveCheckpointTurnDone(t, events) |
| 93 | if second.Err == nil { |
| 94 | t.Fatal("provider failure TurnDone must retain its error") |
| 95 | } |
| 96 | requireCheckpointTurn(t, second, 1) |
| 97 | } |
| 98 | |
| 99 | func TestCancelledTurnDoneCarriesRetainedUserCheckpoint(t *testing.T) { |
| 100 | session := agent.NewSession("system") |
| 101 | started := make(chan struct{}) |
| 102 | runner := &checkpointEventRunner{session: session, started: started, wait: true} |
| 103 | controller, events := newCheckpointEventController(t, runner) |
| 104 | defer controller.Close() |
| 105 | |
| 106 | controller.Send("cancel this prompt") |
| 107 | select { |
| 108 | case <-started: |
| 109 | case <-time.After(5 * time.Second): |
| 110 | t.Fatal("runner did not start") |
| 111 | } |
| 112 | controller.Cancel() |
| 113 | done := receiveCheckpointTurnDone(t, events) |
| 114 | if !done.Cancelled || done.Err != nil || done.Status != event.TurnInterrupted || done.Recovery == nil || done.Recovery.State != "interrupted" || done.Recovery.Reason != "silent_interruption" || done.Recovery.RequiresUserDecision { |
| 115 | t.Fatalf("cancelled TurnDone = %+v, want fact-only silent interruption without send error", done) |
| 116 | } |
| 117 | requireCheckpointTurn(t, done, 0) |
| 118 | } |
| 119 | |
| 120 | func TestCancelBeforeRunnerAddsUserCarriesFallbackCheckpoint(t *testing.T) { |
| 121 | session := agent.NewSession("system") |
| 122 | started := make(chan struct{}) |
| 123 | runner := &checkpointEventRunner{session: session, started: started, wait: true, skipUser: true} |
| 124 | controller, events := newCheckpointEventController(t, runner) |
| 125 | defer controller.Close() |
| 126 | |
| 127 | controller.Send("cancel before user append") |
| 128 | select { |
| 129 | case <-started: |
| 130 | case <-time.After(5 * time.Second): |
| 131 | t.Fatal("runner did not start") |
| 132 | } |
| 133 | controller.Cancel() |
| 134 | done := receiveCheckpointTurnDone(t, events) |
| 135 | requireCheckpointTurn(t, done, 0) |
| 136 | messages := session.Snapshot() |
| 137 | if len(messages) < 2 || messages[1].Role != provider.RoleUser || |
| 138 | !agent.IsUserAuthoredTurnMessage(messages[1]) { |
| 139 | t.Fatalf("cancel fallback messages = %+v, want a retained user prompt at the checkpoint boundary", messages) |
| 140 | } |
| 141 | } |
| 142 | |
| 143 | func TestTurnDoneOmitsUncommittedOrNonVisibleCheckpoint(t *testing.T) { |
| 144 | for _, tc := range []struct { |
| 145 | name string |
| 146 | skipUser bool |
| 147 | localOnly bool |
| 148 | }{ |
| 149 | {name: "no user committed", skipUser: true}, |
| 150 | {name: "local-only user", localOnly: true}, |
| 151 | } { |
| 152 | t.Run(tc.name, func(t *testing.T) { |
| 153 | session := agent.NewSession("system") |
| 154 | runner := &checkpointEventRunner{session: session, skipUser: tc.skipUser, localOnly: tc.localOnly} |
| 155 | controller, events := newCheckpointEventController(t, runner) |
| 156 | defer controller.Close() |
| 157 | |
| 158 | controller.Send("blocked prompt") |
| 159 | if done := receiveCheckpointTurnDone(t, events); done.CheckpointTurn != nil { |
| 160 | t.Fatalf("uncommitted checkpoint leaked into TurnDone: %d", *done.CheckpointTurn) |
| 161 | } |
| 162 | }) |
| 163 | } |
| 164 | } |
| 165 | |
| 166 | func TestTurnDoneRejectsCheckpointAfterSessionSwap(t *testing.T) { |
| 167 | oldSession := agent.NewSession("system") |
| 168 | completion := &guardedTurnCompletion{} |
| 169 | ctx := context.WithValue(context.Background(), guardedTurnCompletionKey{}, completion) |
| 170 | runner := &checkpointEventRunner{session: oldSession} |
| 171 | controller, _ := newCheckpointEventController(t, runner) |
| 172 | defer controller.Close() |
| 173 | |
| 174 | controller.beginCheckpoint(ctx, "old prompt") |
| 175 | oldSession.Add(provider.Message{Role: provider.RoleUser, Content: "old prompt", CreatedAt: time.Now().UnixMilli()}) |
| 176 | controller.executor.SetSession(agent.NewSession("replacement")) |
| 177 | if got := controller.validatedCheckpointTurn(completion); got != nil { |
| 178 | t.Fatalf("session-swapped checkpoint = %d, want nil", *got) |
| 179 | } |
| 180 | } |
| 181 | |
| 182 | func TestTurnDoneRejectsSameSessionCheckpointStoreCollision(t *testing.T) { |
| 183 | session := agent.NewSession("system") |
| 184 | completion := &guardedTurnCompletion{} |
| 185 | ctx := context.WithValue(context.Background(), guardedTurnCompletionKey{}, completion) |
| 186 | runner := &checkpointEventRunner{session: session} |
| 187 | controller, _ := newCheckpointEventController(t, runner) |
| 188 | defer controller.Close() |
| 189 | |
| 190 | controller.beginCheckpoint(ctx, "original prompt") |
| 191 | session.Add(provider.Message{Role: provider.RoleUser, Content: "original prompt", CreatedAt: time.Now().UnixMilli()}) |
| 192 | controller.checkpoints.rebind("", "") |
| 193 | if turn, _, ok := controller.checkpoints.beginWithObserver("collision", 1, nil); !ok || turn != 0 { |
| 194 | t.Fatalf("replacement checkpoint = (%d, %v), want colliding turn zero", turn, ok) |
| 195 | } |
| 196 | if got := controller.validatedCheckpointTurn(completion); got != nil { |
| 197 | t.Fatalf("store-rebound checkpoint = %d, want nil", *got) |
| 198 | } |
| 199 | } |
| 200 | |
| 201 | func TestBlockedCandidateDoesNotLeakIntoNextTurn(t *testing.T) { |
| 202 | session := agent.NewSession("system") |
| 203 | runner := &checkpointEventRunner{session: session, skipUser: true} |
| 204 | controller, events := newCheckpointEventController(t, runner) |
| 205 | defer controller.Close() |
| 206 | |
| 207 | controller.Send("blocked before user append") |
| 208 | if done := receiveCheckpointTurnDone(t, events); done.CheckpointTurn != nil { |
| 209 | t.Fatalf("blocked TurnDone checkpoint = %d, want nil", *done.CheckpointTurn) |
| 210 | } |
| 211 | |
| 212 | runner.skipUser = false |
| 213 | controller.Send("next real prompt") |
| 214 | requireCheckpointTurn(t, receiveCheckpointTurnDone(t, events), 1) |
| 215 | } |
| 216 | |
| 217 | func TestParkedTurnsKeepIndependentCheckpointCandidates(t *testing.T) { |
| 218 | session := agent.NewSession("system") |
| 219 | runner := &checkpointEventRunner{session: session} |
| 220 | events := make(chan event.Event, 2) |
| 221 | firstDelivery := make(chan struct{}) |
| 222 | releaseFirst := make(chan struct{}) |
| 223 | var deliveries atomic.Int32 |
| 224 | executor := agent.New(nil, tool.NewRegistry(), session, agent.Options{}, event.Discard) |
| 225 | dir := t.TempDir() |
| 226 | controller := newOwnedTestController(t, Options{ |
| 227 | Runner: runner, Executor: executor, |
| 228 | SessionDir: dir, SessionPath: dir + "/session.jsonl", |
| 229 | Sink: event.FuncSink(func(e event.Event) { |
| 230 | if e.Kind != event.TurnDone { |
| 231 | return |
| 232 | } |
| 233 | if deliveries.Add(1) == 1 { |
| 234 | close(firstDelivery) |
| 235 | <-releaseFirst |
| 236 | } |
| 237 | events <- e |
| 238 | }), |
| 239 | }) |
| 240 | defer controller.Close() |
| 241 | |
| 242 | controller.Send("first prompt") |
| 243 | select { |
| 244 | case <-firstDelivery: |
| 245 | case <-time.After(5 * time.Second): |
| 246 | t.Fatal("first TurnDone delivery did not start") |
| 247 | } |
| 248 | controller.Send("parked second prompt") |
| 249 | close(releaseFirst) |
| 250 | requireCheckpointTurn(t, receiveCheckpointTurnDone(t, events), 0) |
| 251 | requireCheckpointTurn(t, receiveCheckpointTurnDone(t, events), 1) |
| 252 | } |
| 253 |