@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.
- package/CREDITS.md +26 -0
- package/LICENSE +21 -0
- package/README.md +97 -0
- package/extensions/a2a/auth.go +120 -0
- package/extensions/a2a/auth_test.go +133 -0
- package/extensions/a2a/bench_test.go +70 -0
- package/extensions/a2a/binary_test.go +225 -0
- package/extensions/a2a/cancelqueued_test.go +142 -0
- package/extensions/a2a/card_cache_test.go +42 -0
- package/extensions/a2a/client.go +499 -0
- package/extensions/a2a/client_test.go +402 -0
- package/extensions/a2a/config.go +282 -0
- package/extensions/a2a/config_test.go +194 -0
- package/extensions/a2a/e2e_test.go +311 -0
- package/extensions/a2a/executor.go +318 -0
- package/extensions/a2a/extension.go +382 -0
- package/extensions/a2a/extension_test.go +396 -0
- package/extensions/a2a/fakehost_test.go +548 -0
- package/extensions/a2a/fakepig_test.go +212 -0
- package/extensions/a2a/gaps_test.go +59 -0
- package/extensions/a2a/go.mod +15 -0
- package/extensions/a2a/go.sum +14 -0
- package/extensions/a2a/interop_test.go +163 -0
- package/extensions/a2a/procattr_other.go +25 -0
- package/extensions/a2a/procattr_windows.go +26 -0
- package/extensions/a2a/resubscribe_test.go +113 -0
- package/extensions/a2a/review_test.go +293 -0
- package/extensions/a2a/server.go +267 -0
- package/extensions/a2a/server_test.go +792 -0
- package/extensions/a2a/survivors_test.go +302 -0
- package/extensions/a2a/worker.go +388 -0
- package/extensions/a2a/worker_test.go +257 -0
- package/package.json +41 -0
- package/port/PORT.md +126 -0
- package/port/a2a-go-LICENSE +201 -0
- package/port/golden/flag-without-auth-refused.jsonl +5 -0
- package/port/golden/listener-off.jsonl +4 -0
- package/port/golden/send-missing-message.jsonl +18 -0
- package/port/golden/send-without-remotes.jsonl +18 -0
- package/port/golden/task-unknown-action.jsonl +18 -0
- package/port/golden/tools-visible-to-model.jsonl +11 -0
- package/port/interop/kagent/main.go +73 -0
- package/port/mutation-run.txt +222 -0
- package/port/mutations.json +656 -0
- package/port/red.txt +105 -0
- package/port/scenarios/flag-without-auth-refused.json +8 -0
- package/port/scenarios/listener-off.json +7 -0
- package/port/scenarios/send-missing-message.json +11 -0
- package/port/scenarios/send-without-remotes.json +11 -0
- package/port/scenarios/task-unknown-action.json +11 -0
- package/port/scenarios/tools-visible-to-model.json +8 -0
- 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
|
+
}
|