| 1 | package agent |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "strings" |
| 9 | "testing" |
| 10 | "testing/synctest" |
| 11 | "time" |
| 12 | |
| 13 | "reasonix/internal/attachment" |
| 14 | "reasonix/internal/event" |
| 15 | "reasonix/internal/extension" |
| 16 | "reasonix/internal/extension/protocol" |
| 17 | "reasonix/internal/i18n" |
| 18 | "reasonix/internal/provider" |
| 19 | "reasonix/internal/tool" |
| 20 | ) |
| 21 | |
| 22 | func TestCompactionFinishPreservesUnrelatedErrors(t *testing.T) { |
| 23 | synctest.Test(t, func(t *testing.T) { |
| 24 | a := &Agent{} |
| 25 | ctx, finish := a.beginCompactionRun(t.Context()) |
| 26 | unrelated := errors.New("request preparation failed") |
| 27 | wrapped := fmt.Errorf("save failed: %w", &compactionPersistenceError{context.DeadlineExceeded}) |
| 28 | time.Sleep(compactionBudget) |
| 29 | synctest.Wait() |
| 30 | for _, original := range []error{nil, unrelated, wrapped, context.Canceled, fmt.Errorf("recovery guidance: %w", errSummaryBudget)} { |
| 31 | if got := finish(original); got != original { //nolint:errorlint // Assert exact identity; even an additional wrapper violates this boundary. |
| 32 | t.Errorf("finish(%v) = %v; original error identity lost", original, got) |
| 33 | } |
| 34 | } |
| 35 | if got := compactionError(ctx, context.DeadlineExceeded); !errors.Is(got, errSummaryBudget) { |
| 36 | t.Fatalf("work deadline = %v", got) |
| 37 | } |
| 38 | blocked := fmt.Errorf("%w: %w", ErrCompactionRequired, context.DeadlineExceeded) |
| 39 | if got := compactionError(ctx, blocked); !errors.Is(got, ErrCompactionRequired) || !errors.Is(got, errSummaryBudget) { |
| 40 | t.Fatalf("blocked deadline lost its causes: %v", got) |
| 41 | } |
| 42 | }) |
| 43 | ctx, cancel := context.WithCancel(t.Context()) |
| 44 | cancel() |
| 45 | original := errors.New("unrelated error after stop") |
| 46 | if got := compactionError(ctx, original); got != original { //nolint:errorlint // Assert exact identity, not merely an equivalent wrapped cause. |
| 47 | t.Fatalf("cancellation replaced unrelated error: %v", got) |
| 48 | } |
| 49 | if got := compactionError(t.Context(), context.DeadlineExceeded); got != context.DeadlineExceeded { //nolint:errorlint // A foreign deadline must be returned unchanged. |
| 50 | t.Fatalf("foreign deadline reclassified: %v", got) |
| 51 | } |
| 52 | } |
| 53 | |
| 54 | func TestRequestAndCompressValidationAreNotSummaryFailures(t *testing.T) { |
| 55 | p := &mockProvider{name: "fixture"} |
| 56 | s := NewSession("system") |
| 57 | s.Add(provider.Message{Role: provider.RoleUser, ImageInputs: []attachment.ImageInput{{Kind: attachment.KindURL, URL: "https://example.invalid/image.png"}}}) |
| 58 | a := New(p, nil, s, Options{}, event.Discard) |
| 59 | _, err := a.prepareSamplingRequest(t.Context()) |
| 60 | if err == nil || err.Error() != "image request resolver is unavailable" { |
| 61 | t.Fatalf("image error = %v", err) |
| 62 | } |
| 63 | _, err = a.CompressContext(t.Context(), tool.CompressRequest{Direction: "invalid"}) |
| 64 | if err == nil || err.Error() != "compress: direction must be before or after" { |
| 65 | t.Fatalf("compress validation = %v", err) |
| 66 | } |
| 67 | if len(p.requests) != 0 { |
| 68 | t.Fatal("validation invoked the provider") |
| 69 | } |
| 70 | } |
| 71 | |
| 72 | func TestSafeCompactionTimeoutPreservesLaterExtensionBlock(t *testing.T) { |
| 73 | for _, point := range []extension.InterceptorPoint{extension.PointContextPrepare, extension.PointProviderRequest} { |
| 74 | t.Run(string(point), func(t *testing.T) { |
| 75 | synctest.Test(t, func(t *testing.T) { |
| 76 | p := &slowSummaryProvider{} |
| 77 | client := &fakeDispatchClient{interceptFn: func(protocol.InterceptEvent, json.RawMessage) (protocol.InterceptResult, error) { |
| 78 | return blockWith("policy refused"), nil |
| 79 | }} |
| 80 | a := New(p, nil, foldableSessionOverForce(6), Options{ContextWindow: 5000, CompactRatio: .5, |
| 81 | Extensions: newExtDispatcher(client, false, nil, point)}, event.Discard) |
| 82 | started := time.Now() |
| 83 | _, err := a.prepareSamplingRequest(t.Context()) |
| 84 | var summary *SummaryError |
| 85 | want := extensionBlockedError(point, "policy refused").Error() |
| 86 | if err == nil || err.Error() != want || errors.As(err, &summary) || errors.Is(err, context.DeadlineExceeded) { |
| 87 | t.Fatalf("extension error after safe timeout = %v", err) |
| 88 | } |
| 89 | if p.calls != 1 || time.Since(started) != compactionBudget || a.currentProjectionVersion() != 0 { |
| 90 | t.Fatalf("calls=%d elapsed=%s projection=%d", p.calls, time.Since(started), a.currentProjectionVersion()) |
| 91 | } |
| 92 | }) |
| 93 | }) |
| 94 | } |
| 95 | } |
| 96 | |
| 97 | func TestPendingContextPersistenceIsNotSummaryFailure(t *testing.T) { |
| 98 | saveErr := errors.New("retry flush failure") |
| 99 | recorder := &modelContextRecorderStub{err: saveErr} |
| 100 | s := NewSession("system") |
| 101 | a := New(&mockProvider{name: "fixture"}, nil, s, Options{SessionCheckpointer: recorder}, event.Discard) |
| 102 | a.sess.pendingModelContextCommit = &SessionModelContextCommit{OperationID: "pending-commit"} |
| 103 | _, err := a.prepareSamplingRequest(t.Context()) |
| 104 | var summary *SummaryError |
| 105 | if !errors.Is(err, saveErr) || errors.As(err, &summary) { |
| 106 | t.Fatalf("pending persistence error = %v", err) |
| 107 | } |
| 108 | if a.sess.pendingModelContextCommit == nil { |
| 109 | t.Fatal("failed persistence discarded the pending commit") |
| 110 | } |
| 111 | } |
| 112 | |
| 113 | func TestSummaryRequestClassifiesItsOwnFailures(t *testing.T) { |
| 114 | providerErr := errors.New("summary provider failed") |
| 115 | for _, tc := range []struct { |
| 116 | name string |
| 117 | prov *fakeProvider |
| 118 | code string |
| 119 | cause error |
| 120 | }{ |
| 121 | {"provider", &fakeProvider{streamErr: providerErr}, "summary_provider_error", providerErr}, |
| 122 | {"empty", &fakeProvider{}, "summary_empty", errSummaryEmpty}, |
| 123 | } { |
| 124 | t.Run(tc.name, func(t *testing.T) { |
| 125 | a := New(tc.prov, nil, NewSession("system"), Options{}, event.Discard) |
| 126 | _, _, err := a.summarize(t.Context(), []provider.Message{{Role: provider.RoleUser, Content: "summarize this"}}, "") |
| 127 | var summary *SummaryError |
| 128 | if !errors.As(err, &summary) || summary.Code != tc.code || !errors.Is(err, tc.cause) { |
| 129 | t.Fatalf("summary error = %v, want %s wrapping %v", err, tc.code, tc.cause) |
| 130 | } |
| 131 | }) |
| 132 | } |
| 133 | } |
| 134 | |
| 135 | func TestCompactionTimeoutKeepsOutermostRecoveryGuidance(t *testing.T) { |
| 136 | synctest.Test(t, func(t *testing.T) { |
| 137 | p := &slowSummaryProvider{} |
| 138 | a := New(p, nil, foldableSessionOverForce(12), Options{ContextWindow: 5000, CompactRatio: .5}, event.Discard) |
| 139 | _, err := a.prepareSamplingRequest(t.Context()) |
| 140 | if !errors.Is(err, ErrCompactionRequired) || !errors.Is(err, errSummaryBudget) || !strings.HasPrefix(err.Error(), i18n.M.ContextLimitRecovery+": ") { |
| 141 | t.Fatalf("hard-limit timeout lost recovery guidance: %v", err) |
| 142 | } |
| 143 | if p.calls != 1 { |
| 144 | t.Fatalf("summary calls=%d", p.calls) |
| 145 | } |
| 146 | }) |
| 147 | } |
| 148 |