| 1 | package control |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "path/filepath" |
| 7 | "reasonix/internal/agent" |
| 8 | "reasonix/internal/agent/testutil" |
| 9 | "reasonix/internal/event" |
| 10 | "reasonix/internal/provider" |
| 11 | "reasonix/internal/tool" |
| 12 | "strings" |
| 13 | "testing" |
| 14 | "time" |
| 15 | ) |
| 16 | |
| 17 | type checkpointProbeTool struct { |
| 18 | name string |
| 19 | started chan struct{} |
| 20 | } |
| 21 | |
| 22 | func (t checkpointProbeTool) Name() string { return t.name } |
| 23 | func (t checkpointProbeTool) Description() string { return "checkpoint test" } |
| 24 | func (t checkpointProbeTool) Schema() json.RawMessage { return json.RawMessage(`{"type":"object"}`) } |
| 25 | func (t checkpointProbeTool) ReadOnly() bool { return false } |
| 26 | func (t checkpointProbeTool) Execute(ctx context.Context, _ json.RawMessage) (string, error) { |
| 27 | if t.started != nil { |
| 28 | close(t.started) |
| 29 | <-ctx.Done() |
| 30 | return "", ctx.Err() |
| 31 | } |
| 32 | return "write completed", nil |
| 33 | } |
| 34 | |
| 35 | func TestToolCheckpointSurvivesReloadWhileNextWriterRuns(t *testing.T) { |
| 36 | started := make(chan struct{}) |
| 37 | reg := tool.NewRegistry() |
| 38 | reg.Add(checkpointProbeTool{name: "first"}) |
| 39 | reg.Add(checkpointProbeTool{name: "second", started: started}) |
| 40 | calls := []provider.ToolCall{{ID: "c1", Name: "first", Arguments: `{}`}, {ID: "c2", Name: "second", Arguments: `{}`}} |
| 41 | mock := testutil.NewMock("test", testutil.Turn{Reasoning: "original reasoning", ToolCalls: calls}) |
| 42 | session := agent.NewSession("system") |
| 43 | exec := agent.New(mock, reg, session, agent.Options{}, event.Discard) |
| 44 | dir := t.TempDir() |
| 45 | path := filepath.Join(dir, "session.jsonl") |
| 46 | sink, done, _ := collectSink() |
| 47 | c := newOwnedTestController(t, Options{Runner: exec, Executor: exec, Sink: sink, SessionDir: dir, SessionPath: path}) |
| 48 | t.Cleanup(c.Close) |
| 49 | c.Submit("run both") |
| 50 | select { |
| 51 | case <-started: |
| 52 | case <-time.After(5 * time.Second): |
| 53 | t.Fatal("second tool did not start") |
| 54 | } |
| 55 | loaded := loadDurableSessionProjection(t, path) |
| 56 | completed := false |
| 57 | for _, m := range loaded.Messages { |
| 58 | if m.Role != provider.RoleTool { |
| 59 | continue |
| 60 | } |
| 61 | if m.ToolCallID == "c1" { |
| 62 | completed = strings.HasPrefix(m.Content, "write completed") && provider.ToolResultRunState(m) == provider.ToolRunCompleted |
| 63 | } |
| 64 | } |
| 65 | if !completed || loaded.ActiveTools["c2"] != "second" { |
| 66 | t.Fatalf("completed=%v active=%v history=%+v", completed, loaded.ActiveTools, loaded.Messages) |
| 67 | } |
| 68 | c.Cancel() |
| 69 | waitForDone(t, done) |
| 70 | } |
| 71 |