@pi-in-go/pigpen-a2a 0.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (52) hide show
  1. package/CREDITS.md +26 -0
  2. package/LICENSE +21 -0
  3. package/README.md +97 -0
  4. package/extensions/a2a/auth.go +120 -0
  5. package/extensions/a2a/auth_test.go +133 -0
  6. package/extensions/a2a/bench_test.go +70 -0
  7. package/extensions/a2a/binary_test.go +225 -0
  8. package/extensions/a2a/cancelqueued_test.go +142 -0
  9. package/extensions/a2a/card_cache_test.go +42 -0
  10. package/extensions/a2a/client.go +499 -0
  11. package/extensions/a2a/client_test.go +402 -0
  12. package/extensions/a2a/config.go +282 -0
  13. package/extensions/a2a/config_test.go +194 -0
  14. package/extensions/a2a/e2e_test.go +311 -0
  15. package/extensions/a2a/executor.go +318 -0
  16. package/extensions/a2a/extension.go +382 -0
  17. package/extensions/a2a/extension_test.go +396 -0
  18. package/extensions/a2a/fakehost_test.go +548 -0
  19. package/extensions/a2a/fakepig_test.go +212 -0
  20. package/extensions/a2a/gaps_test.go +59 -0
  21. package/extensions/a2a/go.mod +15 -0
  22. package/extensions/a2a/go.sum +14 -0
  23. package/extensions/a2a/interop_test.go +163 -0
  24. package/extensions/a2a/procattr_other.go +25 -0
  25. package/extensions/a2a/procattr_windows.go +26 -0
  26. package/extensions/a2a/resubscribe_test.go +113 -0
  27. package/extensions/a2a/review_test.go +293 -0
  28. package/extensions/a2a/server.go +267 -0
  29. package/extensions/a2a/server_test.go +792 -0
  30. package/extensions/a2a/survivors_test.go +302 -0
  31. package/extensions/a2a/worker.go +388 -0
  32. package/extensions/a2a/worker_test.go +257 -0
  33. package/package.json +41 -0
  34. package/port/PORT.md +126 -0
  35. package/port/a2a-go-LICENSE +201 -0
  36. package/port/golden/flag-without-auth-refused.jsonl +5 -0
  37. package/port/golden/listener-off.jsonl +4 -0
  38. package/port/golden/send-missing-message.jsonl +18 -0
  39. package/port/golden/send-without-remotes.jsonl +18 -0
  40. package/port/golden/task-unknown-action.jsonl +18 -0
  41. package/port/golden/tools-visible-to-model.jsonl +11 -0
  42. package/port/interop/kagent/main.go +73 -0
  43. package/port/mutation-run.txt +222 -0
  44. package/port/mutations.json +656 -0
  45. package/port/red.txt +105 -0
  46. package/port/scenarios/flag-without-auth-refused.json +8 -0
  47. package/port/scenarios/listener-off.json +7 -0
  48. package/port/scenarios/send-missing-message.json +11 -0
  49. package/port/scenarios/send-without-remotes.json +11 -0
  50. package/port/scenarios/task-unknown-action.json +11 -0
  51. package/port/scenarios/tools-visible-to-model.json +8 -0
  52. package/provenance.json +18 -0
