返回 DeepSeek-Reasonix
request_recovery_test.go
根目录 / internal / agent / request_recovery_test.go
1 package agent
2
3 import (
4 "encoding/json"
5 "errors"
6 "reflect"
7 "strings"
8 "testing"
9
10 "reasonix/internal/event"
11 "reasonix/internal/extension"
12 "reasonix/internal/extension/dispatch"
13 "reasonix/internal/extension/protocol"
14 "reasonix/internal/provider"
15 "reasonix/internal/tool"
16 )
17
18 func TestSummaryRecoversOversizedProtectedToolResult(t *testing.T) {
19 p := &mockProvider{name: "fixture", chunks: []provider.Chunk{
20 {Type: provider.ChunkText, Text: "Read completed; keep the latest request and use the observation to continue."},
21 {Type: provider.ChunkDone},
22 }}
23 s := NewSession("sys")
24 s.Add(provider.Message{Role: provider.RoleUser, Content: "keep the latest request"})
25 s.Add(provider.Message{Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ID: "large", Name: "read", Arguments: `{}`}}})
26 s.Add(provider.Message{Role: provider.RoleTool, ToolCallID: "large", Name: "read", Content: strings.Repeat("large observation ", 10000)})
27 before := s.Snapshot()
28 // The recent tool result is too large for the desired replay, but fits the
29 // summarizer's input window. Recover by summarizing it, never truncating it.
30 a := New(p, tool.NewRegistry(), s, Options{ContextWindow: 64000}, event.Discard)
31 if err := a.CompactNow(t.Context(), ""); err != nil {
32 t.Fatal(err)
33 }
34 if len(p.requests) != 1 {
35 t.Fatalf("summary requests=%d, want 1", len(p.requests))
36 }
37 observed := false
38 for _, msg := range p.requests[0].Messages {
39 if msg.Role == provider.RoleTool && msg.ToolCallID == "large" && msg.Content == before[3].Content {
40 observed = true
41 }
42 }
43 if !observed || !reflect.DeepEqual(before, s.Snapshot()) {
44 t.Fatal("summary lost the protected tool result or changed canonical history")
45 }
46 prepared := a.contextManager().currentPrepared()
47 if prepared.ProjectionVersion == 0 || prepared.InputTokens >= 5000 || latestDigest(prepared.Messages) == "" {
48 t.Fatalf("summary did not recover a compact replay: version=%d tokens=%d", prepared.ProjectionVersion, prepared.InputTokens)
49 }
50 if err := provider.ValidateModelTranscript(provider.ModelMessages(prepared.Messages)); err != nil {
51 t.Fatalf("invalid recovered transcript: %v", err)
52 }
53 if truncatedRescue(a) {
54 t.Fatal("summary recovery used truncation")
55 }
56 }
57
58 func TestRecoveredHistoryReachesModelWithoutExecutingTools(t *testing.T) {
59 mp := &mockProvider{name: "fixture", chunks: []provider.Chunk{{Type: provider.ChunkDone}}}
60 s := NewSession("system")
61 s.Add(provider.Message{Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ID: "old", Name: "write_file", Arguments: `{"body":"["text"]"}`}}})
62 s.Add(provider.Message{Role: provider.RoleTool, ToolCallID: "old", Name: "write_file", ToolRunState: provider.ToolRunUnknown, Content: "original outcome uncertain"})
63 s.Add(provider.Message{Role: provider.RoleUser, Content: "continue"})
64 a := New(mp, tool.NewRegistry(), s, Options{}, event.Discard)
65 req, err := a.buildSamplingRequest(t.Context(), CompactionTriggerPressure)
66 if err != nil {
67 t.Fatal(err)
68 }
69 stream, err := a.streamProviderRequest(t.Context(), req.req)
70 if err != nil {
71 t.Fatal(err)
72 }
73 for range stream {
74 }
75 if len(mp.requests) != 1 || s.Snapshot()[1].ToolCalls[0].Arguments != `{"body":"["text"]"}` {
76 t.Fatal("request did not reach provider or canonical arguments changed")
77 }
78 }
79
80 func TestRequestExtensionRecoveryHonorsRequiredAndExplicitBlocks(t *testing.T) {
81 for _, point := range []extension.InterceptorPoint{extension.PointContextPrepare, extension.PointProviderRequest} {
82 for _, mode := range []string{"optional", "required", "block"} {
83 t.Run(string(point)+"/"+mode, func(t *testing.T) {
84 client := &fakeDispatchClient{interceptFn: func(ev protocol.InterceptEvent, raw json.RawMessage) (protocol.InterceptResult, error) {
85 if mode == "block" {
86 return blockWith("policy refused"), nil
87 }
88 var invalid []protocol.ProviderMessage
89 if err := json.Unmarshal([]byte(`[{"role":"assistant","tool_calls":[{"id":"a","name":"read","arguments":"[]"}]}]`), &invalid); err != nil {
90 t.Fatal(err)
91 }
92 if point == extension.PointContextPrepare {
93 return replaceWith(t, dispatch.ContextPayload{Messages: invalid}), nil
94 }
95 var payload dispatch.ProviderRequestPayload
96 if err := json.Unmarshal(raw, &payload); err != nil {
97 t.Fatal(err)
98 }
99 payload.Request.Messages = invalid
100 return replaceWith(t, payload), nil
101 }}
102 warnings := 0
103 d := newExtDispatcher(client, mode == "required", func(string) { warnings++ }, point)
104 a := New(&mockProvider{name: "fixture"}, tool.NewRegistry(), NewSession("system"), Options{Extensions: d}, event.Discard)
105 got, err := a.prepareSamplingRequest(t.Context())
106 if mode != "optional" {
107 var summary *SummaryError
108 if err == nil || errors.As(err, &summary) {
109 t.Fatalf("required extension or explicit block bypassed or reclassified: %v", err)
110 }
111 return
112 }
113 if err != nil || warnings != 1 || len(got.req.Messages) != 1 || got.req.Messages[0].Content != "system" {
114 t.Fatalf("optional invalid replacement did not preserve original request: %+v %v, warnings=%d", got, err, warnings)
115 }
116 })
117 }
118 }
119 }
120
120 lines GO