@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,302 @@
|
|
|
1
|
+
package a2aext
|
|
2
|
+
|
|
3
|
+
// Tests added after the first mutation run (pigeq mutate --unit): each one kills a mutant that the
|
|
4
|
+
// first suite let survive. The mutant's name is in the test's name or comment.
|
|
5
|
+
|
|
6
|
+
import (
|
|
7
|
+
"bytes"
|
|
8
|
+
"context"
|
|
9
|
+
"encoding/json"
|
|
10
|
+
"net/http"
|
|
11
|
+
"net/http/httptest"
|
|
12
|
+
"runtime"
|
|
13
|
+
"strings"
|
|
14
|
+
"testing"
|
|
15
|
+
"time"
|
|
16
|
+
|
|
17
|
+
"github.com/a2aproject/a2a-go/v2/a2a"
|
|
18
|
+
"github.com/a2aproject/a2a-go/v2/a2asrv"
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
func TestTokenEnvNotSetIsSaidPlainly(t *testing.T) { // missing-token-env-allowed
|
|
22
|
+
dir := t.TempDir()
|
|
23
|
+
writeConfig(t, dir, `{"listen":"127.0.0.1:5555","tokens":[{"name":"a","tokenEnv":"A2A_TOKEN_A"}]}`)
|
|
24
|
+
_, err := LoadConfig(LoadOptions{ConfigHome: dir, Getenv: envFrom(nil)})
|
|
25
|
+
if err == nil || !strings.Contains(err.Error(), "not set") {
|
|
26
|
+
t.Fatalf("%v", err)
|
|
27
|
+
}
|
|
28
|
+
}
|
|
29
|
+
|
|
30
|
+
func TestTwoVariablesWithOneValueAreRejected(t *testing.T) { // duplicate-token-values-allowed
|
|
31
|
+
dir := t.TempDir()
|
|
32
|
+
writeConfig(t, dir, `{"listen":"127.0.0.1:1","tokens":[{"name":"a","tokenEnv":"V1"},{"name":"b","tokenEnv":"V2"}]}`)
|
|
33
|
+
same := "0123456789abcdef0123456789abcdef"
|
|
34
|
+
_, err := LoadConfig(LoadOptions{ConfigHome: dir, Getenv: envFrom(map[string]string{"V1": same, "V2": same})})
|
|
35
|
+
if err == nil || !strings.Contains(err.Error(), "same value") {
|
|
36
|
+
t.Fatalf("two callers with one token cannot be told apart: %v", err)
|
|
37
|
+
}
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
func TestInsecureNoAuthWithTokensIsAmbiguous(t *testing.T) { // insecure-and-tokens-both
|
|
41
|
+
dir := t.TempDir()
|
|
42
|
+
writeConfig(t, dir, `{"listen":"127.0.0.1:1","insecureNoAuth":true,"tokens":[{"name":"a","tokenEnv":"T"}]}`)
|
|
43
|
+
_, err := LoadConfig(LoadOptions{ConfigHome: dir, Getenv: envFrom(map[string]string{"T": "0123456789abcdef0123456789abcdef"})})
|
|
44
|
+
if err == nil || !strings.Contains(err.Error(), "insecureNoAuth") {
|
|
45
|
+
t.Fatalf("%v", err)
|
|
46
|
+
}
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
func TestMalformedListenAddressSaysHostPort(t *testing.T) { // listen-not-host-port
|
|
50
|
+
dir := t.TempDir()
|
|
51
|
+
writeConfig(t, dir, `{"listen":"not-an-address","insecureNoAuth":true}`)
|
|
52
|
+
_, err := LoadConfig(LoadOptions{ConfigHome: dir, Getenv: envFrom(nil)})
|
|
53
|
+
if err == nil || !strings.Contains(err.Error(), "host:port") {
|
|
54
|
+
t.Fatalf("%v", err)
|
|
55
|
+
}
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
func TestOversizedPromptIsRefusedBelowTheBodyLimit(t *testing.T) { // prompt-size-unbounded
|
|
59
|
+
w := &scriptedWorker{}
|
|
60
|
+
s := startServer(t, serverConfig(), w)
|
|
61
|
+
c := a2aClient(t, "http://"+s.Addr(), tokenA)
|
|
62
|
+
if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage(strings.Repeat("a", maxPromptBytes+1))}); err == nil {
|
|
63
|
+
t.Fatal("a prompt over the limit must be refused")
|
|
64
|
+
}
|
|
65
|
+
if len(w.Turns()) != 0 {
|
|
66
|
+
t.Fatal("worker ran")
|
|
67
|
+
}
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
func TestBodyLimitAppliesEvenWithASmallPrompt(t *testing.T) { // body-size-unbounded
|
|
71
|
+
w := &scriptedWorker{}
|
|
72
|
+
s := startServer(t, serverConfig(), w)
|
|
73
|
+
body, _ := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": 1, "method": "SendMessage", "params": map[string]any{
|
|
74
|
+
"message": map[string]any{"messageId": "m", "role": "ROLE_USER", "parts": []any{map[string]any{"text": "small prompt"}},
|
|
75
|
+
"metadata": map[string]any{"padding": strings.Repeat("x", 2<<20)}}}})
|
|
76
|
+
req, _ := http.NewRequest(http.MethodPost, "http://"+s.Addr(), bytes.NewReader(body))
|
|
77
|
+
req.Header.Set("Authorization", "Bearer "+tokenA)
|
|
78
|
+
req.Header.Set("A2A-Version", "1.0")
|
|
79
|
+
if resp, err := http.DefaultClient.Do(req); err == nil {
|
|
80
|
+
resp.Body.Close()
|
|
81
|
+
}
|
|
82
|
+
if len(w.Turns()) != 0 {
|
|
83
|
+
t.Fatal("a 2 MiB request with a small prompt reached the worker; the body cap is the first line of defence")
|
|
84
|
+
}
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
func TestTextPlusFilePartIsRefusedNotTrimmed(t *testing.T) { // file-parts-dropped
|
|
88
|
+
w := &scriptedWorker{}
|
|
89
|
+
s := startServer(t, serverConfig(), w)
|
|
90
|
+
c := a2aClient(t, "http://"+s.Addr(), tokenA)
|
|
91
|
+
m := a2a.NewMessage(a2a.MessageRoleUser, a2a.NewTextPart("summarise this"), a2a.NewRawPart([]byte("attachment")))
|
|
92
|
+
if _, err := c.SendMessage(context.Background(), &a2a.SendMessageRequest{Message: m}); err == nil {
|
|
93
|
+
t.Fatal("an attachment the worker cannot see must not be silently dropped")
|
|
94
|
+
}
|
|
95
|
+
if len(w.Turns()) != 0 {
|
|
96
|
+
t.Fatal("worker ran")
|
|
97
|
+
}
|
|
98
|
+
}
|
|
99
|
+
|
|
100
|
+
func TestMixedVersionHeaderIsRefused(t *testing.T) { // mixed-version-values-accepted
|
|
101
|
+
w := &scriptedWorker{}
|
|
102
|
+
s := startServer(t, serverConfig(), w)
|
|
103
|
+
params := map[string]any{"message": map[string]any{"messageId": "m1", "role": "ROLE_USER", "parts": []any{map[string]any{"text": "hello there"}}}}
|
|
104
|
+
_, out := rawRPC(t, "http://"+s.Addr(), tokenA, "1.0, 0.3", "SendMessage", params)
|
|
105
|
+
if out["error"] == nil {
|
|
106
|
+
t.Fatalf("a header that also names 0.3 is not 1.0: %v", out)
|
|
107
|
+
}
|
|
108
|
+
_, out = rawRPC(t, "http://"+s.Addr(), tokenA, "1.0, 1.0", "SendMessage", params)
|
|
109
|
+
if out["error"] != nil {
|
|
110
|
+
t.Fatalf("kagent's client repeats the header; that must work: %v", out)
|
|
111
|
+
}
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
func TestExecutorRefusesAForgedUser(t *testing.T) { // principal-not-checked
|
|
115
|
+
for name, u := range map[string]*a2asrv.User{
|
|
116
|
+
"nil": nil,
|
|
117
|
+
"unauthenticated": {Name: "token:alice", Authenticated: false, Attributes: map[string]any{"name": "alice"}},
|
|
118
|
+
"key mismatch": {Name: "tenant:other", Authenticated: true, Attributes: map[string]any{"name": "alice", "tenant": "team-a"}},
|
|
119
|
+
} {
|
|
120
|
+
if _, err := principalOf(&a2asrv.ExecutorContext{User: u}); err == nil {
|
|
121
|
+
t.Errorf("%s: a user the interceptor did not vouch for must not run tasks", name)
|
|
122
|
+
}
|
|
123
|
+
}
|
|
124
|
+
p, err := principalOf(&a2asrv.ExecutorContext{User: a2asrv.NewAuthenticatedUser("tenant:team-a", map[string]any{"name": "alice", "tenant": "team-a"})})
|
|
125
|
+
if err != nil || p.Name != "alice" || p.Tenant != "team-a" {
|
|
126
|
+
t.Fatalf("%+v %v", p, err)
|
|
127
|
+
}
|
|
128
|
+
}
|
|
129
|
+
|
|
130
|
+
func TestCancelReturnsOnlyAfterTheWorkerExited(t *testing.T) { // cancel-does-not-wait
|
|
131
|
+
var exited, started = make(chan struct{}), make(chan struct{})
|
|
132
|
+
w := &scriptedWorker{run: func(ctx context.Context, tn Turn, up func(Update)) (Result, error) {
|
|
133
|
+
close(started)
|
|
134
|
+
<-ctx.Done()
|
|
135
|
+
time.Sleep(300 * time.Millisecond) // the worker takes its time to stop
|
|
136
|
+
close(exited)
|
|
137
|
+
return Result{}, ctx.Err()
|
|
138
|
+
}}
|
|
139
|
+
s := startServer(t, serverConfig(), w)
|
|
140
|
+
c := a2aClient(t, "http://"+s.Addr(), tokenA)
|
|
141
|
+
go func() {
|
|
142
|
+
for range c.SendStreamingMessage(context.Background(), &a2a.SendMessageRequest{Message: textMessage("long job")}) {
|
|
143
|
+
}
|
|
144
|
+
}()
|
|
145
|
+
waitChan(t, started, "the worker")
|
|
146
|
+
var id a2a.TaskID
|
|
147
|
+
waitFor(t, func() bool {
|
|
148
|
+
list, err := c.ListTasks(context.Background(), &a2a.ListTasksRequest{})
|
|
149
|
+
if err == nil && len(list.Tasks) == 1 {
|
|
150
|
+
id = list.Tasks[0].ID
|
|
151
|
+
return true
|
|
152
|
+
}
|
|
153
|
+
return false
|
|
154
|
+
}, "the task to be listed")
|
|
155
|
+
if _, err := c.CancelTask(context.Background(), &a2a.CancelTaskRequest{ID: id}); err != nil {
|
|
156
|
+
t.Fatal(err)
|
|
157
|
+
}
|
|
158
|
+
select {
|
|
159
|
+
case <-exited:
|
|
160
|
+
default:
|
|
161
|
+
t.Fatal("CancelTask answered while the worker was still running: a session file would have two writers")
|
|
162
|
+
}
|
|
163
|
+
}
|
|
164
|
+
|
|
165
|
+
func TestSecondStartIsAnError(t *testing.T) { // second-start-allowed
|
|
166
|
+
s := startServer(t, serverConfig(), &scriptedWorker{})
|
|
167
|
+
if err := s.Start(); err == nil {
|
|
168
|
+
t.Fatal("starting a running server twice must fail, not leak a listener")
|
|
169
|
+
}
|
|
170
|
+
}
|
|
171
|
+
|
|
172
|
+
func TestWorkerAsksForTermBeforeKill(t *testing.T) { // no-terminate
|
|
173
|
+
if runtime.GOOS == "windows" {
|
|
174
|
+
t.Skip("POSIX signals")
|
|
175
|
+
}
|
|
176
|
+
w, logPath, _ := newTestWorker(t)
|
|
177
|
+
ctx, cancel := context.WithCancel(context.Background())
|
|
178
|
+
done := make(chan struct{})
|
|
179
|
+
go func() {
|
|
180
|
+
_, _ = w.Run(ctx, Turn{Principal: alice, ContextID: "c", Prompt: "IGNORE_ABORT"}, func(Update) {})
|
|
181
|
+
close(done)
|
|
182
|
+
}()
|
|
183
|
+
waitFakeLog(t, logPath, func(l fakeLog) bool { return l.Started })
|
|
184
|
+
cancel()
|
|
185
|
+
waitChan(t, done, "Run")
|
|
186
|
+
if !readFakeLog(t, logPath).Terminated {
|
|
187
|
+
t.Fatal("after the grace period the worker gets SIGTERM before SIGKILL")
|
|
188
|
+
}
|
|
189
|
+
}
|
|
190
|
+
|
|
191
|
+
func TestWorkerEscalatesToKillWhenTermIsIgnored(t *testing.T) { // no-kill-after-grace
|
|
192
|
+
if runtime.GOOS == "windows" {
|
|
193
|
+
t.Skip("POSIX signals")
|
|
194
|
+
}
|
|
195
|
+
w, logPath, _ := newTestWorker(t)
|
|
196
|
+
ctx, cancel := context.WithCancel(context.Background())
|
|
197
|
+
done := make(chan struct{})
|
|
198
|
+
go func() {
|
|
199
|
+
_, _ = w.Run(ctx, Turn{Principal: alice, ContextID: "c", Prompt: "IGNORE_TERM"}, func(Update) {})
|
|
200
|
+
close(done)
|
|
201
|
+
}()
|
|
202
|
+
l := waitFakeLog(t, logPath, func(l fakeLog) bool { return l.Started })
|
|
203
|
+
cancel()
|
|
204
|
+
select {
|
|
205
|
+
case <-done:
|
|
206
|
+
case <-time.After(20 * time.Second):
|
|
207
|
+
t.Fatal("a worker that ignores abort and SIGTERM must still be killed")
|
|
208
|
+
}
|
|
209
|
+
if processAlive(l.PID) {
|
|
210
|
+
t.Fatal("child survived")
|
|
211
|
+
}
|
|
212
|
+
}
|
|
213
|
+
|
|
214
|
+
func TestWorkerFailsATurnThatNeedsInteractiveInput(t *testing.T) { // ui-request-hangs
|
|
215
|
+
w, _, _ := newTestWorker(t)
|
|
216
|
+
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
217
|
+
defer cancel()
|
|
218
|
+
res, err := w.Run(ctx, Turn{Principal: alice, ContextID: "c", Prompt: "NEEDUI now"}, func(Update) {})
|
|
219
|
+
if err != nil || !strings.Contains(res.Failure, "interactive") {
|
|
220
|
+
t.Fatalf("%+v %v", res, err)
|
|
221
|
+
}
|
|
222
|
+
}
|
|
223
|
+
|
|
224
|
+
func TestWorkerReportsARejectedPromptWithoutTheReason(t *testing.T) { // rejected-prompt-ignored
|
|
225
|
+
w, _, _ := newTestWorker(t)
|
|
226
|
+
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
227
|
+
defer cancel()
|
|
228
|
+
res, err := w.Run(ctx, Turn{Principal: alice, ContextID: "c", Prompt: "REJECT me"}, func(Update) {})
|
|
229
|
+
if err != nil || res.Failure != "PiG rejected the prompt" {
|
|
230
|
+
t.Fatalf("%+v %v", res, err)
|
|
231
|
+
}
|
|
232
|
+
}
|
|
233
|
+
|
|
234
|
+
func TestModelFailureMessageIsExactlyGeneric(t *testing.T) { // model-error-text-copied
|
|
235
|
+
w, _, _ := newTestWorker(t)
|
|
236
|
+
res, err := w.Run(context.Background(), Turn{Principal: alice, ContextID: "c", Prompt: "FAIL now"}, func(Update) {})
|
|
237
|
+
if err != nil || res.Failure != "the model call failed" {
|
|
238
|
+
t.Fatalf("%+v %v", res, err)
|
|
239
|
+
}
|
|
240
|
+
}
|
|
241
|
+
|
|
242
|
+
func TestSendRefusesAnEmptyMessageBeforeAnyRequest(t *testing.T) { // empty-message-sent
|
|
243
|
+
r := newFakeRemote(t)
|
|
244
|
+
_, err := remotes(t, r, nil, nil).Send(context.Background(), SendArgs{Agent: "peer", Message: " "}, nil)
|
|
245
|
+
if err == nil || !strings.Contains(err.Error(), "empty") {
|
|
246
|
+
t.Fatalf("%v", err)
|
|
247
|
+
}
|
|
248
|
+
r.mu.Lock()
|
|
249
|
+
n := len(r.headers)
|
|
250
|
+
r.mu.Unlock()
|
|
251
|
+
if n != 0 {
|
|
252
|
+
t.Fatalf("%d requests for an empty message", n)
|
|
253
|
+
}
|
|
254
|
+
}
|
|
255
|
+
|
|
256
|
+
func TestAuthFailureIsExplainedWithoutTheToken(t *testing.T) { // auth-error-shows-token-context
|
|
257
|
+
r := newFakeRemote(t)
|
|
258
|
+
r.requireAuth = "the-right-one"
|
|
259
|
+
rs := remotes(t, r, func(ra *RemoteAgent) { ra.BearerTokenEnv = "REMOTE_TOKEN" }, map[string]string{"REMOTE_TOKEN": "wrong-secret-token"})
|
|
260
|
+
_, err := rs.Send(context.Background(), SendArgs{Agent: "peer", Message: "hello"}, nil)
|
|
261
|
+
if err == nil || !strings.Contains(err.Error(), "authentication failed") {
|
|
262
|
+
t.Fatalf("the model needs to know the credentials were refused: %v", err)
|
|
263
|
+
}
|
|
264
|
+
}
|
|
265
|
+
|
|
266
|
+
func TestCredentialTransportRefusesAnotherHost(t *testing.T) { // transport-host-not-checked
|
|
267
|
+
hit := false
|
|
268
|
+
other := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { hit = true }))
|
|
269
|
+
defer other.Close()
|
|
270
|
+
// Same scheme as the other server, so only the host check can refuse it.
|
|
271
|
+
tr := headerTransport{h: http.Header{"Authorization": {"Bearer secret-value"}}, host: "configured.example:443", scheme: "http", base: http.DefaultTransport}
|
|
272
|
+
req, _ := http.NewRequest(http.MethodGet, other.URL, nil)
|
|
273
|
+
if _, err := tr.RoundTrip(req); err == nil || hit {
|
|
274
|
+
t.Fatalf("credentials went to %s (hit=%v, err=%v)", other.URL, hit, err)
|
|
275
|
+
}
|
|
276
|
+
}
|
|
277
|
+
|
|
278
|
+
func TestRemoteRedirectsAreNotFollowed(t *testing.T) { // redirects-followed-with-credentials
|
|
279
|
+
rs := NewRemotes(map[string]RemoteAgent{"peer": {URL: "http://127.0.0.1:1", BearerTokenEnv: "T"}}, envFrom(map[string]string{"T": "secret-value"}))
|
|
280
|
+
hc, err := rs.httpClient("peer", rs.cfg["peer"])
|
|
281
|
+
if err != nil {
|
|
282
|
+
t.Fatal(err)
|
|
283
|
+
}
|
|
284
|
+
if got := hc.CheckRedirect(&http.Request{}, nil); got != http.ErrUseLastResponse {
|
|
285
|
+
t.Fatalf("redirects must not be followed: %v", got)
|
|
286
|
+
}
|
|
287
|
+
}
|
|
288
|
+
|
|
289
|
+
func TestSummaryOfAnInputRequiredTaskIsNotTerminal(t *testing.T) { // input-required-called-terminal
|
|
290
|
+
task := &a2a.Task{ID: "t", ContextID: "c", Status: a2a.TaskStatus{State: a2a.TaskStateInputRequired, Message: a2a.NewMessage(a2a.MessageRoleAgent, a2a.NewTextPart("which file?"))}}
|
|
291
|
+
s := summarize(task)
|
|
292
|
+
if s.Terminal || s.State != "input-required" || s.Text != "which file?" {
|
|
293
|
+
t.Fatalf("%+v", s)
|
|
294
|
+
}
|
|
295
|
+
out := formatSummary(s) // unfinished-task-not-flagged
|
|
296
|
+
if !strings.Contains(out, "not finished") || !strings.Contains(out, "which file?") {
|
|
297
|
+
t.Fatalf("%s", out)
|
|
298
|
+
}
|
|
299
|
+
if strings.Contains(formatSummary(TaskSummary{State: "completed", Terminal: true, Text: "done"}), "not finished") {
|
|
300
|
+
t.Fatal("a finished task is not flagged")
|
|
301
|
+
}
|
|
302
|
+
}
|
|
@@ -0,0 +1,388 @@
|
|
|
1
|
+
package a2aext
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"bufio"
|
|
5
|
+
"bytes"
|
|
6
|
+
"context"
|
|
7
|
+
"crypto/sha256"
|
|
8
|
+
"encoding/hex"
|
|
9
|
+
"encoding/json"
|
|
10
|
+
"errors"
|
|
11
|
+
"fmt"
|
|
12
|
+
"io"
|
|
13
|
+
"os"
|
|
14
|
+
"os/exec"
|
|
15
|
+
"path/filepath"
|
|
16
|
+
"strings"
|
|
17
|
+
"time"
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
// Turn is one task's run of PiG.
|
|
21
|
+
type Turn struct {
|
|
22
|
+
Principal Principal
|
|
23
|
+
ContextID string
|
|
24
|
+
TaskID string
|
|
25
|
+
Prompt string
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
// Update is incremental output from a running turn.
|
|
29
|
+
type Update struct {
|
|
30
|
+
Text string // assistant text delta
|
|
31
|
+
Tool string // name of a tool that started
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
// Result is a finished turn. Text is exactly the concatenation of the streamed text updates.
|
|
35
|
+
type Result struct {
|
|
36
|
+
Text string
|
|
37
|
+
// Failure is a short, peer-safe reason the turn failed. Provider error text is never copied into it.
|
|
38
|
+
Failure string
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
// Worker runs turns.
|
|
42
|
+
type Worker interface {
|
|
43
|
+
Run(ctx context.Context, turn Turn, onUpdate func(Update)) (Result, error)
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
// ProcessWorker runs each turn in a `pig --mode rpc` child process. A context is one PiG
|
|
47
|
+
// session file, so a later task in the same context continues the conversation.
|
|
48
|
+
type ProcessWorker struct {
|
|
49
|
+
cfg WorkerConfig
|
|
50
|
+
stateDir string
|
|
51
|
+
getenv func(string) string
|
|
52
|
+
}
|
|
53
|
+
|
|
54
|
+
const maxFrameBytes = 16 << 20
|
|
55
|
+
|
|
56
|
+
// baseEnv is what a worker inherits. Provider keys and everything else must be named in passEnv.
|
|
57
|
+
var baseEnv = []string{
|
|
58
|
+
"PATH", "HOME", "USER", "LANG", "LC_ALL", "TMPDIR", "TERM",
|
|
59
|
+
"PIG_HOME", "PIG_CODING_AGENT_DIR", "PI_CODING_AGENT_DIR",
|
|
60
|
+
"XDG_CONFIG_HOME", "XDG_STATE_HOME", "XDG_CACHE_HOME", "XDG_DATA_HOME",
|
|
61
|
+
"HTTP_PROXY", "HTTPS_PROXY", "NO_PROXY", "http_proxy", "https_proxy", "no_proxy",
|
|
62
|
+
"SSL_CERT_FILE", "SSL_CERT_DIR", "SystemRoot", "USERPROFILE", "APPDATA", "LOCALAPPDATA",
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
// NewProcessWorker builds a ProcessWorker whose per-principal sessions live under stateDir.
|
|
66
|
+
func NewProcessWorker(cfg WorkerConfig, stateDir string, getenv func(string) string) (*ProcessWorker, error) {
|
|
67
|
+
if stateDir == "" {
|
|
68
|
+
return nil, errors.New("a2a: worker needs a state directory")
|
|
69
|
+
}
|
|
70
|
+
if getenv == nil {
|
|
71
|
+
getenv = os.Getenv
|
|
72
|
+
}
|
|
73
|
+
if cfg.Command == "" {
|
|
74
|
+
cfg.Command = getenv("PIG_A2A_PIG")
|
|
75
|
+
}
|
|
76
|
+
if cfg.Command == "" {
|
|
77
|
+
cfg.Command = "pig"
|
|
78
|
+
}
|
|
79
|
+
for _, name := range cfg.PassEnv {
|
|
80
|
+
if !envRE.MatchString(name) {
|
|
81
|
+
return nil, fmt.Errorf("a2a: worker passEnv %q is not an environment variable name", name)
|
|
82
|
+
}
|
|
83
|
+
}
|
|
84
|
+
if cfg.GraceSeconds <= 0 {
|
|
85
|
+
cfg.GraceSeconds = defaultGraceSeconds
|
|
86
|
+
}
|
|
87
|
+
return &ProcessWorker{cfg: cfg, stateDir: stateDir, getenv: getenv}, nil
|
|
88
|
+
}
|
|
89
|
+
|
|
90
|
+
func keyDigest(s string) string {
|
|
91
|
+
sum := sha256.Sum256([]byte(s))
|
|
92
|
+
return hex.EncodeToString(sum[:])
|
|
93
|
+
}
|
|
94
|
+
|
|
95
|
+
// SessionID derives the PiG session id for a principal's context. It is hex, so a hostile
|
|
96
|
+
// context id cannot reach a file name, and it depends on the principal's tenant boundary,
|
|
97
|
+
// so two tenants that pick the same context id never share a session.
|
|
98
|
+
func SessionID(p Principal, contextID string) string {
|
|
99
|
+
return keyDigest("a2a-session\x00" + p.Key() + "\x00" + contextID)[:32]
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
func (w *ProcessWorker) sessionDir(p Principal) string {
|
|
103
|
+
return filepath.Join(w.stateDir, "sessions", keyDigest("a2a-dir\x00" + p.Key())[:16])
|
|
104
|
+
}
|
|
105
|
+
|
|
106
|
+
func (w *ProcessWorker) cwd(p Principal) string {
|
|
107
|
+
if w.cfg.Cwd != "" {
|
|
108
|
+
return w.cfg.Cwd
|
|
109
|
+
}
|
|
110
|
+
// Without an explicit workspace the worker sees an empty private directory, not the operator's files.
|
|
111
|
+
return filepath.Join(w.stateDir, "workspaces", keyDigest("a2a-dir\x00" + p.Key())[:16])
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
func (w *ProcessWorker) args(t Turn) []string {
|
|
115
|
+
args := []string{"--mode", "rpc", "--offline", "--no-extensions", "--no-skills", "--no-context-files", "--no-prompt-templates", "--no-themes"}
|
|
116
|
+
if len(w.cfg.Tools) > 0 {
|
|
117
|
+
args = append(args, "--tools", strings.Join(w.cfg.Tools, ","))
|
|
118
|
+
} else {
|
|
119
|
+
// PiG's file tools are not confined to the working directory, so a worker gets no tools unless named.
|
|
120
|
+
args = append(args, "--no-tools")
|
|
121
|
+
}
|
|
122
|
+
if w.cfg.Provider != "" {
|
|
123
|
+
args = append(args, "--provider", w.cfg.Provider)
|
|
124
|
+
}
|
|
125
|
+
if w.cfg.Model != "" {
|
|
126
|
+
args = append(args, "--model", w.cfg.Model)
|
|
127
|
+
}
|
|
128
|
+
args = append(args, "--session-dir", w.sessionDir(t.Principal), "--session-id", SessionID(t.Principal, t.ContextID))
|
|
129
|
+
return append(args, w.cfg.Args...)
|
|
130
|
+
}
|
|
131
|
+
|
|
132
|
+
func (w *ProcessWorker) env() []string {
|
|
133
|
+
var env []string
|
|
134
|
+
seen := map[string]bool{}
|
|
135
|
+
for _, name := range append(append([]string{}, baseEnv...), w.cfg.PassEnv...) {
|
|
136
|
+
if v := w.getenv(name); v != "" && !seen[name] {
|
|
137
|
+
seen[name] = true
|
|
138
|
+
env = append(env, name+"="+v)
|
|
139
|
+
}
|
|
140
|
+
}
|
|
141
|
+
return append(env, "PIG_A2A_WORKER=1", "PI_SKIP_VERSION_CHECK=1")
|
|
142
|
+
}
|
|
143
|
+
|
|
144
|
+
// wireEvent is the part of PiG's RPC JSONL the worker consumes.
|
|
145
|
+
type wireEvent struct {
|
|
146
|
+
Type string `json:"type"`
|
|
147
|
+
ID string `json:"id"`
|
|
148
|
+
Command string `json:"command"`
|
|
149
|
+
Success bool `json:"success"`
|
|
150
|
+
ToolName string `json:"toolName"`
|
|
151
|
+
AssistantMessageEvent *struct {
|
|
152
|
+
Type string `json:"type"`
|
|
153
|
+
Delta string `json:"delta"`
|
|
154
|
+
} `json:"assistantMessageEvent"`
|
|
155
|
+
Message *struct {
|
|
156
|
+
Role string `json:"role"`
|
|
157
|
+
StopReason string `json:"stopReason"`
|
|
158
|
+
Content json.RawMessage `json:"content"`
|
|
159
|
+
} `json:"message"`
|
|
160
|
+
}
|
|
161
|
+
|
|
162
|
+
type frame struct {
|
|
163
|
+
event wireEvent
|
|
164
|
+
err error
|
|
165
|
+
}
|
|
166
|
+
|
|
167
|
+
func readFrames(r io.Reader, out chan<- frame) {
|
|
168
|
+
defer close(out)
|
|
169
|
+
br := bufio.NewReaderSize(r, 1<<20)
|
|
170
|
+
for {
|
|
171
|
+
var line []byte
|
|
172
|
+
for {
|
|
173
|
+
chunk, err := br.ReadSlice('\n')
|
|
174
|
+
line = append(line, chunk...)
|
|
175
|
+
if len(line) > maxFrameBytes {
|
|
176
|
+
out <- frame{err: fmt.Errorf("PiG frame exceeds %d bytes", maxFrameBytes)}
|
|
177
|
+
return
|
|
178
|
+
}
|
|
179
|
+
if err == bufio.ErrBufferFull {
|
|
180
|
+
continue
|
|
181
|
+
}
|
|
182
|
+
if err != nil {
|
|
183
|
+
if len(bytes.TrimSpace(line)) > 0 {
|
|
184
|
+
out <- frame{err: fmt.Errorf("PiG stream closed inside a record: %w", err)}
|
|
185
|
+
}
|
|
186
|
+
return
|
|
187
|
+
}
|
|
188
|
+
break
|
|
189
|
+
}
|
|
190
|
+
if len(bytes.TrimSpace(line)) == 0 {
|
|
191
|
+
continue
|
|
192
|
+
}
|
|
193
|
+
var ev wireEvent
|
|
194
|
+
if err := json.Unmarshal(line, &ev); err != nil || ev.Type == "" {
|
|
195
|
+
out <- frame{err: errors.New("PiG sent a record that is not an event")}
|
|
196
|
+
return
|
|
197
|
+
}
|
|
198
|
+
out <- frame{event: ev}
|
|
199
|
+
}
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
type tailBuffer struct{ b []byte }
|
|
203
|
+
|
|
204
|
+
func (t *tailBuffer) Write(p []byte) (int, error) {
|
|
205
|
+
t.b = append(t.b, p...)
|
|
206
|
+
if len(t.b) > 4096 {
|
|
207
|
+
t.b = t.b[len(t.b)-4096:]
|
|
208
|
+
}
|
|
209
|
+
return len(p), nil
|
|
210
|
+
}
|
|
211
|
+
|
|
212
|
+
// Run implements Worker.
|
|
213
|
+
func (w *ProcessWorker) Run(ctx context.Context, t Turn, onUpdate func(Update)) (Result, error) {
|
|
214
|
+
if strings.TrimSpace(t.Prompt) == "" {
|
|
215
|
+
return Result{}, errors.New("a2a: empty prompt")
|
|
216
|
+
}
|
|
217
|
+
if err := ctx.Err(); err != nil {
|
|
218
|
+
return Result{}, err
|
|
219
|
+
}
|
|
220
|
+
for _, dir := range []string{w.sessionDir(t.Principal), w.cwd(t.Principal)} {
|
|
221
|
+
if err := os.MkdirAll(dir, 0o700); err != nil {
|
|
222
|
+
return Result{}, fmt.Errorf("a2a: prepare worker directory: %w", err)
|
|
223
|
+
}
|
|
224
|
+
}
|
|
225
|
+
cmd := exec.Command(w.cfg.Command, w.args(t)...)
|
|
226
|
+
cmd.Dir, cmd.Env = w.cwd(t.Principal), w.env()
|
|
227
|
+
configureProcess(cmd)
|
|
228
|
+
stderr := &tailBuffer{}
|
|
229
|
+
cmd.Stderr = stderr
|
|
230
|
+
stdin, err := cmd.StdinPipe()
|
|
231
|
+
if err != nil {
|
|
232
|
+
return Result{}, err
|
|
233
|
+
}
|
|
234
|
+
stdout, err := cmd.StdoutPipe()
|
|
235
|
+
if err != nil {
|
|
236
|
+
return Result{}, err
|
|
237
|
+
}
|
|
238
|
+
if err := cmd.Start(); err != nil {
|
|
239
|
+
return Result{}, fmt.Errorf("a2a: start worker %q: %w", w.cfg.Command, err)
|
|
240
|
+
}
|
|
241
|
+
frames := make(chan frame, 64)
|
|
242
|
+
go readFrames(stdout, frames)
|
|
243
|
+
exited := make(chan error, 1)
|
|
244
|
+
go func() { exited <- cmd.Wait() }()
|
|
245
|
+
grace := time.Duration(w.cfg.GraceSeconds) * time.Second
|
|
246
|
+
|
|
247
|
+
write := func(v map[string]any) error {
|
|
248
|
+
b, _ := json.Marshal(v)
|
|
249
|
+
_, err := stdin.Write(append(b, '\n'))
|
|
250
|
+
return err
|
|
251
|
+
}
|
|
252
|
+
var result Result
|
|
253
|
+
var runErr error
|
|
254
|
+
defer func() {
|
|
255
|
+
_ = stdin.Close()
|
|
256
|
+
select {
|
|
257
|
+
case <-exited:
|
|
258
|
+
case <-time.After(grace):
|
|
259
|
+
_ = terminateTree(cmd.Process)
|
|
260
|
+
select {
|
|
261
|
+
case <-exited:
|
|
262
|
+
case <-time.After(grace):
|
|
263
|
+
_ = killTree(cmd.Process)
|
|
264
|
+
<-exited
|
|
265
|
+
}
|
|
266
|
+
}
|
|
267
|
+
// A tool the worker started can outlive it.
|
|
268
|
+
_ = killTree(cmd.Process)
|
|
269
|
+
go func() {
|
|
270
|
+
for range frames { // let the reader finish
|
|
271
|
+
}
|
|
272
|
+
}()
|
|
273
|
+
}()
|
|
274
|
+
|
|
275
|
+
if err := write(map[string]any{"id": "prompt", "type": "prompt", "message": t.Prompt}); err != nil {
|
|
276
|
+
return Result{}, fmt.Errorf("a2a: write prompt: %w", err)
|
|
277
|
+
}
|
|
278
|
+
|
|
279
|
+
var (
|
|
280
|
+
text strings.Builder
|
|
281
|
+
msgText string // streamed text of the current assistant message
|
|
282
|
+
msgHasText bool
|
|
283
|
+
lastStop string
|
|
284
|
+
accepted bool
|
|
285
|
+
aborting bool
|
|
286
|
+
abortTimer <-chan time.Time
|
|
287
|
+
ctxDone = ctx.Done()
|
|
288
|
+
)
|
|
289
|
+
emit := func(s string) {
|
|
290
|
+
if s == "" {
|
|
291
|
+
return
|
|
292
|
+
}
|
|
293
|
+
if !msgHasText && text.Len() > 0 {
|
|
294
|
+
text.WriteString("\n\n")
|
|
295
|
+
onUpdate(Update{Text: "\n\n"})
|
|
296
|
+
}
|
|
297
|
+
msgHasText = true
|
|
298
|
+
msgText += s
|
|
299
|
+
text.WriteString(s)
|
|
300
|
+
onUpdate(Update{Text: s})
|
|
301
|
+
}
|
|
302
|
+
for {
|
|
303
|
+
select {
|
|
304
|
+
case <-ctxDone:
|
|
305
|
+
ctxDone = nil
|
|
306
|
+
aborting = true
|
|
307
|
+
_ = write(map[string]any{"id": "abort", "type": "abort"})
|
|
308
|
+
abortTimer = time.After(grace)
|
|
309
|
+
case <-abortTimer:
|
|
310
|
+
return Result{}, ctx.Err()
|
|
311
|
+
case f, ok := <-frames:
|
|
312
|
+
if !ok {
|
|
313
|
+
if aborting {
|
|
314
|
+
return Result{}, ctx.Err()
|
|
315
|
+
}
|
|
316
|
+
return Result{}, fmt.Errorf("a2a: worker exited before the turn settled (%s)", strings.TrimSpace(string(stderr.b)))
|
|
317
|
+
}
|
|
318
|
+
if f.err != nil {
|
|
319
|
+
return Result{}, fmt.Errorf("a2a: worker protocol: %w", f.err)
|
|
320
|
+
}
|
|
321
|
+
ev := f.event
|
|
322
|
+
switch ev.Type {
|
|
323
|
+
case "response":
|
|
324
|
+
if ev.ID == "prompt" {
|
|
325
|
+
if !ev.Success {
|
|
326
|
+
return Result{Failure: "PiG rejected the prompt"}, nil
|
|
327
|
+
}
|
|
328
|
+
accepted = true
|
|
329
|
+
}
|
|
330
|
+
case "message_start":
|
|
331
|
+
if ev.Message != nil && ev.Message.Role == "assistant" {
|
|
332
|
+
msgText, msgHasText, lastStop = "", false, ""
|
|
333
|
+
}
|
|
334
|
+
case "message_update":
|
|
335
|
+
if !aborting && ev.AssistantMessageEvent != nil && ev.AssistantMessageEvent.Type == "text_delta" {
|
|
336
|
+
emit(ev.AssistantMessageEvent.Delta)
|
|
337
|
+
}
|
|
338
|
+
case "message_end":
|
|
339
|
+
if ev.Message == nil || ev.Message.Role != "assistant" {
|
|
340
|
+
break
|
|
341
|
+
}
|
|
342
|
+
lastStop = ev.Message.StopReason
|
|
343
|
+
if lastStop == "error" || lastStop == "aborted" || aborting {
|
|
344
|
+
break // a failed attempt may retry; never copy provider text
|
|
345
|
+
}
|
|
346
|
+
var blocks []struct {
|
|
347
|
+
Type string `json:"type"`
|
|
348
|
+
Text string `json:"text"`
|
|
349
|
+
}
|
|
350
|
+
if json.Unmarshal(ev.Message.Content, &blocks) == nil {
|
|
351
|
+
var final strings.Builder
|
|
352
|
+
for _, b := range blocks {
|
|
353
|
+
if b.Type == "text" {
|
|
354
|
+
final.WriteString(b.Text)
|
|
355
|
+
}
|
|
356
|
+
}
|
|
357
|
+
if f := final.String(); len(f) > len(msgText) && strings.HasPrefix(f, msgText) {
|
|
358
|
+
emit(f[len(msgText):])
|
|
359
|
+
}
|
|
360
|
+
}
|
|
361
|
+
case "tool_execution_start":
|
|
362
|
+
if !aborting && ev.ToolName != "" {
|
|
363
|
+
onUpdate(Update{Tool: ev.ToolName})
|
|
364
|
+
}
|
|
365
|
+
case "error":
|
|
366
|
+
result.Failure = "PiG reported an error"
|
|
367
|
+
return result, runErr
|
|
368
|
+
case "extension_ui_request":
|
|
369
|
+
result.Failure = "PiG asked for interactive input, which an A2A task cannot provide"
|
|
370
|
+
return result, runErr
|
|
371
|
+
case "agent_settled":
|
|
372
|
+
if aborting {
|
|
373
|
+
return Result{}, ctx.Err()
|
|
374
|
+
}
|
|
375
|
+
if !accepted {
|
|
376
|
+
return Result{}, errors.New("a2a: worker settled before accepting the prompt")
|
|
377
|
+
}
|
|
378
|
+
switch lastStop {
|
|
379
|
+
case "error":
|
|
380
|
+
return Result{Text: text.String(), Failure: "the model call failed"}, nil
|
|
381
|
+
case "aborted":
|
|
382
|
+
return Result{Text: text.String(), Failure: "the turn was aborted"}, nil
|
|
383
|
+
}
|
|
384
|
+
return Result{Text: text.String()}, nil
|
|
385
|
+
}
|
|
386
|
+
}
|
|
387
|
+
}
|
|
388
|
+
}
|