| 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 |