(t *testing.T)
| 2266 | } |
| 2267 | |
| 2268 | func TestToolErrorMiddleware(t *testing.T) { |
| 2269 | ctx := context.Background() |
| 2270 | ct, st := NewInMemoryTransports() |
| 2271 | |
| 2272 | s := NewServer(testImpl, nil) |
| 2273 | AddTool(s, &Tool{ |
| 2274 | Name: "greet", |
| 2275 | Description: "say hi", |
| 2276 | }, sayHi) |
| 2277 | AddTool(s, &Tool{Name: "fail", InputSchema: &jsonschema.Schema{Type: "object"}}, |
| 2278 | func(context.Context, *CallToolRequest, map[string]any) (*CallToolResult, any, error) { |
| 2279 | return nil, nil, errTestFailure |
| 2280 | }) |
| 2281 | |
| 2282 | var middleErr error |
| 2283 | s.AddReceivingMiddleware(func(h MethodHandler) MethodHandler { |
| 2284 | return func(ctx context.Context, method string, req Request) (Result, error) { |
| 2285 | res, err := h(ctx, method, req) |
| 2286 | if err == nil { |
| 2287 | if ctr, ok := res.(*CallToolResult); ok { |
| 2288 | middleErr = ctr.GetError() |
| 2289 | } |
| 2290 | } |
| 2291 | return res, err |
| 2292 | } |
| 2293 | }) |
| 2294 | _, err := s.Connect(ctx, st, nil) |
| 2295 | if err != nil { |
| 2296 | t.Fatal(err) |
| 2297 | } |
| 2298 | client := NewClient(&Implementation{Name: "test-client"}, nil) |
| 2299 | clientSession, err := client.Connect(ctx, ct, nil) |
| 2300 | if err != nil { |
| 2301 | t.Fatal(err) |
| 2302 | } |
| 2303 | defer clientSession.Close() |
| 2304 | |
| 2305 | _, err = clientSession.CallTool(ctx, &CallToolParams{ |
| 2306 | Name: "greet", |
| 2307 | Arguments: map[string]any{"Name": "al"}, |
| 2308 | }) |
| 2309 | if err != nil { |
| 2310 | t.Errorf("CallTool() failed: %v", err) |
| 2311 | } |
| 2312 | if middleErr != nil { |
| 2313 | t.Errorf("middleware got error %v, want nil", middleErr) |
| 2314 | } |
| 2315 | res, err := clientSession.CallTool(ctx, &CallToolParams{ |
| 2316 | Name: "fail", |
| 2317 | }) |
| 2318 | if err != nil { |
| 2319 | t.Errorf("CallTool() failed: %v", err) |
| 2320 | } |
| 2321 | if !res.IsError { |
| 2322 | t.Fatal("want error, got none") |
| 2323 | } |
| 2324 | // Clients can't see the error, because it isn't marshaled. |
| 2325 | if err := res.GetError(); err != nil { |
nothing calls this directly
no test coverage detected
searching dependent graphs…