| 1 | package control |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "path/filepath" |
| 9 | "testing" |
| 10 | "time" |
| 11 | |
| 12 | "reasonix/internal/agent" |
| 13 | "reasonix/internal/agent/testutil" |
| 14 | "reasonix/internal/event" |
| 15 | "reasonix/internal/provider" |
| 16 | "reasonix/internal/session" |
| 17 | "reasonix/internal/tool" |
| 18 | "reasonix/internal/transcript" |
| 19 | ) |
| 20 | |
| 21 | type followTally struct { |
| 22 | polls, changes, maxBatch int |
| 23 | resets, gaps, changeResets int |
| 24 | } |
| 25 | |
| 26 | // desktopFollower applies transcriptFollowClient's acceptance rules: a |
| 27 | // response reset, a change reset or a skipped revision each cost a baseline. |
| 28 | type desktopFollower struct { |
| 29 | c *Controller |
| 30 | subscription string |
| 31 | revision uint64 |
| 32 | tally followTally |
| 33 | } |
| 34 | |
| 35 | func (f *desktopFollower) baseline(ctx context.Context) error { |
| 36 | response, err := f.c.TranscriptFollow(ctx, transcript.FollowRequest{}) |
| 37 | if err == nil && response.Snapshot == nil { |
| 38 | err = errors.New("baseline without snapshot") |
| 39 | } |
| 40 | if err != nil { |
| 41 | return err |
| 42 | } |
| 43 | f.subscription, f.revision = response.Subscription, response.Snapshot.ProjectionRevision |
| 44 | return nil |
| 45 | } |
| 46 | |
| 47 | // run polls until done is closed and the queue is drained; work stands in for |
| 48 | // the renderer applying one delivered batch before it asks for the next. |
| 49 | func (f *desktopFollower) run(ctx context.Context, done <-chan struct{}, work time.Duration) error { |
| 50 | finished := false |
| 51 | for { |
| 52 | select { |
| 53 | case <-done: |
| 54 | finished = true |
| 55 | default: |
| 56 | } |
| 57 | pollCtx, cancel := context.WithTimeout(ctx, 200*time.Millisecond) |
| 58 | response, err := f.c.TranscriptFollow(pollCtx, transcript.FollowRequest{Subscription: f.subscription, AfterRevision: f.revision}) |
| 59 | cancel() |
| 60 | if err != nil && !errors.Is(err, context.DeadlineExceeded) { |
| 61 | return err |
| 62 | } |
| 63 | f.tally.polls++ |
| 64 | f.tally.changes += len(response.Changes) |
| 65 | f.tally.maxBatch = max(f.tally.maxBatch, len(response.Changes)) |
| 66 | broken := response.ResetRequired |
| 67 | if broken { |
| 68 | f.tally.resets++ |
| 69 | } |
| 70 | for _, change := range response.Changes { |
| 71 | if broken || change.Revision <= f.revision { |
| 72 | continue |
| 73 | } |
| 74 | switch { |
| 75 | case change.ResetRequired: |
| 76 | f.tally.changeResets++ |
| 77 | broken = true |
| 78 | case change.Revision != f.revision+1: |
| 79 | f.tally.gaps++ |
| 80 | broken = true |
| 81 | default: |
| 82 | f.revision = change.Revision |
| 83 | } |
| 84 | } |
| 85 | if broken { |
| 86 | _, _ = f.c.TranscriptFollow(context.Background(), transcript.FollowRequest{Subscription: f.subscription, Close: true}) |
| 87 | if err := f.baseline(ctx); err != nil { |
| 88 | return err |
| 89 | } |
| 90 | continue |
| 91 | } |
| 92 | if finished && len(response.Changes) == 0 { |
| 93 | _, _ = f.c.TranscriptFollow(context.Background(), transcript.FollowRequest{Subscription: f.subscription, Close: true}) |
| 94 | return nil |
| 95 | } |
| 96 | time.Sleep(work) |
| 97 | } |
| 98 | } |
| 99 | |
| 100 | // shellFloodTool writes output the way a verbose build does: many short |
| 101 | // lines, each reaching the agent as its own progress chunk. |
| 102 | type shellFloodTool struct{ lines int } |
| 103 | |
| 104 | func (shellFloodTool) Name() string { return "flood" } |
| 105 | func (shellFloodTool) Description() string { return "prints build output" } |
| 106 | func (shellFloodTool) Schema() json.RawMessage { return json.RawMessage(`{"type":"object"}`) } |
| 107 | func (shellFloodTool) ReadOnly() bool { return true } |
| 108 | func (f shellFloodTool) Execute(ctx context.Context, _ json.RawMessage) (string, error) { |
| 109 | emit, _ := tool.ProgressFrom(ctx) |
| 110 | for i := range f.lines { |
| 111 | if emit != nil { |
| 112 | emit(fmt.Sprintf("compiling unit %d\n", i)) |
| 113 | } |
| 114 | if i%2 == 1 { |
| 115 | time.Sleep(time.Millisecond) |
| 116 | } |
| 117 | } |
| 118 | return "ok", nil |
| 119 | } |
| 120 | |
| 121 | func TestTranscriptFollowSteadyStreamingNeverResetsFollower(t *testing.T) { |
| 122 | chunks := make([]provider.Chunk, 0, 3001) |
| 123 | for range 1500 { |
| 124 | chunks = append(chunks, provider.Chunk{Type: provider.ChunkReasoning, Text: "think "}) |
| 125 | } |
| 126 | for range 1500 { |
| 127 | chunks = append(chunks, provider.Chunk{Type: provider.ChunkText, Text: "word "}) |
| 128 | } |
| 129 | chunks = append(chunks, provider.Chunk{Type: provider.ChunkDone}) |
| 130 | flood := []provider.ToolCall{{ID: "call-flood", Name: "flood", Arguments: "{}"}} |
| 131 | cases := map[string][]testutil.Turn{ |
| 132 | "reasoning and text": {{Chunks: chunks}}, |
| 133 | "shell output": {{ToolCalls: flood}, {Text: "built"}}, |
| 134 | } |
| 135 | for name, turns := range cases { |
| 136 | t.Run(name, func(t *testing.T) { |
| 137 | service, err := session.NewService("desktop", session.NewFilesystemPersistence(filepath.Join(t.TempDir(), "sessions"))) |
| 138 | if err != nil { |
| 139 | t.Fatal(err) |
| 140 | } |
| 141 | runtime, err := service.Create(t.Context(), session.CreateOptions{SessionID: "transcript-follow"}) |
| 142 | if err != nil { |
| 143 | t.Fatal(err) |
| 144 | } |
| 145 | registry := tool.NewRegistry() |
| 146 | registry.Add(shellFloodTool{lines: 3000}) |
| 147 | executor := agent.New(testutil.NewMock("test", turns...), registry, agent.NewSession("system"), agent.Options{}, event.Discard) |
| 148 | c := newOwnedTestController(t, Options{Runner: executor, Executor: executor, Sink: event.Discard, |
| 149 | SessionService: service, SessionRuntime: runtime, ExclusiveSession: true}) |
| 150 | |
| 151 | follower := &desktopFollower{c: c} |
| 152 | if err := follower.baseline(t.Context()); err != nil { |
| 153 | t.Fatal(err) |
| 154 | } |
| 155 | done, result := make(chan struct{}), make(chan error, 1) |
| 156 | go func() { result <- follower.run(t.Context(), done, 400*time.Millisecond) }() |
| 157 | started := time.Now() |
| 158 | if err := c.RunTurn(t.Context(), "question"); err != nil { |
| 159 | t.Fatal(err) |
| 160 | } |
| 161 | elapsed := time.Since(started) |
| 162 | close(done) |
| 163 | if err := <-result; err != nil { |
| 164 | t.Fatal(err) |
| 165 | } |
| 166 | tally := follower.tally |
| 167 | t.Logf("turn=%v polls=%d changes=%d maxBatch=%d resets=%d gaps=%d changeResets=%d", |
| 168 | elapsed, tally.polls, tally.changes, tally.maxBatch, tally.resets, tally.gaps, tally.changeResets) |
| 169 | if tally.resets+tally.gaps+tally.changeResets != 0 { |
| 170 | t.Fatalf("steady output forced %d follower baselines: %+v", tally.resets+tally.gaps+tally.changeResets, tally) |
| 171 | } |
| 172 | }) |
| 173 | } |
| 174 | } |
| 175 |