| 1 | package plugin |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "reflect" |
| 8 | "strings" |
| 9 | "sync/atomic" |
| 10 | "testing" |
| 11 | "time" |
| 12 | |
| 13 | mcpsdk "github.com/modelcontextprotocol/go-sdk/mcp" |
| 14 | ) |
| 15 | |
| 16 | func TestSDKListsConsumeEveryPageOnOneSession(t *testing.T) { |
| 17 | var connections atomic.Int32 |
| 18 | transport := newInMemorySDKTransport(t, func() *mcpsdk.Server { |
| 19 | connections.Add(1) |
| 20 | server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "paged", Version: "1"}, &mcpsdk.ServerOptions{PageSize: 1}) |
| 21 | for _, name := range []string{"zeta", "alpha", "middle"} { |
| 22 | server.AddTool(&mcpsdk.Tool{Name: name, InputSchema: map[string]any{"type": "object"}}, nil) |
| 23 | server.AddPrompt(&mcpsdk.Prompt{Name: "prompt_" + name}, nil) |
| 24 | server.AddResource(&mcpsdk.Resource{URI: "test://" + name, Name: name}, nil) |
| 25 | } |
| 26 | return server |
| 27 | }) |
| 28 | |
| 29 | assertListSize := func(method, key string, want int) { |
| 30 | t.Helper() |
| 31 | result, err := transport.call(t.Context(), method, map[string]any{}) |
| 32 | if err != nil { |
| 33 | t.Fatalf("%s: %v", method, err) |
| 34 | } |
| 35 | var payload map[string]json.RawMessage |
| 36 | if err := json.Unmarshal(result, &payload); err != nil { |
| 37 | t.Fatal(err) |
| 38 | } |
| 39 | var values []json.RawMessage |
| 40 | if err := json.Unmarshal(payload[key], &values); err != nil { |
| 41 | t.Fatal(err) |
| 42 | } |
| 43 | if len(values) != want { |
| 44 | t.Fatalf("%s returned %d items, want %d", method, len(values), want) |
| 45 | } |
| 46 | } |
| 47 | assertListSize("tools/list", "tools", 3) |
| 48 | assertListSize("prompts/list", "prompts", 3) |
| 49 | assertListSize("resources/list", "resources", 3) |
| 50 | if got := connections.Load(); got != 1 { |
| 51 | t.Fatalf("tool/prompt/resource lists opened %d sessions, want one", got) |
| 52 | } |
| 53 | } |
| 54 | |
| 55 | func TestSDKPromptAndResourceListChangesRefreshSharedSession(t *testing.T) { |
| 56 | var connections atomic.Int32 |
| 57 | server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "surfaces", Version: "1"}, nil) |
| 58 | server.AddPrompt(&mcpsdk.Prompt{Name: "prompt-one"}, nil) |
| 59 | server.AddResource(&mcpsdk.Resource{URI: "test://resource-one", Name: "resource-one"}, nil) |
| 60 | transport := newInMemorySDKTransport(t, func() *mcpsdk.Server { |
| 61 | connections.Add(1) |
| 62 | return server |
| 63 | }) |
| 64 | refreshCtx, cancelRefresh := context.WithCancel(t.Context()) |
| 65 | client := &Client{ |
| 66 | name: "surfaces", spec: Spec{Name: "surfaces", Type: "http"}, transport: "http", t: transport, |
| 67 | refresh: toolListRefreshState{ctx: refreshCtx, cancel: cancelRefresh}, |
| 68 | } |
| 69 | if err := client.initialize(t.Context()); err != nil { |
| 70 | t.Fatal(err) |
| 71 | } |
| 72 | if !client.capabilities.promptsListChanged || !client.capabilities.resourcesListChanged { |
| 73 | t.Fatalf("list-changed capabilities = prompts:%v resources:%v", client.capabilities.promptsListChanged, client.capabilities.resourcesListChanged) |
| 74 | } |
| 75 | host := NewHost() |
| 76 | host.bindToolListChanges(client) |
| 77 | if _, err := host.registerStartedClient(client, nil); err != nil { |
| 78 | t.Fatal(err) |
| 79 | } |
| 80 | t.Cleanup(host.Close) |
| 81 | host.StartPhaseB(t.Context(), nil) |
| 82 | |
| 83 | waitSurfaceCounts := func(wantPrompts, wantResources int) { |
| 84 | t.Helper() |
| 85 | deadline := time.Now().Add(2 * time.Second) |
| 86 | for { |
| 87 | host.mu.RLock() |
| 88 | gotPrompts, gotResources := len(host.prompts), len(host.resources) |
| 89 | host.mu.RUnlock() |
| 90 | if gotPrompts == wantPrompts && gotResources == wantResources { |
| 91 | return |
| 92 | } |
| 93 | if time.Now().After(deadline) { |
| 94 | t.Fatalf("surface counts = %d/%d, want %d/%d", gotPrompts, gotResources, wantPrompts, wantResources) |
| 95 | } |
| 96 | time.Sleep(5 * time.Millisecond) |
| 97 | } |
| 98 | } |
| 99 | waitSurfaceCounts(1, 1) |
| 100 | |
| 101 | server.AddPrompt(&mcpsdk.Prompt{Name: "prompt-two"}, nil) |
| 102 | server.AddResource(&mcpsdk.Resource{URI: "test://resource-two", Name: "resource-two"}, nil) |
| 103 | waitSurfaceCounts(2, 2) |
| 104 | if got := connections.Load(); got != 1 { |
| 105 | t.Fatalf("prompt/resource refresh opened %d sessions, want one shared session", got) |
| 106 | } |
| 107 | } |
| 108 | |
| 109 | func TestSDKToolConversionPreservesProviderCatalogFingerprint(t *testing.T) { |
| 110 | const wireFixture = `{"tools":[ |
| 111 | {"name":"zeta","description":"Z","inputSchema":{"required":["b","a"],"properties":{"b":{"type":"number"},"a":{"type":"string"}},"type":"object"},"outputSchema":{"type":"object","properties":{"ok":{"type":"boolean"}}},"annotations":{"readOnlyHint":true,"destructiveHint":false}}, |
| 112 | {"name":"alpha","description":"A","inputSchema":{"type":"object","properties":{}}} |
| 113 | ]}` |
| 114 | wireClient := &Client{name: "fixture", t: &countingToolsTransport{raw: json.RawMessage(wireFixture)}} |
| 115 | wireCatalog, err := wireClient.fetchToolCatalog(t.Context(), false) |
| 116 | if err != nil { |
| 117 | t.Fatal(err) |
| 118 | } |
| 119 | |
| 120 | transport := newInMemorySDKTransport(t, func() *mcpsdk.Server { |
| 121 | server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "fixture", Version: "1"}, nil) |
| 122 | falseValue := false |
| 123 | server.AddTool(&mcpsdk.Tool{ |
| 124 | Name: "zeta", Description: "Z", |
| 125 | InputSchema: map[string]any{ |
| 126 | "required": []any{"b", "a"}, "properties": map[string]any{ |
| 127 | "b": map[string]any{"type": "number"}, "a": map[string]any{"type": "string"}, |
| 128 | }, "type": "object", |
| 129 | }, |
| 130 | OutputSchema: map[string]any{"type": "object", "properties": map[string]any{"ok": map[string]any{"type": "boolean"}}}, |
| 131 | Annotations: &mcpsdk.ToolAnnotations{ReadOnlyHint: true, DestructiveHint: &falseValue}, |
| 132 | }, nil) |
| 133 | server.AddTool(&mcpsdk.Tool{Name: "alpha", Description: "A", InputSchema: map[string]any{"type": "object", "properties": map[string]any{}}}, nil) |
| 134 | return server |
| 135 | }) |
| 136 | sdkClient := &Client{name: "fixture", t: transport} |
| 137 | sdkCatalog, err := sdkClient.fetchToolCatalog(t.Context(), false) |
| 138 | if err != nil { |
| 139 | t.Fatal(err) |
| 140 | } |
| 141 | |
| 142 | if wireCatalog.fingerprint != sdkCatalog.fingerprint { |
| 143 | t.Fatalf("catalog fingerprint changed across SDK conversion:\nwire=%x\nsdk =%x", wireCatalog.fingerprint, sdkCatalog.fingerprint) |
| 144 | } |
| 145 | if !reflect.DeepEqual(wireCatalog.infos, sdkCatalog.infos) { |
| 146 | t.Fatalf("tool info changed:\nwire=%+v\nsdk =%+v", wireCatalog.infos, sdkCatalog.infos) |
| 147 | } |
| 148 | if got, want := toolCatalogBytes(sdkCatalog), toolCatalogBytes(wireCatalog); !reflect.DeepEqual(got, want) { |
| 149 | t.Fatalf("provider-visible tool catalog changed:\nwire=%s\nsdk =%s", want, got) |
| 150 | } |
| 151 | } |
| 152 | |
| 153 | func toolCatalogBytes(catalog toolCatalogSnapshot) [][]byte { |
| 154 | out := make([][]byte, 0, len(catalog.adapters)) |
| 155 | for _, adapter := range catalog.adapters { |
| 156 | out = append(out, []byte(adapter.Name()+"\x00"+adapter.Description()+"\x00"+string(adapter.Schema()))) |
| 157 | } |
| 158 | return out |
| 159 | } |
| 160 | |
| 161 | func TestSDKSessionWaiterCancellationDoesNotCancelSharedBuild(t *testing.T) { |
| 162 | lifeCtx, cancelLife := context.WithCancel(context.Background()) |
| 163 | transport := &sdkSessionTransport{ |
| 164 | name: "waiter", spec: Spec{Name: "waiter", Type: "http"}, lifeCtx: lifeCtx, cancel: cancelLife, |
| 165 | state: SessionStateConnecting, |
| 166 | } |
| 167 | release := make(chan struct{}) |
| 168 | transport.endpointFactory = func(ctx context.Context) (sdkEndpoint, error) { |
| 169 | <-release |
| 170 | clientSide, serverSide := mcpsdk.NewInMemoryTransports() |
| 171 | server := mcpsdk.NewServer(&mcpsdk.Implementation{Name: "waiter", Version: "1"}, nil) |
| 172 | go func() { _ = server.Run(ctx, serverSide) }() |
| 173 | return sdkEndpoint{transport: clientSide}, nil |
| 174 | } |
| 175 | t.Cleanup(transport.close) |
| 176 | |
| 177 | waitCtx, cancelWait := context.WithCancel(t.Context()) |
| 178 | done := make(chan error, 1) |
| 179 | go func() { |
| 180 | _, err := transport.acquire(waitCtx) |
| 181 | done <- err |
| 182 | }() |
| 183 | cancelWait() |
| 184 | if err := <-done; !errors.Is(err, context.Canceled) { |
| 185 | t.Fatalf("cancelled waiter error = %v, want context.Canceled", err) |
| 186 | } |
| 187 | close(release) |
| 188 | if _, err := transport.acquire(t.Context()); err != nil { |
| 189 | t.Fatalf("shared build was cancelled with its first waiter: %v", err) |
| 190 | } |
| 191 | } |
| 192 | |
| 193 | func TestSDKSessionDiagnosticsRedactSessionAndConfiguredValues(t *testing.T) { |
| 194 | const ( |
| 195 | sessionID = "session-secret-123" |
| 196 | projectPath = "/workspace/private-project" |
| 197 | headerToken = "header-secret-456" |
| 198 | envToken = "environment-secret-789" |
| 199 | ) |
| 200 | transport := &sdkSessionTransport{spec: Spec{ |
| 201 | WorkspaceRoot: projectPath, |
| 202 | Headers: map[string]string{"IJ_MCP_SERVER_PROJECT_PATH": projectPath, "Authorization": headerToken}, |
| 203 | Env: map[string]string{"MCP_TOKEN": envToken}, |
| 204 | }} |
| 205 | message := transport.safeErrorText(errors.New( |
| 206 | "request failed (session ID: "+sessionID+"): project="+projectPath+" auth="+headerToken+" env="+envToken, |
| 207 | ), sessionID) |
| 208 | for _, secret := range []string{sessionID, projectPath, headerToken, envToken} { |
| 209 | if strings.Contains(message, secret) { |
| 210 | t.Fatalf("diagnostic leaked %q: %s", secret, message) |
| 211 | } |
| 212 | } |
| 213 | } |
| 214 |