返回 DeepSeek-Reasonix
stdio_cancel_test.go
根目录 / internal / plugin / stdio_cancel_test.go
1 package plugin
2
3 import (
4 "context"
5 "encoding/json"
6 "errors"
7 "sync/atomic"
8 "testing"
9 "time"
10
11 mcpsdk "github.com/modelcontextprotocol/go-sdk/mcp"
12 "reasonix/internal/tool"
13 )
14
15 func newInMemorySDKTransport(t *testing.T, serverFactory func() *mcpsdk.Server) *sdkSessionTransport {
16 t.Helper()
17 lifeCtx, cancel := context.WithCancel(context.Background())
18 transport := &sdkSessionTransport{
19 name: "test",
20 spec: Spec{Name: "test", Type: "http", StartupTimeout: 2 * time.Second},
21 lifeCtx: lifeCtx,
22 cancel: cancel,
23 state: SessionStateConnecting,
24 reconnectDelays: []time.Duration{time.Millisecond},
25 }
26 transport.endpointFactory = func(ctx context.Context) (sdkEndpoint, error) {
27 clientSide, serverSide := mcpsdk.NewInMemoryTransports()
28 server := serverFactory()
29 go func() { _ = server.Run(ctx, serverSide) }()
30 return sdkEndpoint{transport: clientSide}, nil
31 }
32 t.Cleanup(transport.close)
33 return transport
34 }
35
36 func TestSDKIOCallReturnsOnContextCancelAndNotifiesServer(t *testing.T) {
37 started := make(chan struct{})
38 cancelled := make(chan struct{})
39 transport := newInMemorySDKTransport(t, func() *mcpsdk.Server {
40 server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "hung", Version: "1"}, nil)
41 mcpsdk.AddTool(server, &mcpsdk.Tool{Name: "wait"}, func(ctx context.Context, _ *mcpsdk.CallToolRequest, _ map[string]any) (*mcpsdk.CallToolResult, any, error) {
42 close(started)
43 <-ctx.Done()
44 close(cancelled)
45 return nil, nil, ctx.Err()
46 })
47 return server
48 })
49
50 ctx, cancel := context.WithCancel(context.Background())
51 done := make(chan error, 1)
52 go func() {
53 _, err := transport.call(ctx, "tools/call", map[string]any{"name": "wait", "arguments": map[string]any{}})
54 done <- err
55 }()
56 select {
57 case <-started:
58 case <-time.After(2 * time.Second):
59 t.Fatal("server did not receive tools/call")
60 }
61 cancel()
62 select {
63 case err := <-done:
64 if !errors.Is(err, context.Canceled) {
65 t.Fatalf("cancelled call error = %v, want context.Canceled", err)
66 }
67 case <-time.After(2 * time.Second):
68 t.Fatal("SDK call did not return after context cancellation")
69 }
70 select {
71 case <-cancelled:
72 case <-time.After(2 * time.Second):
73 t.Fatal("server did not receive notifications/cancelled")
74 }
75 }
76
77 func TestSDKSessionRoutesConcurrentResponsesByRequestID(t *testing.T) {
78 slowStarted := make(chan struct{})
79 releaseSlow := make(chan struct{})
80 transport := newInMemorySDKTransport(t, func() *mcpsdk.Server {
81 server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "parallel", Version: "1"}, nil)
82 mcpsdk.AddTool(server, &mcpsdk.Tool{Name: "work"}, func(_ context.Context, _ *mcpsdk.CallToolRequest, input map[string]any) (*mcpsdk.CallToolResult, any, error) {
83 label, _ := input["label"].(string)
84 if label == "slow" {
85 close(slowStarted)
86 <-releaseSlow
87 }
88 return &mcpsdk.CallToolResult{Content: []mcpsdk.Content{&mcpsdk.TextContent{Text: label}}}, nil, nil
89 })
90 return server
91 })
92
93 call := func(label string) (json.RawMessage, error) {
94 return transport.call(t.Context(), "tools/call", map[string]any{"name": "work", "arguments": map[string]any{"label": label}})
95 }
96 slowDone := make(chan json.RawMessage, 1)
97 go func() {
98 result, _ := call("slow")
99 slowDone <- result
100 }()
101 <-slowStarted
102 fastDone := make(chan json.RawMessage, 1)
103 go func() {
104 result, _ := call("fast")
105 fastDone <- result
106 }()
107 select {
108 case result := <-fastDone:
109 if !json.Valid(result) || !containsJSONText(result, "fast") {
110 t.Fatalf("fast result = %s", result)
111 }
112 case <-time.After(time.Second):
113 t.Fatal("fast request was serialized behind slow request")
114 }
115 close(releaseSlow)
116 select {
117 case result := <-slowDone:
118 if !containsJSONText(result, "slow") {
119 t.Fatalf("slow result = %s", result)
120 }
121 case <-time.After(time.Second):
122 t.Fatal("slow request did not finish")
123 }
124 }
125
126 func TestSDKSessionRoutesProgressNotification(t *testing.T) {
127 transport := newInMemorySDKTransport(t, func() *mcpsdk.Server {
128 server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "progress", Version: "1"}, nil)
129 mcpsdk.AddTool(server, &mcpsdk.Tool{Name: "index"}, func(ctx context.Context, req *mcpsdk.CallToolRequest, _ map[string]any) (*mcpsdk.CallToolResult, any, error) {
130 if token := req.Params.GetProgressToken(); token != nil {
131 _ = req.Session.NotifyProgress(ctx, &mcpsdk.ProgressNotificationParams{
132 ProgressToken: token, Progress: 2, Total: 5, Message: "Indexing",
133 })
134 }
135 return &mcpsdk.CallToolResult{}, nil, nil
136 })
137 return server
138 })
139 client := &Client{name: "progress", t: transport}
140 progress := make(chan string, 1)
141 ctx := tool.WithProgress(t.Context(), func(chunk string) { progress <- chunk })
142 if _, err := client.call(ctx, "tools/call", map[string]any{"name": "index", "arguments": map[string]any{}}); err != nil {
143 t.Fatalf("tools/call: %v", err)
144 }
145 select {
146 case got := <-progress:
147 if got != "Indexing (2/5)\n" {
148 t.Fatalf("progress = %q", got)
149 }
150 case <-time.After(time.Second):
151 t.Fatal("progress notification was not routed")
152 }
153 }
154
155 func TestSDKSessionConcurrentRebuildIsSingleflight(t *testing.T) {
156 var connections atomic.Int32
157 transport := newInMemorySDKTransport(t, func() *mcpsdk.Server {
158 connections.Add(1)
159 server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "singleflight", Version: "1"}, nil)
160 mcpsdk.AddTool(server, &mcpsdk.Tool{Name: "read"}, func(context.Context, *mcpsdk.CallToolRequest, map[string]any) (*mcpsdk.CallToolResult, any, error) {
161 return &mcpsdk.CallToolResult{}, nil, nil
162 })
163 return server
164 })
165 first, err := transport.acquire(t.Context())
166 if err != nil {
167 t.Fatal(err)
168 }
169 transport.invalidate(first)
170
171 const callers = 12
172 errs := make(chan error, callers)
173 for range callers {
174 go func() {
175 _, err := transport.acquire(t.Context())
176 errs <- err
177 }()
178 }
179 for range callers {
180 if err := <-errs; err != nil {
181 t.Fatal(err)
182 }
183 }
184 if got := connections.Load(); got != 2 {
185 t.Fatalf("connections = %d, want initial + one shared rebuild", got)
186 }
187 transport.mu.Lock()
188 current := transport.current
189 transport.mu.Unlock()
190 if current == nil || current == first {
191 t.Fatal("rebuild did not publish a new generation")
192 }
193
194 // A stale Wait callback can arrive after the replacement has already been
195 // published. Re-run that exact callback path and prove its generation fence
196 // cannot clear the healthy current session.
197 transport.handleSessionEnd(first, mcpsdk.ErrConnectionClosed)
198 transport.mu.Lock()
199 stillCurrent := transport.current == current
200 transport.mu.Unlock()
201 if !stillCurrent {
202 t.Fatal("stale generation callback cleared the replacement session")
203 }
204 }
205
206 func TestSDKSessionDropsStaleGenerationProgress(t *testing.T) {
207 transport := newInMemorySDKTransport(t, func() *mcpsdk.Server {
208 return mcpsdk.NewServer(&mcpsdk.Implementation{Name: "progress-generation", Version: "1"}, nil)
209 })
210 first, err := transport.acquire(t.Context())
211 if err != nil {
212 t.Fatal(err)
213 }
214 progress := make(chan string, 2)
215 stop := transport.registerProgress("same-token", func(chunk string) { progress <- chunk })
216 defer stop()
217
218 transport.dispatchSDKProgress(first.generation, &mcpsdk.ProgressNotificationParams{
219 ProgressToken: "same-token", Progress: 1, Total: 2, Message: "first",
220 })
221 if got := <-progress; got != "first (1/2)\n" {
222 t.Fatalf("first progress = %q", got)
223 }
224
225 transport.invalidate(first)
226 second, err := transport.acquire(t.Context())
227 if err != nil {
228 t.Fatal(err)
229 }
230 transport.dispatchSDKProgress(first.generation, &mcpsdk.ProgressNotificationParams{
231 ProgressToken: "same-token", Progress: 2, Total: 2, Message: "stale",
232 })
233 if len(progress) != 0 {
234 t.Fatalf("stale generation delivered progress: %q", <-progress)
235 }
236 transport.dispatchSDKProgress(second.generation, &mcpsdk.ProgressNotificationParams{
237 ProgressToken: "same-token", Progress: 2, Total: 2, Message: "current",
238 })
239 if got := <-progress; got != "current (2/2)\n" {
240 t.Fatalf("current progress = %q", got)
241 }
242 }
243
244 func containsJSONText(result json.RawMessage, want string) bool {
245 var decoded struct {
246 Content []struct {
247 Text string `json:"text"`
248 } `json:"content"`
249 }
250 return json.Unmarshal(result, &decoded) == nil && len(decoded.Content) == 1 && decoded.Content[0].Text == want
251 }
252
252 lines GO