@@ -0,0 +1,792 @@
1
+ package a2aext
2
+
3
+ import (
4
+ "bytes"
5
+ "context"
6
+ "crypto/ecdsa"
7
+ "crypto/elliptic"
8
+ "crypto/rand"
9
+ "crypto/tls"
10
+ "crypto/x509"
11
+ "crypto/x509/pkix"
12
+ "encoding/json"
13
+ "encoding/pem"
14
+ "errors"
15
+ "io"
16
+ "math/big"
17
+ "net"
18
+ "net/http"
19
+ "os"
20
+ "path/filepath"
21
+ "strings"
22
+ "sync"
23
+ "sync/atomic"
24
+ "testing"
25
+ "time"
26
+
27
+ "github.com/a2aproject/a2a-go/v2/a2a"
28
+ "github.com/a2aproject/a2a-go/v2/a2aclient"
29
+ )
30
+
31
+ // scriptedWorker is a Worker whose behaviour a test sets.
32
+ type scriptedWorker struct {
33
+ mu sync.Mutex
34
+ turns []Turn
35
+ run func(ctx context.Context, t Turn, up func(Update)) (Result, error)
36
+ running atomic.Int32
37
+ maxRun atomic.Int32
38
+ }
39
+
40
+ func (w *scriptedWorker) Run(ctx context.Context, t Turn, up func(Update)) (Result, error) {
41
+ w.mu.Lock()
42
+ w.turns = append(w.turns, t)
43
+ w.mu.Unlock()
44
+ n := w.running.Add(1)
45
+ defer w.running.Add(-1)
46
+ for {
47
+ m := w.maxRun.Load()
48
+ if n <= m || w.maxRun.CompareAndSwap(m, n) {
49
+ break
50
+ }
51
+ }
52
+ if w.run != nil {
53
+ return w.run(ctx, t, up)
54
+ }
55
+ up(Update{Text: "echo: "})
56
+ up(Update{Text: t.Prompt})
57
+ return Result{Text: "echo: " + t.Prompt}, nil
58
+ }
59
+
60
+ func (w *scriptedWorker) Turns() []Turn {
61
+ w.mu.Lock()
62
+ defer w.mu.Unlock()
63
+ return append([]Turn(nil), w.turns...)
64
+ }
65
+
66
+ func serverConfig() Config {
67
+ return Config{
68
+ Listen: "127.0.0.1:0",
69
+ Tokens: []TokenConfig{{Name: "alice", TokenEnv: "TOKEN_A", Tenant: "team-a"}, {Name: "bob", TokenEnv: "TOKEN_B"}},
70
+ Name: "pig-under-test", MaxConcurrentTasks: 4, TaskTimeoutSeconds: 60,
71
+ }
72
+ }
73
+
74
+ var serverEnv = envFrom(map[string]string{"TOKEN_A": tokenA, "TOKEN_B": tokenB})
75
+
76
+ func startServer(t *testing.T, cfg Config, w Worker) *Server {
77
+ t.Helper()
78
+ s, err := NewServer(cfg, w, serverEnv)
79
+ if err != nil {
80
+ t.Fatalf("NewServer: %v", err)
81
+ }
82
+ if err := s.Start(); err != nil {
83
+ t.Fatalf("Start: %v", err)
84
+ }
85
+ t.Cleanup(func() {
86
+ ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
87
+ defer cancel()
88
+ _ = s.Shutdown(ctx)
89
+ })
90
+ return s
91
+ }
92
+
93
+ type bearerTransport struct {
94
+ token string
95
+ base http.RoundTripper
96
+ }
97
+
98
+ func (b bearerTransport) RoundTrip(r *http.Request) (*http.Response, error) {
99
+ r = r.Clone(r.Context())
100
+ if b.token != "" {
101
+ r.Header.Set("Authorization", "Bearer "+b.token)
102
+ }
103
+ base := b.base
104
+ if base == nil {
105
+ base = http.DefaultTransport
106
+ }
107
+ return base.RoundTrip(r)
108
+ }
109
+
110
+ func a2aClient(t *testing.T, baseURL, token string) *a2aclient.Client {
111
+ t.Helper()
112
+ hc := &http.Client{Transport: bearerTransport{token: token}, Timeout: 30 * time.Second}
113
+ c, err := a2aclient.NewFromEndpoints(context.Background(),
114
+ []*a2a.AgentInterface{a2a.NewAgentInterface(baseURL, a2a.TransportProtocolJSONRPC)},
115
+ a2aclient.WithJSONRPCTransport(hc))
116
+ if err != nil {
117
+ t.Fatal(err)
118
+ }
119
+ return c
120
+ }
121
+
122
+ func textMessage(text string) *a2a.Message {
123
+ return a2a.NewMessage(a2a.MessageRoleUser, a2a.NewTextPart(text))
124
+ }
125
+
126
+ func taskText(task *a2a.Task) string {
127
+ var b strings.Builder
128
+ for _, art := range task.Artifacts {
129
+ for _, p := range art.Parts {
130
+ b.WriteString(p.Text())
131
+ }
132
+ }
133
+ return b.String()
134
+ }
135
+
136
+ func mustTask(t *testing.T, res a2a.SendMessageResult, err error) *a2a.Task {
137
+ t.Helper()
138
+ if err != nil {
139
+ t.Fatalf("send: %v", err)
140
+ }
141
+ task, ok := res.(*a2a.Task)
142
+ if !ok {
143
+ t.Fatalf("want a task, got %T", res)
144
+ }
145
+ return task
146
+ }
147
+
148
+ func TestAgentCardIsPublicAndPinnedToProtocol10(t *testing.T) {
149
+ cfg := serverConfig()
150
+ s := startServer(t, cfg, &scriptedWorker{})
151
+ resp, err := http.Get("http://" + s.Addr() + "/.well-known/agent-card.json")
152
+ if err != nil {
153
+ t.Fatal(err)
154
+ }
155
+ defer resp.Body.Close()
156
+ if resp.StatusCode != 200 {
157
+ t.Fatalf("the agent card is public: status %d", resp.StatusCode)
158
+ }
159
+ var card a2a.AgentCard
160
+ if err := json.NewDecoder(resp.Body).Decode(&card); err != nil {
161
+ t.Fatal(err)
162
+ }
163
+ if card.Name != "pig-under-test" || !card.Capabilities.Streaming || card.Capabilities.PushNotifications {
164
+ t.Fatalf("card %+v", card)
165
+ }
166
+ if len(card.SupportedInterfaces) != 1 {
167
+ t.Fatalf("interfaces %+v", card.SupportedInterfaces)
168
+ }
169
+ i := card.SupportedInterfaces[0]
170
+ if i.ProtocolVersion != "1.0" || i.ProtocolBinding != a2a.TransportProtocolJSONRPC || i.URL != "http://"+s.Addr() {
171
+ t.Fatalf("interface %+v", i)
172
+ }
173
+ if len(card.SecuritySchemes) == 0 || len(card.SecurityRequirements) == 0 {
174
+ t.Fatalf("the card must declare bearer authentication: %+v", card)
175
+ }
176
+ if len(card.Skills) == 0 {
177
+ t.Fatal("a card needs at least one skill")
178
+ }
179
+ }
180
+
181
+ func TestAgentCardAdvertisesExternalURL(t *testing.T) {
182
+ cfg := serverConfig()
183
+ cfg.ExternalURL = "https://pig.example.test/a2a"
184
+ s := startServer(t, cfg, &scriptedWorker{})
185
+ resp, err := http.Get("http://" + s.Addr() + "/.well-known/agent-card.json")
186
+ if err != nil {
187
+ t.Fatal(err)
188
+ }
189
+ defer resp.Body.Close()
190
+ var card a2a.AgentCard
191
+ _ = json.NewDecoder(resp.Body).Decode(&card)
192
+ if len(card.SupportedInterfaces) != 1 || card.SupportedInterfaces[0].URL != "https://pig.example.test/a2a" {
193
+ t.Fatalf("%+v", card.SupportedInterfaces)
194
+ }
195
+ }
196
+
197
+ func TestUnauthenticatedRequestIsRejectedBeforeTheWorker(t *testing.T) {
198
+ w := &scriptedWorker{}
199
+ s := startServer(t, serverConfig(), w)
200
+ for name, token := range map[string]string{"none": "", "wrong": "nope-nope-nope-nope-nope-nope-nope"} {
201
+ c := a2aClient(t, "http://"+s.Addr(), token)
202
+ if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("hi")}); err == nil {
203
+ t.Fatalf("%s: want an error", name)
204
+ }
205
+ }
206
+ if len(w.Turns()) != 0 {
207
+ t.Fatal("an unauthenticated request reached the worker")
208
+ }
209
+ }
210
+
211
+ func TestSendMessageRunsATurnAndReturnsACompletedTask(t *testing.T) {
212
+ w := &scriptedWorker{}
213
+ s := startServer(t, serverConfig(), w)
214
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
215
+ res, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("hello pig")})
216
+ task := mustTask(t, res, err)
217
+ if task.Status.State != a2a.TaskStateCompleted {
218
+ t.Fatalf("state %s", task.Status.State)
219
+ }
220
+ if taskText(task) != "echo: hello pig" {
221
+ t.Fatalf("artifact text %q", taskText(task))
222
+ }
223
+ if task.ID == "" || task.ContextID == "" {
224
+ t.Fatalf("ids %q %q", task.ID, task.ContextID)
225
+ }
226
+ turns := w.Turns()
227
+ if len(turns) != 1 || turns[0].Prompt != "hello pig" || turns[0].Principal.Name != "alice" || turns[0].ContextID != task.ContextID || turns[0].TaskID != string(task.ID) {
228
+ t.Fatalf("turn %+v", turns)
229
+ }
230
+ }
231
+
232
+ func TestSameContextIDContinuesTheSameWorkerContext(t *testing.T) {
233
+ w := &scriptedWorker{}
234
+ s := startServer(t, serverConfig(), w)
235
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
236
+ first := sendTask(t, c, &a2a.SendMessageRequest{Message: textMessage("one")})
237
+ m2 := textMessage("two")
238
+ m2.ContextID = first.ContextID
239
+ second := sendTask(t, c, &a2a.SendMessageRequest{Message: m2})
240
+ if second.ContextID != first.ContextID || second.ID == first.ID {
241
+ t.Fatalf("second task %s/%s in first %s/%s: a context spans tasks", second.ID, second.ContextID, first.ID, first.ContextID)
242
+ }
243
+ turns := w.Turns()
244
+ if len(turns) != 2 || turns[0].ContextID != turns[1].ContextID {
245
+ t.Fatalf("turns %+v", turns)
246
+ }
247
+ }
248
+
249
+ func mustSend(c *a2aclient.Client, req *a2a.SendMessageRequest) (a2a.SendMessageResult, error) {
250
+ return c.SendMessage(context.Background(), req)
251
+ }
252
+
253
+ func TestTenantsDoNotShareContextsOrTasks(t *testing.T) {
254
+ w := &scriptedWorker{}
255
+ s := startServer(t, serverConfig(), w)
256
+ ca := a2aClient(t, "http://"+s.Addr(), tokenA)
257
+ cb := a2aClient(t, "http://"+s.Addr(), tokenB)
258
+ ma := textMessage("from alice")
259
+ ma.ContextID = "shared-context"
260
+ taskA := sendTask(t, ca, &a2a.SendMessageRequest{Message: ma})
261
+ mb := textMessage("from bob")
262
+ mb.ContextID = "shared-context"
263
+ sendTask(t, cb, &a2a.SendMessageRequest{Message: mb})
264
+ turns := w.Turns()
265
+ if len(turns) != 2 || turns[0].Principal.Key() == turns[1].Principal.Key() {
266
+ t.Fatalf("the same context id under two principals must reach the worker as two identities: %+v", turns)
267
+ }
268
+ if _, err := cb.GetTask(context.Background(), &a2a.GetTaskRequest{ID: taskA.ID}); !errors.Is(err, a2a.ErrTaskNotFound) {
269
+ t.Fatalf("bob read alice's task: %v", err)
270
+ }
271
+ if _, err := cb.CancelTask(context.Background(), &a2a.CancelTaskRequest{ID: taskA.ID}); err == nil {
272
+ t.Fatal("bob cancelled alice's task")
273
+ }
274
+ list, err := cb.ListTasks(context.Background(), &a2a.ListTasksRequest{})
275
+ if err != nil {
276
+ t.Fatal(err)
277
+ }
278
+ for _, task := range list.Tasks {
279
+ if task.ID == taskA.ID {
280
+ t.Fatal("bob's task list shows alice's task")
281
+ }
282
+ }
283
+ if got, err := ca.GetTask(context.Background(), &a2a.GetTaskRequest{ID: taskA.ID}); err != nil || got.ID != taskA.ID {
284
+ t.Fatalf("alice must read her own task: %v", err)
285
+ }
286
+ }
287
+
288
+ func TestBobCannotAttachToAliceTaskWithAMessage(t *testing.T) {
289
+ w := &scriptedWorker{}
290
+ s := startServer(t, serverConfig(), w)
291
+ ca := a2aClient(t, "http://"+s.Addr(), tokenA)
292
+ cb := a2aClient(t, "http://"+s.Addr(), tokenB)
293
+ taskA := sendTask(t, ca, &a2a.SendMessageRequest{Message: textMessage("x y")})
294
+ m := textMessage("hijack")
295
+ m.TaskID, m.ContextID = taskA.ID, taskA.ContextID
296
+ before := len(w.Turns())
297
+ if _, err := cb.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: m}); err == nil {
298
+ t.Fatal("bob continued alice's task")
299
+ }
300
+ if len(w.Turns()) != before {
301
+ t.Fatal("the worker ran for a foreign task")
302
+ }
303
+ }
304
+
305
+ func TestRequestTenantMustMatchTheTokensTenant(t *testing.T) {
306
+ w := &scriptedWorker{}
307
+ s := startServer(t, serverConfig(), w)
308
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
309
+ if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Tenant: "team-b", Message: textMessage("x y")}); err == nil {
310
+ t.Fatal("a caller must not name another tenant than its credential's")
311
+ }
312
+ if len(w.Turns()) != 0 {
313
+ t.Fatal("worker ran for a foreign tenant")
314
+ }
315
+ if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Tenant: "team-a", Message: textMessage("x y")}); err != nil {
316
+ t.Fatalf("its own tenant is fine: %v", err)
317
+ }
318
+ }
319
+
320
+ func rawRPC(t *testing.T, url, token, version, method string, params any) (int, map[string]any) {
321
+ t.Helper()
322
+ body, _ := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": 1, "method": method, "params": params})
323
+ req, _ := http.NewRequest(http.MethodPost, url, bytes.NewReader(body))
324
+ req.Header.Set("Content-Type", "application/json")
325
+ req.Header.Set("Authorization", "Bearer "+token)
326
+ if version != "" {
327
+ req.Header.Set("A2A-Version", version)
328
+ }
329
+ resp, err := http.DefaultClient.Do(req)
330
+ if err != nil {
331
+ t.Fatal(err)
332
+ }
333
+ defer resp.Body.Close()
334
+ var out map[string]any
335
+ _ = json.NewDecoder(resp.Body).Decode(&out)
336
+ return resp.StatusCode, out
337
+ }
338
+
339
+ func TestProtocolVersionIsPinned(t *testing.T) {
340
+ w := &scriptedWorker{}
341
+ s := startServer(t, serverConfig(), w)
342
+ params := map[string]any{"message": map[string]any{"messageId": "m1", "role": "ROLE_USER", "parts": []any{map[string]any{"text": "hello there"}}}}
343
+ for name, version := range map[string]string{"0.3": "0.3", "2.0": "2.0", "absent (spec: 0.3)": "", "garbage": "one"} {
344
+ _, out := rawRPC(t, "http://"+s.Addr(), tokenA, version, "SendMessage", params)
345
+ e, _ := out["error"].(map[string]any)
346
+ if e == nil {
347
+ t.Fatalf("%s: want a JSON-RPC error, got %v", name, out)
348
+ }
349
+ if msg, _ := e["message"].(string); !strings.Contains(strings.ToLower(msg), "version") {
350
+ t.Fatalf("%s: error should say the version is unsupported: %v", name, e)
351
+ }
352
+ }
353
+ if len(w.Turns()) != 0 {
354
+ t.Fatal("a request with an unsupported protocol version reached the worker")
355
+ }
356
+ _, out := rawRPC(t, "http://"+s.Addr(), tokenA, "1.0", "SendMessage", params)
357
+ if out["error"] != nil || out["result"] == nil {
358
+ t.Fatalf("1.0 must work: %v", out)
359
+ }
360
+ }
361
+
362
+ func TestStreamingDeliversArtifactChunksBeforeCompletion(t *testing.T) {
363
+ gate := make(chan struct{})
364
+ w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
365
+ up(Update{Text: "first "})
366
+ select {
367
+ case <-gate:
368
+ case <-ctx.Done():
369
+ return Result{}, ctx.Err()
370
+ }
371
+ up(Update{Text: "second"})
372
+ return Result{Text: "first second"}, nil
373
+ }}
374
+ s := startServer(t, serverConfig(), w)
375
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
376
+ var seenChunkBeforeGate bool
377
+ var final a2a.TaskState
378
+ var text strings.Builder
379
+ for ev, err := range c.SendStreamingMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("stream me")}) {
380
+ if err != nil {
381
+ t.Fatal(err)
382
+ }
383
+ switch e := ev.(type) {
384
+ case *a2a.TaskArtifactUpdateEvent:
385
+ for _, p := range e.Artifact.Parts {
386
+ text.WriteString(p.Text())
387
+ }
388
+ if strings.Contains(text.String(), "first") && !seenChunkBeforeGate {
389
+ seenChunkBeforeGate = true
390
+ close(gate)
391
+ }
392
+ case *a2a.TaskStatusUpdateEvent:
393
+ final = e.Status.State
394
+ }
395
+ }
396
+ if !seenChunkBeforeGate {
397
+ t.Fatal("no artifact chunk arrived before the worker finished")
398
+ }
399
+ if final != a2a.TaskStateCompleted || text.String() != "first second" {
400
+ t.Fatalf("final %s text %q", final, text.String())
401
+ }
402
+ }
403
+
404
+ func TestCancelTaskAbortsTheWorkerAndEndsCanceled(t *testing.T) {
405
+ started := make(chan struct{})
406
+ var ctxErr atomic.Value
407
+ w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
408
+ close(started)
409
+ <-ctx.Done()
410
+ ctxErr.Store(ctx.Err())
411
+ return Result{}, ctx.Err()
412
+ }}
413
+ s := startServer(t, serverConfig(), w)
414
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
415
+ var taskID a2a.TaskID
416
+ events := c.SendStreamingMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("long job")})
417
+ done := make(chan a2a.TaskState, 1)
418
+ go func() {
419
+ var last a2a.TaskState
420
+ for ev, err := range events {
421
+ if err != nil {
422
+ break
423
+ }
424
+ switch e := ev.(type) {
425
+ case *a2a.Task:
426
+ taskID, last = e.ID, e.Status.State
427
+ case *a2a.TaskStatusUpdateEvent:
428
+ last = e.Status.State
429
+ }
430
+ }
431
+ done <- last
432
+ }()
433
+ waitChan(t, started, "the worker to start")
434
+ var task *a2a.Task
435
+ deadline := time.Now().Add(5 * time.Second)
436
+ for task == nil && time.Now().Before(deadline) {
437
+ list, err := c.ListTasks(context.Background(), &a2a.ListTasksRequest{})
438
+ if err == nil && len(list.Tasks) == 1 {
439
+ task = list.Tasks[0]
440
+ }
441
+ time.Sleep(10 * time.Millisecond)
442
+ }
443
+ if task == nil {
444
+ t.Fatal("running task not listed")
445
+ }
446
+ got, err := c.CancelTask(context.Background(), &a2a.CancelTaskRequest{ID: task.ID})
447
+ if err != nil {
448
+ t.Fatalf("cancel: %v", err)
449
+ }
450
+ if got.Status.State != a2a.TaskStateCanceled {
451
+ t.Fatalf("state %s", got.Status.State)
452
+ }
453
+ select {
454
+ case last := <-done:
455
+ if last != a2a.TaskStateCanceled {
456
+ t.Fatalf("stream ended in %s", last)
457
+ }
458
+ case <-time.After(10 * time.Second):
459
+ t.Fatal("stream did not end")
460
+ }
461
+ if !errors.Is(ctxErr.Load().(error), context.Canceled) {
462
+ t.Fatalf("worker context error %v", ctxErr.Load())
463
+ }
464
+ _ = taskID
465
+ waitFor(t, func() bool { return s.ActiveTasks() == 0 }, "active tasks to drain")
466
+ }
467
+
468
+ func waitFor(t *testing.T, ok func() bool, what string) {
469
+ t.Helper()
470
+ deadline := time.Now().Add(10 * time.Second)
471
+ for !ok() {
472
+ if time.Now().After(deadline) {
473
+ t.Fatalf("timed out waiting for %s", what)
474
+ }
475
+ time.Sleep(10 * time.Millisecond)
476
+ }
477
+ }
478
+
479
+ func TestWorkerFailureBecomesAFailedTask(t *testing.T) {
480
+ w := &scriptedWorker{run: func(context.Context, Turn, func(Update)) (Result, error) {
481
+ return Result{Failure: "the model call failed"}, nil
482
+ }}
483
+ s := startServer(t, serverConfig(), w)
484
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
485
+ task := sendTask(t, c, &a2a.SendMessageRequest{Message: textMessage("x y")})
486
+ if task.Status.State != a2a.TaskStateFailed {
487
+ t.Fatalf("state %s", task.Status.State)
488
+ }
489
+ if task.Status.Message == nil || !strings.Contains(task.Status.Message.Parts[0].Text(), "the model call failed") {
490
+ t.Fatalf("status message %+v", task.Status.Message)
491
+ }
492
+ }
493
+
494
+ func TestWorkerErrorDetailIsNotSentToThePeer(t *testing.T) {
495
+ w := &scriptedWorker{run: func(context.Context, Turn, func(Update)) (Result, error) {
496
+ return Result{}, errors.New("exec /opt/secret/pig: permission denied")
497
+ }}
498
+ s := startServer(t, serverConfig(), w)
499
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
500
+ task := sendTask(t, c, &a2a.SendMessageRequest{Message: textMessage("x y")})
501
+ if task.Status.State != a2a.TaskStateFailed {
502
+ t.Fatalf("state %s", task.Status.State)
503
+ }
504
+ raw, _ := json.Marshal(task)
505
+ if strings.Contains(string(raw), "/opt/secret") {
506
+ t.Fatalf("internal error detail leaked to the peer: %s", raw)
507
+ }
508
+ }
509
+
510
+ func TestConcurrencyLimitQueuesTasks(t *testing.T) {
511
+ release := make(chan struct{})
512
+ w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
513
+ select {
514
+ case <-release:
515
+ case <-ctx.Done():
516
+ }
517
+ return Result{Text: "ok"}, nil
518
+ }}
519
+ cfg := serverConfig()
520
+ cfg.MaxConcurrentTasks = 1
521
+ s := startServer(t, cfg, w)
522
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
523
+ var wg sync.WaitGroup
524
+ for i := 0; i < 3; i++ {
525
+ wg.Add(1)
526
+ go func() {
527
+ defer wg.Done()
528
+ _, _ = c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("job job")})
529
+ }()
530
+ }
531
+ waitFor(t, func() bool { return w.running.Load() == 1 }, "the first task to start")
532
+ time.Sleep(200 * time.Millisecond)
533
+ if w.running.Load() != 1 {
534
+ t.Fatalf("%d tasks running with a limit of 1", w.running.Load())
535
+ }
536
+ close(release)
537
+ wg.Wait()
538
+ if w.maxRun.Load() != 1 {
539
+ t.Fatalf("max concurrent %d", w.maxRun.Load())
540
+ }
541
+ }
542
+
543
+ func TestTasksInOneContextRunOneAtATime(t *testing.T) {
544
+ w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
545
+ time.Sleep(100 * time.Millisecond)
546
+ return Result{Text: "ok"}, nil
547
+ }}
548
+ s := startServer(t, serverConfig(), w)
549
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
550
+ var wg sync.WaitGroup
551
+ for i := 0; i < 3; i++ {
552
+ wg.Add(1)
553
+ go func() {
554
+ defer wg.Done()
555
+ m := textMessage("same ctx")
556
+ m.ContextID = "one-context"
557
+ _, _ = c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: m})
558
+ }()
559
+ }
560
+ wg.Wait()
561
+ if w.maxRun.Load() != 1 || len(w.Turns()) != 3 {
562
+ t.Fatalf("a session file has one writer: max concurrent %d, turns %d", w.maxRun.Load(), len(w.Turns()))
563
+ }
564
+ }
565
+
566
+ func TestNonTextPartsAreRejected(t *testing.T) {
567
+ w := &scriptedWorker{}
568
+ s := startServer(t, serverConfig(), w)
569
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
570
+ m := a2a.NewMessage(a2a.MessageRoleUser, a2a.NewRawPart([]byte("binary")))
571
+ if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: m}); err == nil {
572
+ t.Fatal("PiG tasks take text; a file part must be refused, not dropped")
573
+ }
574
+ if len(w.Turns()) != 0 {
575
+ t.Fatal("worker ran")
576
+ }
577
+ }
578
+
579
+ func TestEmptyMessageIsRejected(t *testing.T) {
580
+ w := &scriptedWorker{}
581
+ s := startServer(t, serverConfig(), w)
582
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
583
+ if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage(" ")}); err == nil {
584
+ t.Fatal("empty prompt")
585
+ }
586
+ if len(w.Turns()) != 0 {
587
+ t.Fatal("worker ran")
588
+ }
589
+ }
590
+
591
+ func TestTaskTimeoutFailsTheTask(t *testing.T) {
592
+ var ctxDone atomic.Bool
593
+ w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
594
+ <-ctx.Done()
595
+ ctxDone.Store(true)
596
+ return Result{}, ctx.Err()
597
+ }}
598
+ cfg := serverConfig()
599
+ cfg.TaskTimeoutSeconds = 1
600
+ s := startServer(t, cfg, w)
601
+ c := a2aClient(t, "http://"+s.Addr(), tokenA)
602
+ task := sendTask(t, c, &a2a.SendMessageRequest{Message: textMessage("never ends")})
603
+ if task.Status.State != a2a.TaskStateFailed || !ctxDone.Load() {
604
+ t.Fatalf("state %s ctxDone %v", task.Status.State, ctxDone.Load())
605
+ }
606
+ if task.Status.Message == nil || !strings.Contains(strings.ToLower(task.Status.Message.Parts[0].Text()), "time") {
607
+ t.Fatalf("the failure should say the task timed out: %+v", task.Status.Message)
608
+ }
609
+ }
610
+
611
+ func TestShutdownCancelsRunningTasksAndStopsListening(t *testing.T) {
612
+ started := make(chan struct{})
613
+ var canceled atomic.Bool
614
+ w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
615
+ close(started)
616
+ <-ctx.Done()
617
+ canceled.Store(true)
618
+ return Result{}, ctx.Err()
619
+ }}
620
+ s, err := NewServer(serverConfig(), w, serverEnv)
621
+ if err != nil {
622
+ t.Fatal(err)
623
+ }
624
+ if err := s.Start(); err != nil {
625
+ t.Fatal(err)
626
+ }
627
+ addr := s.Addr()
628
+ c := a2aClient(t, "http://"+addr, tokenA)
629
+ go func() {
630
+ _, _ = c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("long job")})
631
+ }()
632
+ waitChan(t, started, "the worker to start")
633
+ ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
634
+ defer cancel()
635
+ if err := s.Shutdown(ctx); err != nil {
636
+ t.Fatalf("shutdown: %v", err)
637
+ }
638
+ if !canceled.Load() {
639
+ t.Fatal("shutdown must cancel running turns")
640
+ }
641
+ if conn, err := net.DialTimeout("tcp", addr, time.Second); err == nil {
642
+ conn.Close()
643
+ t.Fatal("still listening after shutdown")
644
+ }
645
+ if s.ActiveTasks() != 0 {
646
+ t.Fatalf("%d active tasks after shutdown", s.ActiveTasks())
647
+ }
648
+ }
649
+
650
+ func TestStartFailsOnAnAddressInUse(t *testing.T) {
651
+ ln, err := net.Listen("tcp", "127.0.0.1:0")
652
+ if err != nil {
653
+ t.Fatal(err)
654
+ }
655
+ defer ln.Close()
656
+ cfg := serverConfig()
657
+ cfg.Listen = ln.Addr().String()
658
+ s, err := NewServer(cfg, &scriptedWorker{}, serverEnv)
659
+ if err != nil {
660
+ t.Fatal(err)
661
+ }
662
+ if err := s.Start(); err == nil {
663
+ t.Fatal("want a bind error, not a silent no-op")
664
+ }
665
+ }
666
+
667
+ func TestOversizedRequestBodyIsRefused(t *testing.T) {
668
+ w := &scriptedWorker{}
669
+ s := startServer(t, serverConfig(), w)
670
+ big := strings.Repeat("a", 3<<20)
671
+ body, _ := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": 1, "method": "SendMessage", "params": map[string]any{
672
+ "message": map[string]any{"messageId": "m", "role": "ROLE_USER", "parts": []any{map[string]any{"text": big}}}}})
673
+ req, _ := http.NewRequest(http.MethodPost, "http://"+s.Addr(), bytes.NewReader(body))
674
+ req.Header.Set("Authorization", "Bearer "+tokenA)
675
+ req.Header.Set("A2A-Version", "1.0")
676
+ resp, err := http.DefaultClient.Do(req)
677
+ if err == nil {
678
+ defer resp.Body.Close()
679
+ io.Copy(io.Discard, resp.Body)
680
+ if resp.StatusCode == 200 {
681
+ var out map[string]any
682
+ _ = out
683
+ }
684
+ }
685
+ if len(w.Turns()) != 0 {
686
+ t.Fatal("a 3 MiB request reached the worker; bodies are capped")
687
+ }
688
+ }
689
+
690
+ func selfSignedPair(t *testing.T) (certFile, keyFile string, pool *x509.CertPool) {
691
+ t.Helper()
692
+ key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
693
+ if err != nil {
694
+ t.Fatal(err)
695
+ }
696
+ tmpl := &x509.Certificate{
697
+ SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "127.0.0.1"},
698
+ NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour),
699
+ KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
700
+ IPAddresses: []net.IP{net.ParseIP("127.0.0.1")},
701
+ }
702
+ der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
703
+ if err != nil {
704
+ t.Fatal(err)
705
+ }
706
+ dir := t.TempDir()
707
+ certFile, keyFile = filepath.Join(dir, "c.pem"), filepath.Join(dir, "k.pem")
708
+ if err := os.WriteFile(certFile, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600); err != nil {
709
+ t.Fatal(err)
710
+ }
711
+ kb, _ := x509.MarshalECPrivateKey(key)
712
+ if err := os.WriteFile(keyFile, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: kb}), 0o600); err != nil {
713
+ t.Fatal(err)
714
+ }
715
+ cert, _ := x509.ParseCertificate(der)
716
+ pool = x509.NewCertPool()
717
+ pool.AddCert(cert)
718
+ return
719
+ }
720
+
721
+ func TestTLSListener(t *testing.T) {
722
+ certFile, keyFile, pool := selfSignedPair(t)
723
+ cfg := serverConfig()
724
+ cfg.TLS = &TLSConfig{CertFile: certFile, KeyFile: keyFile}
725
+ s := startServer(t, cfg, &scriptedWorker{})
726
+ hc := &http.Client{Transport: bearerTransport{token: tokenA, base: &http.Transport{TLSClientConfig: &tls.Config{RootCAs: pool}}}}
727
+ c, err := a2aclient.NewFromEndpoints(context.Background(),
728
+ []*a2a.AgentInterface{a2a.NewAgentInterface("https://"+s.Addr(), a2a.TransportProtocolJSONRPC)}, a2aclient.WithJSONRPCTransport(hc))
729
+ if err != nil {
730
+ t.Fatal(err)
731
+ }
732
+ task := sendTask(t, c, &a2a.SendMessageRequest{Message: textMessage("over tls")})
733
+ if task.Status.State != a2a.TaskStateCompleted {
734
+ t.Fatalf("state %s", task.Status.State)
735
+ }
736
+ if plain, err := http.Get("http://" + s.Addr() + "/.well-known/agent-card.json"); err == nil {
737
+ defer plain.Body.Close()
738
+ if plain.StatusCode == 200 {
739
+ t.Fatal("plain HTTP must not be served on a TLS listener")
740
+ }
741
+ }
742
+ }
743
+
744
+ func sendTask(t *testing.T, c *a2aclient.Client, req *a2a.SendMessageRequest) *a2a.Task {
745
+ t.Helper()
746
+ res, err := c.SendMessage(context.Background(), req)
747
+ return mustTask(t, res, err)
748
+ }
749
+
750
+ func waitChan(t *testing.T, ch <-chan struct{}, what string) {
751
+ t.Helper()
752
+ select {
753
+ case <-ch:
754
+ case <-time.After(10 * time.Second):
755
+ t.Fatalf("timed out waiting for %s", what)
756
+ }
757
+ }
758
+
759
+ func TestUnreadableTLSKeyPairFailsStart(t *testing.T) { // tls-load-failure-ignored
760
+ cfg := serverConfig()
761
+ cfg.TLS = &TLSConfig{CertFile: filepath.Join(t.TempDir(), "missing.pem"), KeyFile: filepath.Join(t.TempDir(), "missing.key")}
762
+ s, err := NewServer(cfg, &scriptedWorker{}, serverEnv)
763
+ if err != nil {
764
+ return // refused at construction is fine too
765
+ }
766
+ if err := s.Start(); err == nil {
767
+ _ = s.Shutdown(context.Background())
768
+ t.Fatal("a listener that cannot load its certificate must not start (nor fall back to plain HTTP)")
769
+ }
770
+ }
771
+
772
+ // http.Server only closes listeners its Serve goroutine has registered; a Shutdown that wins the race with that
773
+ // goroutine must still close the port before it returns (seen as a flaky reload test at -race -count=24 under load).
774
+ func TestShutdownClosesThePortEvenRightAfterStart(t *testing.T) {
775
+ for i := 0; i < 300; i++ {
776
+ s, err := NewServer(serverConfig(), &scriptedWorker{}, serverEnv)
777
+ if err != nil {
778
+ t.Fatal(err)
779
+ }
780
+ if err := s.Start(); err != nil {
781
+ t.Fatal(err)
782
+ }
783
+ addr := s.Addr()
784
+ if err := s.Shutdown(context.Background()); err != nil {
785
+ t.Fatal(err)
786
+ }
787
+ if c, err := net.DialTimeout("tcp", addr, time.Second); err == nil {
788
+ c.Close()
789
+ t.Fatalf("iteration %d: %s still accepts connections after Shutdown returned", i, addr)
790
+ }
791
+ }
792
+ }