| 1 | package agent |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "path/filepath" |
| 9 | "strings" |
| 10 | "sync/atomic" |
| 11 | "testing" |
| 12 | |
| 13 | "reasonix/internal/agent/testutil" |
| 14 | "reasonix/internal/event" |
| 15 | "reasonix/internal/provider" |
| 16 | "reasonix/internal/tool" |
| 17 | ) |
| 18 | |
| 19 | func echoRegistry() *tool.Registry { |
| 20 | reg := tool.NewRegistry() |
| 21 | reg.Add(echoTool{}) |
| 22 | return reg |
| 23 | } |
| 24 | |
| 25 | func TestStreamIdentityMatchesPersistedAssistant(t *testing.T) { |
| 26 | prov := testutil.NewMock("m", testutil.Turn{Text: "answer"}) |
| 27 | session := NewSession("system") |
| 28 | sink := &recordSink{} |
| 29 | a := New(prov, tool.NewRegistry(), session, Options{}, sink) |
| 30 | if err := a.Run(withNoClosedLoop(context.Background()), "question"); err != nil { |
| 31 | t.Fatal(err) |
| 32 | } |
| 33 | messages := session.Snapshot() |
| 34 | users := sink.kinds(event.UserMessage) |
| 35 | if len(users) != 1 || users[0].MessageID == "" || users[0].Text != "question" { |
| 36 | t.Fatalf("missing admitted user identity: %+v", users) |
| 37 | } |
| 38 | if messages[len(messages)-2].ID != users[0].MessageID { |
| 39 | t.Fatal("user event does not identify the persisted user message") |
| 40 | } |
| 41 | assistant := messages[len(messages)-1] |
| 42 | if assistant.Role != provider.RoleAssistant || assistant.ID == "" { |
| 43 | t.Fatalf("missing persisted assistant identity: %+v", assistant) |
| 44 | } |
| 45 | for _, kind := range []event.Kind{event.Text, event.Message, event.StreamAttempt} { |
| 46 | events := sink.kinds(kind) |
| 47 | if len(events) == 0 { |
| 48 | t.Fatalf("no %v events", kind) |
| 49 | } |
| 50 | for _, e := range events { |
| 51 | if e.MessageID != assistant.ID || e.AttemptID != assistant.ID { |
| 52 | t.Fatalf("event identity differs from committed message %q: %+v", assistant.ID, e) |
| 53 | } |
| 54 | } |
| 55 | } |
| 56 | } |
| 57 | |
| 58 | func TestToolEventsRetainCommittedMessageIdentityAcrossRounds(t *testing.T) { |
| 59 | prov := testutil.NewMock("m", |
| 60 | testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "first", Name: "echo", Arguments: `{"text":"one"}`}}}, |
| 61 | testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "second", Name: "echo", Arguments: `{"text":"two"}`}}}, |
| 62 | testutil.Turn{Text: "done"}, |
| 63 | ) |
| 64 | session := NewSession("system") |
| 65 | sink := &recordSink{} |
| 66 | a := New(prov, echoRegistry(), session, Options{}, sink) |
| 67 | if err := a.Run(withNoClosedLoop(context.Background()), "run both"); err != nil { |
| 68 | t.Fatal(err) |
| 69 | } |
| 70 | owners := map[string]string{} |
| 71 | for _, m := range session.Snapshot() { |
| 72 | for _, call := range m.ToolCalls { |
| 73 | owners[call.ID] = m.ID |
| 74 | } |
| 75 | } |
| 76 | if owners["first"] == "" || owners["second"] == "" || owners["first"] == owners["second"] { |
| 77 | t.Fatalf("invalid committed owners: %v", owners) |
| 78 | } |
| 79 | for _, kind := range []event.Kind{event.ToolDispatch, event.ToolResult} { |
| 80 | for _, e := range sink.kinds(kind) { |
| 81 | if expected := owners[e.Tool.ID]; expected != "" && e.MessageID != expected { |
| 82 | t.Fatalf("tool %s event %v owner %q, want %q", e.Tool.ID, kind, e.MessageID, expected) |
| 83 | } |
| 84 | } |
| 85 | } |
| 86 | } |
| 87 | |
| 88 | func TestRunPersistsUserCreatedAtWithoutSendingItToProvider(t *testing.T) { |
| 89 | const existingCreatedAt int64 = 1_718_000_000_000 |
| 90 | prov := testutil.NewMock("m", testutil.Turn{Text: "done"}) |
| 91 | session := NewSession("system") |
| 92 | session.Add(provider.Message{Role: provider.RoleUser, Content: "existing", CreatedAt: existingCreatedAt}) |
| 93 | agent := New(prov, tool.NewRegistry(), session, Options{}, event.Discard) |
| 94 | |
| 95 | if err := agent.Run(withNoClosedLoop(context.Background()), "new prompt"); err != nil { |
| 96 | t.Fatalf("Run: %v", err) |
| 97 | } |
| 98 | request := prov.LastRequest() |
| 99 | if request == nil { |
| 100 | t.Fatal("provider received no request") |
| 101 | } |
| 102 | for i, message := range request.Messages { |
| 103 | if message.CreatedAt != 0 { |
| 104 | t.Fatalf("provider message %d leaked createdAt %d", i, message.CreatedAt) |
| 105 | } |
| 106 | } |
| 107 | |
| 108 | messages := session.Snapshot() |
| 109 | if len(messages) < 3 || messages[1].CreatedAt != existingCreatedAt { |
| 110 | t.Fatalf("persisted existing timestamp changed: %+v", messages) |
| 111 | } |
| 112 | if messages[2].Role != provider.RoleUser || messages[2].CreatedAt <= 0 { |
| 113 | t.Fatalf("new user timestamp was not persisted: %+v", messages[2]) |
| 114 | } |
| 115 | } |
| 116 | |
| 117 | func TestRunPersistsResponsesItemsAcrossSessionReload(t *testing.T) { |
| 118 | raw := json.RawMessage(`{"id":"ws_1","type":"web_search_call","status":"completed","action":{"type":"search","query":"latest"}}`) |
| 119 | prov := testutil.NewMock("deepseek-responses", testutil.Turn{Chunks: []provider.Chunk{ |
| 120 | {Type: provider.ChunkResponsesItem, ResponsesItem: raw}, |
| 121 | {Type: provider.ChunkText, Text: "answer"}, |
| 122 | {Type: provider.ChunkDone}, |
| 123 | }}) |
| 124 | session := NewSession("system") |
| 125 | agent := New(prov, tool.NewRegistry(), session, Options{}, event.Discard) |
| 126 | if err := agent.Run(withNoClosedLoop(context.Background()), "search"); err != nil { |
| 127 | t.Fatalf("Run: %v", err) |
| 128 | } |
| 129 | |
| 130 | messages := session.Snapshot() |
| 131 | assistant := messages[len(messages)-1] |
| 132 | if assistant.Role != provider.RoleAssistant || len(assistant.ResponsesItems) != 1 || string(assistant.ResponsesItems[0]) != string(raw) { |
| 133 | t.Fatalf("assistant Responses items = %#v, want persisted search item", assistant.ResponsesItems) |
| 134 | } |
| 135 | |
| 136 | path := filepath.Join(t.TempDir(), "responses-items.jsonl") |
| 137 | if err := session.Save(path); err != nil { |
| 138 | t.Fatalf("Save: %v", err) |
| 139 | } |
| 140 | loaded, err := LoadSession(path) |
| 141 | if err != nil { |
| 142 | t.Fatalf("LoadSession: %v", err) |
| 143 | } |
| 144 | loadedAssistant := loaded.Messages[len(loaded.Messages)-1] |
| 145 | if len(loadedAssistant.ResponsesItems) != 1 || string(loadedAssistant.ResponsesItems[0]) != string(raw) { |
| 146 | t.Fatalf("reloaded Responses items = %#v, want original item", loadedAssistant.ResponsesItems) |
| 147 | } |
| 148 | } |
| 149 | |
| 150 | // TestRunMultiToolRoundEmptyIDsSurvivePairing drives the real loop through a turn |
| 151 | // that fans out two tool calls carrying no id (a gateway that streams by index), |
| 152 | // then asserts both results still pair back after SanitizeToolPairing — the repair |
| 153 | // that runs on every send. Keying on tool_call_id alone collapsed them into one, |
| 154 | // dropping a result from the model's context on the very next turn. |
| 155 | func TestRunMultiToolRoundEmptyIDsSurvivePairing(t *testing.T) { |
| 156 | mp := testutil.NewMock("m", |
| 157 | testutil.Turn{ToolCalls: []provider.ToolCall{ |
| 158 | {ID: "", Name: "echo", Arguments: `{"text":"alpha"}`}, |
| 159 | {ID: "", Name: "echo", Arguments: `{"text":"beta"}`}, |
| 160 | }}, |
| 161 | testutil.Turn{Text: "done"}, |
| 162 | ) |
| 163 | a := New(mp, echoRegistry(), NewSession(""), Options{}, event.Discard) |
| 164 | if err := a.Run(withNoClosedLoop(context.Background()), "go"); err != nil { |
| 165 | t.Fatalf("Run: %v", err) |
| 166 | } |
| 167 | |
| 168 | repaired := provider.SanitizeToolPairing(a.Session().Messages) |
| 169 | var results []string |
| 170 | for _, m := range repaired { |
| 171 | if m.Role == provider.RoleTool { |
| 172 | results = append(results, m.Content) |
| 173 | } |
| 174 | } |
| 175 | if len(results) != 2 { |
| 176 | t.Fatalf("want 2 tool results after pairing, got %d: %v", len(results), results) |
| 177 | } |
| 178 | if results[0] == results[1] { |
| 179 | t.Fatalf("both results collapsed to %q — one was lost from the model's context", results[0]) |
| 180 | } |
| 181 | if !strings.Contains(results[0], "alpha") || !strings.Contains(results[1], "beta") { |
| 182 | t.Errorf("results lost their identity: %v", results) |
| 183 | } |
| 184 | } |
| 185 | |
| 186 | func TestRunPersistsCumulativeAssistantWorkDuration(t *testing.T) { |
| 187 | mp := testutil.NewMock("m", |
| 188 | testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "call-1", Name: "echo", Arguments: `{"text":"hello"}`}}}, |
| 189 | testutil.Turn{Text: "done"}, |
| 190 | ) |
| 191 | a := New(mp, echoRegistry(), NewSession(""), Options{}, event.Discard) |
| 192 | if err := a.Run(withNoClosedLoop(context.Background()), "go"); err != nil { |
| 193 | t.Fatalf("Run: %v", err) |
| 194 | } |
| 195 | |
| 196 | var durations []int64 |
| 197 | for _, message := range a.Session().Messages { |
| 198 | if message.Role == provider.RoleAssistant { |
| 199 | durations = append(durations, message.WorkDurationMs) |
| 200 | } |
| 201 | } |
| 202 | if len(durations) != 2 { |
| 203 | t.Fatalf("assistant durations = %v, want two rounds", durations) |
| 204 | } |
| 205 | if durations[0] <= 0 || durations[1] < durations[0] { |
| 206 | t.Fatalf("assistant durations must be positive and cumulative: %v", durations) |
| 207 | } |
| 208 | } |
| 209 | |
| 210 | // TestRunCancelledMidStreamLeavesResumableSession proves a turn cancelled before |
| 211 | // the model answered leaves the session well-formed: the user message stands, |
| 212 | // nothing dangling, and the repaired history is sendable as-is on resume. |
| 213 | func TestRunCancelledMidStreamLeavesResumableSession(t *testing.T) { |
| 214 | mp := testutil.NewMock("m", testutil.ErrorTurn(context.Canceled)) |
| 215 | a := New(mp, echoRegistry(), NewSession("sys"), Options{}, event.Discard) |
| 216 | |
| 217 | err := a.Run(withNoClosedLoop(context.Background()), "do the thing") |
| 218 | if !errors.Is(err, context.Canceled) { |
| 219 | t.Fatalf("Run should surface the cancellation, got %v", err) |
| 220 | } |
| 221 | |
| 222 | repaired := provider.SanitizeToolPairing(a.Session().Messages) |
| 223 | for i, m := range repaired { |
| 224 | if m.Role == provider.RoleTool { |
| 225 | t.Fatalf("a cancelled turn left a dangling tool message at %d: %+v", i, m) |
| 226 | } |
| 227 | } |
| 228 | last := repaired[len(repaired)-1] |
| 229 | if last.Role != provider.RoleUser || StripTransientUserBlocks(last.Content) != "do the thing" { |
| 230 | t.Errorf("the pending user message should survive a cancel, got %+v", last) |
| 231 | } |
| 232 | } |
| 233 | |
| 234 | func TestInterruptedStreamStopsUntilUserRetries(t *testing.T) { |
| 235 | interrupted := &provider.StreamInterruptedError{Err: errors.New("unexpected EOF"), Reason: provider.StreamInterruptPrematureEOF} |
| 236 | for _, partialTool := range []bool{false, true} { |
| 237 | t.Run(fmt.Sprintf("partialTool=%v", partialTool), func(t *testing.T) { |
| 238 | first := testutil.Turn{Text: "partial ", ChunkError: interrupted} |
| 239 | if partialTool { |
| 240 | first = testutil.Turn{Chunks: []provider.Chunk{ |
| 241 | {Type: provider.ChunkToolCallStart, ToolCall: &provider.ToolCall{ID: "incomplete", Name: "echo"}}, |
| 242 | {Type: provider.ChunkError, Err: interrupted}, |
| 243 | }} |
| 244 | } |
| 245 | p := testutil.NewMock("m", first, testutil.Turn{Text: "continued"}) |
| 246 | sink := &recordSink{} |
| 247 | a := New(p, echoRegistry(), NewSession(""), Options{}, sink) |
| 248 | if err := a.Run(withNoClosedLoop(t.Context()), "go"); !errors.Is(err, interrupted) { |
| 249 | t.Fatalf("lost stream failure: %v", err) |
| 250 | } |
| 251 | if p.CallCount() != 1 || len(sink.kinds(event.Retrying)) != 0 || len(sink.kinds(event.ToolResult)) != 0 { |
| 252 | t.Fatalf("calls=%d retries=%v tools=%v", p.CallCount(), sink.kinds(event.Retrying), sink.kinds(event.ToolResult)) |
| 253 | } |
| 254 | if !partialTool { |
| 255 | found := false |
| 256 | for _, m := range a.Session().Snapshot() { |
| 257 | if m.LocalOnly && m.Content == "partial " { |
| 258 | found = true |
| 259 | } |
| 260 | } |
| 261 | if !found { |
| 262 | t.Fatal("partial reply lost") |
| 263 | } |
| 264 | } |
| 265 | if err := a.Run(withNoClosedLoop(t.Context()), "try again"); err != nil { |
| 266 | t.Fatal(err) |
| 267 | } |
| 268 | if p.CallCount() != 2 || lastAssistantContent(a.Session()) != "continued" { |
| 269 | t.Fatalf("manual retry calls=%d", p.CallCount()) |
| 270 | } |
| 271 | for _, m := range provider.ModelMessages(p.Requests()[1].Messages) { |
| 272 | if m.Content == "partial " || m.ToolCallID == "incomplete" { |
| 273 | t.Fatalf("uncommitted output leaked into manual retry: %+v", m) |
| 274 | } |
| 275 | } |
| 276 | }) |
| 277 | } |
| 278 | } |
| 279 | |
| 280 | func TestInterruptedStreamAccountsOnlyAttemptedRequest(t *testing.T) { |
| 281 | interrupted := provider.StreamInterrupt(errors.New("eof"), provider.StreamInterruptPrematureEOF) |
| 282 | p := testutil.NewMock("m", testutil.Turn{Text: "a", Usage: &provider.Usage{PromptTokens: 30, CompletionTokens: 1, TotalTokens: 31, CacheMissTokens: 30}, ChunkError: interrupted}, testutil.Turn{Text: "must not retry"}) |
| 283 | sink := &recordSink{} |
| 284 | a := New(p, echoRegistry(), NewSession(""), Options{}, sink) |
| 285 | if err := a.Run(withNoClosedLoop(t.Context()), "go"); !errors.Is(err, interrupted) { |
| 286 | t.Fatal(err) |
| 287 | } |
| 288 | usages := sink.kinds(event.Usage) |
| 289 | if p.CallCount() != 1 || len(usages) != 1 || usages[0].Usage == nil { |
| 290 | t.Fatalf("calls=%d usage=%+v", p.CallCount(), usages) |
| 291 | } |
| 292 | u := usages[0].Usage |
| 293 | if u.RequestCount != 1 || u.PromptTokens != 30 || u.CompletionTokens != 1 || u.CacheHitTokens+u.CacheMissTokens != u.PromptTokens { |
| 294 | t.Fatalf("usage=%+v", u) |
| 295 | } |
| 296 | if last := a.sess.output.lastUsage.Load(); last == nil || last.PromptTokens != 30 { |
| 297 | t.Fatalf("latest usage=%+v", last) |
| 298 | } |
| 299 | } |
| 300 | |
| 301 | func TestRunInterruptedStreamPersistsPendingLocalOnly(t *testing.T) { |
| 302 | interrupted := &provider.StreamInterruptedError{Err: errors.New("eof"), Reason: provider.StreamInterruptPrematureEOF} |
| 303 | turns := make([]testutil.Turn, 0, maxSamplingAttempts) |
| 304 | for range maxSamplingAttempts { |
| 305 | turns = append(turns, testutil.Turn{Text: "half", ChunkError: interrupted}) |
| 306 | } |
| 307 | mp := testutil.NewMock("m", turns...) |
| 308 | a := New(mp, echoRegistry(), NewSession(""), Options{}, event.Discard) |
| 309 | |
| 310 | err := a.Run(withNoClosedLoop(context.Background()), "go") |
| 311 | if !provider.IsStreamInterrupted(err) { |
| 312 | t.Fatalf("Run error = %v, want original StreamInterruptedError", err) |
| 313 | } |
| 314 | if mp.CallCount() != 1 { |
| 315 | t.Fatalf("provider calls = %d, want 1", mp.CallCount()) |
| 316 | } |
| 317 | var pending *provider.InterruptedTurnRecovery |
| 318 | var local provider.Message |
| 319 | for _, m := range a.Session().Messages { |
| 320 | if m.LocalOnly && m.InterruptedTurn != nil && m.InterruptedTurn.Pending { |
| 321 | pending = m.InterruptedTurn |
| 322 | local = m |
| 323 | } |
| 324 | } |
| 325 | if pending == nil || local.Content != "half" { |
| 326 | t.Fatalf("interruption must leave one pending LocalOnly record: local=%+v pending=%+v", local, pending) |
| 327 | } |
| 328 | // No synthetic recovery user messages mid-turn. |
| 329 | for _, m := range a.Session().Messages { |
| 330 | if m.Role == provider.RoleUser && strings.Contains(m.Content, "previous assistant response was interrupted") { |
| 331 | t.Fatalf("must not inject synthetic stream recovery: %+v", m) |
| 332 | } |
| 333 | } |
| 334 | } |
| 335 | |
| 336 | func TestRunCompleteUncommittedToolCallNeverExecutes(t *testing.T) { |
| 337 | // Full tool block arrived, but the stream was interrupted before a clean |
| 338 | // terminal — the call stays speculative and must never reach executeBatch. |
| 339 | interrupted := &provider.StreamInterruptedError{Err: errors.New("eof"), Reason: provider.StreamInterruptPrematureEOF} |
| 340 | writer := &countingWriterTool{} |
| 341 | reg := tool.NewRegistry() |
| 342 | reg.Add(writer) |
| 343 | mp := testutil.NewMock("m", |
| 344 | testutil.Turn{Chunks: []provider.Chunk{ |
| 345 | {Type: provider.ChunkToolCall, ToolCall: &provider.ToolCall{ID: "w1", Name: "write_file", Arguments: `{"path":"x.txt","content":"from-writer"}`}}, |
| 346 | {Type: provider.ChunkError, Err: interrupted}, |
| 347 | }}, |
| 348 | testutil.Turn{Text: "recovered without write"}, |
| 349 | ) |
| 350 | a := New(mp, reg, NewSession(""), Options{}, event.Discard) |
| 351 | if err := a.Run(withNoClosedLoop(context.Background()), "write it"); !errors.Is(err, interrupted) { |
| 352 | t.Fatalf("Run: %v", err) |
| 353 | } |
| 354 | if mp.CallCount() != 1 { |
| 355 | t.Fatalf("interrupted request retried: %d calls", mp.CallCount()) |
| 356 | } |
| 357 | if writer.calls.Load() != 0 { |
| 358 | t.Fatalf("writer executed %d times, want 0 (uncommitted tool call)", writer.calls.Load()) |
| 359 | } |
| 360 | } |
| 361 | |
| 362 | type countingWriterTool struct{ calls atomic.Int32 } |
| 363 | |
| 364 | func (c *countingWriterTool) Name() string { return "write_file" } |
| 365 | func (c *countingWriterTool) Description() string { return "count writes" } |
| 366 | func (c *countingWriterTool) Schema() json.RawMessage { |
| 367 | return json.RawMessage(`{"type":"object","properties":{"path":{"type":"string"},"content":{"type":"string"}}}`) |
| 368 | } |
| 369 | func (c *countingWriterTool) ReadOnly() bool { return false } |
| 370 | func (c *countingWriterTool) Execute(context.Context, json.RawMessage) (string, error) { |
| 371 | c.calls.Add(1) |
| 372 | return "wrote", nil |
| 373 | } |
| 374 | |
| 375 | func TestRunGenericStreamErrorPersistsLocalDisplayAndInjectsBoundedRecovery(t *testing.T) { |
| 376 | apiErr := errors.New("upstream reset") |
| 377 | mp := testutil.NewMock("m", |
| 378 | testutil.Turn{Reasoning: "private partial reasoning", Text: "visible partial", ChunkError: apiErr}, |
| 379 | testutil.Turn{Text: "continued safely"}, |
| 380 | ) |
| 381 | session := NewSession("system") |
| 382 | a := New(mp, echoRegistry(), session, Options{}, event.Discard) |
| 383 | |
| 384 | if err := a.Run(withNoClosedLoop(context.Background()), "change the file"); !errors.Is(err, apiErr) { |
| 385 | t.Fatalf("first Run error = %v, want %v", err, apiErr) |
| 386 | } |
| 387 | msgs := session.Snapshot() |
| 388 | last := msgs[len(msgs)-1] |
| 389 | if !last.LocalOnly || last.InterruptedTurn == nil || !last.InterruptedTurn.Pending { |
| 390 | t.Fatalf("terminal stream error did not leave pending local recovery: %+v", last) |
| 391 | } |
| 392 | if last.Content != "visible partial" || last.ReasoningContent != "private partial reasoning" { |
| 393 | t.Fatalf("local display lost streamed output: %+v", last) |
| 394 | } |
| 395 | |
| 396 | if err := a.Run(withNoClosedLoop(context.Background()), "continue"); err != nil { |
| 397 | t.Fatalf("second Run: %v", err) |
| 398 | } |
| 399 | req := mp.Requests()[1] |
| 400 | for _, message := range req.Messages { |
| 401 | if message.LocalOnly || strings.Contains(message.Content, "visible partial") || strings.Contains(message.ReasoningContent, "private partial reasoning") { |
| 402 | t.Fatalf("unsafe partial output leaked to provider: %+v", req.Messages) |
| 403 | } |
| 404 | } |
| 405 | lastUser := req.Messages[len(req.Messages)-1] |
| 406 | if lastUser.Role != provider.RoleUser || !strings.Contains(lastUser.Content, "<interrupted-turn-recovery>") || |
| 407 | !strings.Contains(lastUser.Content, "unsafe_partial_output: excluded") || !strings.Contains(lastUser.Content, "continue") { |
| 408 | t.Fatalf("next user turn missing bounded recovery block: %+v", lastUser) |
| 409 | } |
| 410 | if got := StripTransientUserBlocks(lastUser.Content); got != "continue" { |
| 411 | t.Fatalf("recovery block leaked into user display: %q", got) |
| 412 | } |
| 413 | } |
| 414 | |
| 415 | func TestRunRecoveryKeepsCompletedToolPairAndSummarizesChangedFile(t *testing.T) { |
| 416 | session := NewSession("system") |
| 417 | session.Add(provider.Message{Role: provider.RoleUser, Content: "update config"}) |
| 418 | session.Add(provider.Message{Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ |
| 419 | ID: "done-1", Name: "write_file", Arguments: `{"path":"config.json","content":"{}"}`, Added: 1, |
| 420 | }}}) |
| 421 | session.Add(provider.Message{Role: provider.RoleTool, ToolCallID: "done-1", Name: "write_file", Content: "wrote config.json"}) |
| 422 | session.Add(provider.Message{ |
| 423 | Role: provider.RoleTool, ToolCallID: provider.LocalOnlyToolID, Name: provider.LocalOnlyToolName, LocalOnly: true, |
| 424 | ReasoningContent: "unsafe partial reasoning", |
| 425 | InterruptedTurn: &provider.InterruptedTurnRecovery{ |
| 426 | Pending: true, |
| 427 | CompletedTools: []provider.InterruptedToolSummary{{ |
| 428 | ID: "done-1", Name: "write_file", Files: []string{"config.json"}, Added: 1, |
| 429 | }}, |
| 430 | InterruptedTools: []string{"bash"}, |
| 431 | DroppedPartialReasoning: true, |
| 432 | }, |
| 433 | }) |
| 434 | mp := testutil.NewMock("m", testutil.Turn{Text: "done"}) |
| 435 | a := New(mp, echoRegistry(), session, Options{}, event.Discard) |
| 436 | if err := a.Run(withNoClosedLoop(context.Background()), "continue"); err != nil { |
| 437 | t.Fatalf("Run: %v", err) |
| 438 | } |
| 439 | |
| 440 | req := mp.Requests()[0] |
| 441 | if len(req.Messages) != 5 { |
| 442 | t.Fatalf("provider request should contain system + user + complete pair + recovery user, got %+v", req.Messages) |
| 443 | } |
| 444 | if req.Messages[2].Role != provider.RoleAssistant || req.Messages[3].Role != provider.RoleTool { |
| 445 | t.Fatalf("completed tool pair was not replayed canonically: %+v", req.Messages) |
| 446 | } |
| 447 | last := req.Messages[len(req.Messages)-1] |
| 448 | for _, want := range []string{"write_file files=config.json diff=+1/-0", "interrupted_tools: bash", "Use these facts", "continue"} { |
| 449 | if !strings.Contains(last.Content, want) { |
| 450 | t.Fatalf("recovery user message missing %q: %s", want, last.Content) |
| 451 | } |
| 452 | } |
| 453 | if strings.Contains(last.Content, "unsafe partial reasoning") { |
| 454 | t.Fatalf("raw partial reasoning leaked into recovery summary: %s", last.Content) |
| 455 | } |
| 456 | } |
| 457 | |
| 458 | // TestRunWellFormedToolLoopRoundTrips is the happy-path baseline: a tool round |
| 459 | // then a final answer. The session must end with the assistant answer and pair |
| 460 | // cleanly (the repair is a no-op on well-formed histories). |
| 461 | func TestRunWellFormedToolLoopRoundTrips(t *testing.T) { |
| 462 | mp := testutil.NewMock("m", |
| 463 | testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "c1", Name: "echo", Arguments: `{"text":"hi"}`}}}, |
| 464 | testutil.Turn{Text: "all set"}, |
| 465 | ) |
| 466 | a := New(mp, echoRegistry(), NewSession(""), Options{}, event.Discard) |
| 467 | if err := a.Run(withNoClosedLoop(context.Background()), "go"); err != nil { |
| 468 | t.Fatalf("Run: %v", err) |
| 469 | } |
| 470 | |
| 471 | msgs := a.Session().Messages |
| 472 | last := msgs[len(msgs)-1] |
| 473 | if last.Role != provider.RoleAssistant || last.Content != "all set" { |
| 474 | t.Fatalf("final message should be the assistant answer, got %+v", last) |
| 475 | } |
| 476 | before := len(msgs) |
| 477 | if after := len(provider.SanitizeToolPairing(msgs)); after != before { |
| 478 | t.Errorf("repair mutated a well-formed session: %d -> %d", before, after) |
| 479 | } |
| 480 | } |
| 481 | |
| 482 | // TestRunNotifiesWhenStreamFails pins the #9560 visibility fix: |
| 483 | // when a model request ends in a stream interruption, |
| 484 | // the run must surface a user-readable warn notice explaining the failure — |
| 485 | // not only the generic interrupted-turn record. |
| 486 | func TestRunNotifiesWhenStreamFails(t *testing.T) { |
| 487 | interrupted := &provider.StreamInterruptedError{Err: errors.New("dial tcp: lookup gw.invalid: no such host"), Reason: provider.StreamInterruptIdleTimeout} |
| 488 | script := make([]testutil.Turn, maxSamplingAttempts) |
| 489 | for i := range script { |
| 490 | script[i] = testutil.Turn{ChunkError: interrupted} |
| 491 | } |
| 492 | mp := testutil.NewMock("m", script...) |
| 493 | sink := &recordSink{} |
| 494 | a := New(mp, echoRegistry(), NewSession(""), Options{}, sink) |
| 495 | |
| 496 | err := a.Run(withNoClosedLoop(context.Background()), "go") |
| 497 | if err == nil { |
| 498 | t.Fatal("Run must fail immediately on stream interruption") |
| 499 | } |
| 500 | if !provider.IsStreamInterrupted(err) { |
| 501 | t.Fatalf("terminal error = %v, want a stream interruption", err) |
| 502 | } |
| 503 | var sawExplanation bool |
| 504 | for _, e := range sink.kinds(event.Notice) { |
| 505 | if e.Level == event.LevelWarn && strings.Contains(e.Text, "idle timeout") { |
| 506 | sawExplanation = true |
| 507 | if e.Code != event.NoticeCodeStreamInterruptedIdleTimeout { |
| 508 | t.Fatalf("stream interruption notice code = %q", e.Code) |
| 509 | } |
| 510 | if strings.Contains(e.Text, "gw.invalid") || strings.Contains(e.Text, "dial tcp") { |
| 511 | t.Fatalf("notice leaks raw transport error text: %q", e.Text) |
| 512 | } |
| 513 | } |
| 514 | } |
| 515 | if !sawExplanation { |
| 516 | notices := sink.kinds(event.Notice) |
| 517 | t.Fatalf("no warn notice explains the exhausted stream; notices = %+v", notices) |
| 518 | } |
| 519 | } |
| 520 |