返回 DeepSeek-Reasonix
prompt_result_test.go
根目录 / internal / acp / prompt_result_test.go
1 package acp
2
3 import (
4 "context"
5 "encoding/json"
6 "errors"
7 "strings"
8 "sync/atomic"
9 "testing"
10
11 "reasonix/internal/agent"
12 "reasonix/internal/control"
13 "reasonix/internal/event"
14 )
15
16 func TestServePromptFailureReturnsRedactedJSONRPCError(t *testing.T) {
17 const secret = "ghp_abcdefghijklmnopqrstuvwxyz"
18 const opaqueSecret = "relayKeyAbcdefghijkl"
19 const maskedSuffix = "ae54"
20 var attempts atomic.Int32
21 factory := &fakeFactory{behavior: func(_ context.Context, sink event.Sink, _ string) error {
22 if attempts.Add(1) == 1 {
23 return errors.New("provider failed: Authorization: Bearer " + secret + " credential " + opaqueSecret +
24 " rejected token ****" + maskedSuffix + "\ndetails=" + strings.Repeat("x", 3_000))
25 }
26 sink.Emit(event.Event{Kind: event.Text, Text: "recovered"})
27 return nil
28 }}
29 client, stop := startServer(t, factory)
30 defer stop()
31
32 client.call(t, "initialize", InitializeParams{ProtocolVersion: 1})
33 newResp := client.call(t, "session/new", SessionNewParams{})
34 var nr SessionNewResult
35 if err := json.Unmarshal(newResp.Result, &nr); err != nil {
36 t.Fatalf("session/new result: %v", err)
37 }
38
39 promptCh := client.callAsync("session/prompt", SessionPromptParams{
40 SessionID: nr.SessionID,
41 Prompt: []ContentBlock{{Type: "text", Text: "fail"}},
42 })
43 notifications, resp := drainPrompt(t, client, promptCh)
44 if resp.Error == nil || resp.Error.Code != ErrInternal {
45 t.Fatalf("prompt response = %+v, want %d JSON-RPC error", resp, ErrInternal)
46 }
47 if !strings.HasPrefix(resp.Error.Message, "session/prompt: provider failed:") {
48 t.Errorf("error message = %q, want underlying cause", resp.Error.Message)
49 }
50 if strings.Contains(resp.Error.Message, secret) || strings.Contains(resp.Error.Message, opaqueSecret) || strings.Contains(resp.Error.Message, maskedSuffix) {
51 t.Errorf("error message leaked credential: %q", resp.Error.Message)
52 }
53 if len(resp.Error.Message) > len("session/prompt: ")+2_048 {
54 t.Errorf("error message length = %d, want at most %d", len(resp.Error.Message), len("session/prompt: ")+2_048)
55 }
56 if len(resp.Result) != 0 {
57 t.Errorf("result = %s, want no successful prompt result", resp.Result)
58 }
59
60 wantReason := strings.TrimPrefix(resp.Error.Message, "session/prompt: ")
61 foundStatus := false
62 for _, notification := range notifications {
63 if notification.Method != sessionStatusUpdateMethod {
64 continue
65 }
66 var update ReasonixStatusUpdate
67 if err := json.Unmarshal(notification.Params, &update); err != nil {
68 t.Fatalf("status update: %v", err)
69 }
70 if update.Event == "error" {
71 foundStatus = true
72 if update.Status.TurnOutcome.Kind != "error" || update.Status.TurnOutcome.Reason != wantReason {
73 t.Errorf("error status = %+v, want reason %q", update.Status.TurnOutcome, wantReason)
74 }
75 }
76 }
77 if !foundStatus {
78 t.Fatal("missing error status update before prompt response")
79 }
80
81 retryCh := client.callAsync("session/prompt", SessionPromptParams{
82 SessionID: nr.SessionID,
83 Prompt: []ContentBlock{{Type: "text", Text: "retry"}},
84 })
85 _, retryResp := drainPrompt(t, client, retryCh)
86 if retryResp.Error != nil {
87 t.Fatalf("retry prompt errored: %+v", retryResp.Error)
88 }
89 var retryResult SessionPromptResult
90 if err := json.Unmarshal(retryResp.Result, &retryResult); err != nil {
91 t.Fatalf("retry prompt result: %v", err)
92 }
93 if retryResult.StopReason != StopEndTurn {
94 t.Errorf("retry stopReason = %q, want end_turn", retryResult.StopReason)
95 }
96 }
97
98 func TestServeCancelWhenRunnerReturnsNil(t *testing.T) {
99 started := make(chan struct{})
100 factory := &fakeFactory{behavior: func(ctx context.Context, _ event.Sink, _ string) error {
101 close(started)
102 <-ctx.Done()
103 return nil
104 }}
105 client, stop := startServer(t, factory)
106 defer stop()
107
108 client.call(t, "initialize", InitializeParams{ProtocolVersion: 1})
109 newResp := client.call(t, "session/new", SessionNewParams{})
110 var nr SessionNewResult
111 if err := json.Unmarshal(newResp.Result, &nr); err != nil {
112 t.Fatalf("session/new result: %v", err)
113 }
114 promptCh := client.callAsync("session/prompt", SessionPromptParams{
115 SessionID: nr.SessionID,
116 Prompt: []ContentBlock{{Type: "text", Text: "loop"}},
117 })
118 <-started
119 client.notify("session/cancel", SessionCancelParams{SessionID: nr.SessionID})
120
121 _, resp := drainPrompt(t, client, promptCh)
122 if resp.Error != nil {
123 t.Fatalf("cancelled prompt errored: %+v", resp.Error)
124 }
125 var result SessionPromptResult
126 if err := json.Unmarshal(resp.Result, &result); err != nil {
127 t.Fatalf("prompt result: %v", err)
128 }
129 if result.StopReason != StopCancelled {
130 t.Errorf("stopReason = %q, want cancelled", result.StopReason)
131 }
132 }
133
134 func TestServePromptRecoveryPauseReturnsEndTurn(t *testing.T) {
135 factory := &fakeFactory{behavior: func(context.Context, event.Sink, string) error {
136 return &agent.RecoveryPauseError{Message: "automatic recovery paused"}
137 }}
138 client, stop := startServer(t, factory)
139 defer stop()
140
141 client.call(t, "initialize", InitializeParams{ProtocolVersion: 1})
142 newResp := client.call(t, "session/new", SessionNewParams{})
143 var nr SessionNewResult
144 if err := json.Unmarshal(newResp.Result, &nr); err != nil {
145 t.Fatalf("session/new result: %v", err)
146 }
147 promptCh := client.callAsync("session/prompt", SessionPromptParams{
148 SessionID: nr.SessionID,
149 Prompt: []ContentBlock{{Type: "text", Text: "pause"}},
150 })
151 _, resp := drainPrompt(t, client, promptCh)
152 if resp.Error != nil {
153 t.Fatalf("prompt error = %+v, want controlled completion", resp.Error)
154 }
155 var result SessionPromptResult
156 if err := json.Unmarshal(resp.Result, &result); err != nil {
157 t.Fatalf("prompt result: %v", err)
158 }
159 if result.StopReason != StopEndTurn {
160 t.Errorf("stopReason = %q, want end_turn", result.StopReason)
161 }
162 }
163
164 func TestServeStaleFinalReadinessRecoveryIsInvalidWithoutStatusTurn(t *testing.T) {
165 factory := &fakeFactory{behavior: func(context.Context, event.Sink, string) error { return nil }}
166 client, stop := startServer(t, factory)
167 defer stop()
168
169 client.call(t, "initialize", InitializeParams{ProtocolVersion: 1})
170 newResp := client.call(t, "session/new", SessionNewParams{})
171 var nr SessionNewResult
172 if err := json.Unmarshal(newResp.Result, &nr); err != nil {
173 t.Fatalf("session/new result: %v", err)
174 }
175 before := getStatus(t, client, nr.SessionID)
176
177 promptCh := client.callAsync("session/prompt", SessionPromptParams{
178 SessionID: nr.SessionID,
179 Action: control.FinalReadinessRecoveryAction,
180 Prompt: []ContentBlock{{Type: "text", Text: "continue checks"}},
181 })
182 notifications, resp := drainPrompt(t, client, promptCh)
183 if resp.Error == nil || resp.Error.Code != ErrInvalidRequest {
184 t.Fatalf("stale recovery response = %+v, want %d", resp, ErrInvalidRequest)
185 }
186 for _, notification := range notifications {
187 if notification.Method == sessionStatusUpdateMethod {
188 t.Fatalf("stale recovery published a status turn: %+v", notification)
189 }
190 }
191 after := getStatus(t, client, nr.SessionID)
192 if after.Sequence != before.Sequence || after.Phase != before.Phase || after.TurnOutcome != before.TurnOutcome {
193 t.Fatalf("stale recovery changed status: before=%+v after=%+v", before, after)
194 }
195
196 retryCh := client.callAsync("session/prompt", SessionPromptParams{
197 SessionID: nr.SessionID,
198 Prompt: []ContentBlock{{Type: "text", Text: "ordinary turn"}},
199 })
200 _, retryResp := drainPrompt(t, client, retryCh)
201 if retryResp.Error != nil {
202 t.Fatalf("ordinary prompt after stale recovery errored: %+v", retryResp.Error)
203 }
204 }
205
205 lines GO