Skip to content

Commit c0fa9d6

Browse files
committed
test: add active AgentInstance interaction e2e
Signed-off-by: Eitan Yarmush <eitan.yarmush@solo.io>
1 parent 89d0655 commit c0fa9d6

1 file changed

Lines changed: 262 additions & 72 deletions

File tree

go/core/test/e2e/interaction_test.go

Lines changed: 262 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,15 @@ package e2e_test
33
import (
44
"context"
55
"embed"
6+
"encoding/json"
67
"net"
8+
"net/http"
9+
"net/http/httptest"
710
"net/url"
811
"os"
912
goruntime "runtime"
1013
"strings"
14+
"sync"
1115
"testing"
1216
"time"
1317

@@ -37,25 +41,195 @@ var interactionMocks embed.FS
3741
// TestAgentInstanceInteraction verifies the complete public interaction path:
3842
// gateway routing, Substrate Actor transport, Go ADK execution, and the model call.
3943
func TestAgentInstanceInteraction(t *testing.T) {
44+
fixture := newInteractionFixture(t, interactionTarget(t), startInteractionMock(t))
45+
_, _, task := fixture.send(t, "What is 2+2?")
46+
if task.Status.State != a2atype.TaskStateCompleted {
47+
t.Fatalf("A2A task state = %s, want COMPLETED", task.Status.State)
48+
}
49+
if text := taskText(task); !strings.Contains(text, "The answer is 4.") {
50+
t.Fatalf("A2A response text = %q, want mock LLM response", text)
51+
}
52+
}
53+
54+
func TestAgentInstanceTaskPersistenceAndIdempotency(t *testing.T) {
55+
fixture := newInteractionFixture(t, interactionTarget(t), startInteractionMock(t))
56+
message, request, task := fixture.send(t, "What is 2+2?")
57+
58+
getRequest, err := pbconv.ToProtoGetTaskRequest(&a2atype.GetTaskRequest{ID: task.ID})
59+
if err != nil {
60+
t.Fatalf("build GetTask request: %v", err)
61+
}
62+
gotProto, err := fixture.client.GetTask(fixture.ctx, getRequest)
63+
if err != nil {
64+
t.Fatalf("get persisted task: %v", err)
65+
}
66+
got, err := pbconv.FromProtoTask(gotProto)
67+
if err != nil {
68+
t.Fatalf("decode persisted task: %v", err)
69+
}
70+
if got.ID != task.ID || got.ContextID != fixture.instanceID || got.Status.State != a2atype.TaskStateCompleted {
71+
t.Fatalf("persisted task = %#v, want completed task %s in context %s", got, task.ID, fixture.instanceID)
72+
}
73+
74+
listRequest, err := pbconv.ToProtoListTasksRequest(&a2atype.ListTasksRequest{ContextID: fixture.instanceID})
75+
if err != nil {
76+
t.Fatalf("build ListTasks request: %v", err)
77+
}
78+
listedProto, err := fixture.client.ListTasks(fixture.ctx, listRequest)
79+
if err != nil {
80+
t.Fatalf("list persisted tasks: %v", err)
81+
}
82+
listed, err := pbconv.FromProtoListTasksResponse(listedProto)
83+
if err != nil {
84+
t.Fatalf("decode listed tasks: %v", err)
85+
}
86+
if listed.TotalSize != 1 || len(listed.Tasks) != 1 || listed.Tasks[0].ID != task.ID || listed.Tasks[0].ContextID != fixture.instanceID {
87+
t.Fatalf("listed tasks = %#v, want only task %s in context %s", listed, task.ID, fixture.instanceID)
88+
}
89+
90+
replayedProto, err := fixture.client.SendMessage(fixture.ctx, request)
91+
if err != nil {
92+
t.Fatalf("replay A2A message: %v", err)
93+
}
94+
replayed, err := pbconv.FromProtoSendMessageResponse(replayedProto)
95+
if err != nil {
96+
t.Fatalf("decode replayed response: %v", err)
97+
}
98+
replayedTask, ok := replayed.(*a2atype.Task)
99+
if !ok || replayedTask.ID != task.ID {
100+
t.Fatalf("replayed response = %#v, want task %s", replayed, task.ID)
101+
}
102+
103+
conflictingMessage := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("What is 3+3?"))
104+
conflictingMessage.ID = message.ID
105+
conflictingRequest, err := pbconv.ToProtoSendMessageRequest(&a2atype.SendMessageRequest{Message: conflictingMessage})
106+
if err != nil {
107+
t.Fatalf("build conflicting A2A request: %v", err)
108+
}
109+
if _, err := fixture.client.SendMessage(fixture.ctx, conflictingRequest); status.Code(err) != codes.InvalidArgument {
110+
t.Fatalf("conflicting message error = %v, want %s", err, codes.InvalidArgument)
111+
}
112+
}
113+
114+
func TestAgentInstanceActiveTask(t *testing.T) {
115+
target := interactionTarget(t)
116+
modelURL, started := startBlockingInteractionMock(t)
117+
fixture := newInteractionFixture(t, target, modelURL)
118+
_, request := newMessageRequest(t, "Wait for cancellation")
119+
stream, err := fixture.client.SendStreamingMessage(fixture.ctx, request)
120+
if err != nil {
121+
t.Fatalf("start streaming A2A message: %v", err)
122+
}
123+
select {
124+
case <-started:
125+
case <-time.After(time.Minute):
126+
t.Fatal("runtime did not call the blocking model")
127+
}
128+
129+
listRequest, err := pbconv.ToProtoListTasksRequest(&a2atype.ListTasksRequest{ContextID: fixture.instanceID})
130+
if err != nil {
131+
t.Fatalf("build ListTasks request: %v", err)
132+
}
133+
listedProto, err := fixture.client.ListTasks(fixture.ctx, listRequest)
134+
if err != nil {
135+
t.Fatalf("list active tasks: %v", err)
136+
}
137+
listed, err := pbconv.FromProtoListTasksResponse(listedProto)
138+
if err != nil {
139+
t.Fatalf("decode active tasks: %v", err)
140+
}
141+
if len(listed.Tasks) != 1 || listed.Tasks[0].Status.State.Terminal() {
142+
t.Fatalf("active tasks = %#v, want one non-terminal task", listed.Tasks)
143+
}
144+
task := listed.Tasks[0]
145+
146+
_, busyRequest := newMessageRequest(t, "Second concurrent request")
147+
if _, err := fixture.client.SendMessage(fixture.ctx, busyRequest); status.Code(err) != codes.FailedPrecondition {
148+
t.Fatalf("concurrent message error = %v, want %s", err, codes.FailedPrecondition)
149+
}
150+
151+
subscribeRequest, err := pbconv.ToProtoSubscribeToTaskRequest(&a2atype.SubscribeToTaskRequest{ID: task.ID})
152+
if err != nil {
153+
t.Fatalf("build SubscribeToTask request: %v", err)
154+
}
155+
subscription, err := fixture.client.SubscribeToTask(fixture.ctx, subscribeRequest)
156+
if err != nil {
157+
t.Fatalf("subscribe to active task: %v", err)
158+
}
159+
firstEventProto, err := subscription.Recv()
160+
if err != nil {
161+
t.Fatalf("receive initial subscribed task event: %v", err)
162+
}
163+
firstEvent, err := pbconv.FromProtoStreamResponse(firstEventProto)
164+
if err != nil {
165+
t.Fatalf("decode initial subscribed task event: %v", err)
166+
}
167+
if firstEvent.TaskInfo().TaskID != task.ID {
168+
t.Fatalf("subscribed task = %s, want %s", firstEvent.TaskInfo().TaskID, task.ID)
169+
}
170+
171+
cancelRequest, err := pbconv.ToProtoCancelTaskRequest(&a2atype.CancelTaskRequest{ID: task.ID})
172+
if err != nil {
173+
t.Fatalf("build CancelTask request: %v", err)
174+
}
175+
canceledProto, err := fixture.client.CancelTask(fixture.ctx, cancelRequest)
176+
if err != nil {
177+
t.Fatalf("cancel active task: %v", err)
178+
}
179+
canceled, err := pbconv.FromProtoTask(canceledProto)
180+
if err != nil {
181+
t.Fatalf("decode canceled task: %v", err)
182+
}
183+
if canceled.ID != task.ID || canceled.Status.State != a2atype.TaskStateCanceled {
184+
t.Fatalf("canceled task = %#v, want task %s in CANCELED", canceled, task.ID)
185+
}
186+
waitForTaskState(t, subscription, a2atype.TaskStateCanceled)
187+
waitForTaskState(t, stream, a2atype.TaskStateCanceled)
188+
getRequest, err := pbconv.ToProtoGetTaskRequest(&a2atype.GetTaskRequest{ID: task.ID})
189+
if err != nil {
190+
t.Fatalf("build GetTask request: %v", err)
191+
}
192+
persistedProto, err := fixture.client.GetTask(fixture.ctx, getRequest)
193+
if err != nil {
194+
t.Fatalf("get canceled task: %v", err)
195+
}
196+
persisted, err := pbconv.FromProtoTask(persistedProto)
197+
if err != nil {
198+
t.Fatalf("decode canceled task: %v", err)
199+
}
200+
if persisted.Status.State != a2atype.TaskStateCanceled {
201+
t.Fatalf("persisted task state = %s, want CANCELED", persisted.Status.State)
202+
}
203+
}
204+
205+
type interactionFixture struct {
206+
ctx context.Context
207+
client a2apb.A2AServiceClient
208+
instanceID string
209+
}
210+
211+
func interactionTarget(t *testing.T) string {
212+
t.Helper()
40213
target := os.Getenv("KAGENT_E2E_GRPC_TARGET")
41214
if target == "" {
42215
target = os.Getenv("KAGENT_GRPC_URL")
43216
}
44217
if target == "" {
45218
t.Skip("KAGENT_E2E_GRPC_TARGET is not set")
46219
}
220+
return target
221+
}
47222

48-
modelURL := startInteractionMock(t)
223+
func newInteractionFixture(t *testing.T, target, modelURL string) *interactionFixture {
224+
t.Helper()
49225
templateName := createInteractionTemplate(t, modelURL)
50-
51226
conn, err := grpc.NewClient(target, grpc.WithTransportCredentials(insecure.NewCredentials()))
52227
if err != nil {
53228
t.Fatalf("connect to kagent gRPC API: %v", err)
54229
}
55230
t.Cleanup(func() { _ = conn.Close() })
56-
57231
ctx, cancel := context.WithTimeout(metadata.AppendToOutgoingContext(t.Context(), "x-user-id", "e2e"), 4*time.Minute)
58-
defer cancel()
232+
t.Cleanup(cancel)
59233
instances := apiv1alpha1.NewAgentInstanceServiceClient(conn)
60234
created, err := instances.CreateAgentInstance(ctx, &apiv1alpha1.CreateAgentInstanceRequest{
61235
Namespace: "kagent", AgentTemplate: templateName, Harness: "kagent", RequestId: uuid.NewString(),
@@ -77,20 +251,20 @@ func TestAgentInstanceInteraction(t *testing.T) {
77251
if instance.GetState() != apiv1alpha1.AgentInstanceState_AGENT_INSTANCE_STATE_READY {
78252
t.Fatalf("created AgentInstance state = %s, want READY", instance.GetState())
79253
}
80-
81-
message := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("What is 2+2?"))
82-
request, err := pbconv.ToProtoSendMessageRequest(&a2atype.SendMessageRequest{
83-
Message: message,
84-
})
85-
if err != nil {
86-
t.Fatalf("build A2A request: %v", err)
254+
return &interactionFixture{
255+
ctx: metadata.AppendToOutgoingContext(ctx,
256+
"x-kagent-agent-instance-namespace", "kagent",
257+
"x-kagent-agent-instance-id", instance.GetId(),
258+
),
259+
client: a2apb.NewA2AServiceClient(conn),
260+
instanceID: instance.GetId(),
87261
}
88-
interactionCtx := metadata.AppendToOutgoingContext(ctx,
89-
"x-kagent-agent-instance-namespace", "kagent",
90-
"x-kagent-agent-instance-id", instance.GetId(),
91-
)
92-
a2aClient := a2apb.NewA2AServiceClient(conn)
93-
response, err := a2aClient.SendMessage(interactionCtx, request)
262+
}
263+
264+
func (f *interactionFixture) send(t *testing.T, text string) (*a2atype.Message, *a2apb.SendMessageRequest, *a2atype.Task) {
265+
t.Helper()
266+
message, request := newMessageRequest(t, text)
267+
response, err := f.client.SendMessage(f.ctx, request)
94268
if err != nil {
95269
t.Fatalf("send A2A message: %v", err)
96270
}
@@ -102,66 +276,44 @@ func TestAgentInstanceInteraction(t *testing.T) {
102276
if !ok {
103277
t.Fatalf("A2A response = %T, want Task", result)
104278
}
105-
if task.Status.State != a2atype.TaskStateCompleted {
106-
t.Fatalf("A2A task state = %s, want COMPLETED", task.Status.State)
107-
}
108-
if text := taskText(task); !strings.Contains(text, "The answer is 4.") {
109-
t.Fatalf("A2A response text = %q, want mock LLM response", text)
110-
}
111-
112-
getRequest, err := pbconv.ToProtoGetTaskRequest(&a2atype.GetTaskRequest{ID: task.ID})
113-
if err != nil {
114-
t.Fatalf("build GetTask request: %v", err)
115-
}
116-
gotProto, err := a2aClient.GetTask(interactionCtx, getRequest)
117-
if err != nil {
118-
t.Fatalf("get persisted task: %v", err)
119-
}
120-
got, err := pbconv.FromProtoTask(gotProto)
121-
if err != nil {
122-
t.Fatalf("decode persisted task: %v", err)
123-
}
124-
if got.ID != task.ID || got.ContextID != instance.GetId() || got.Status.State != a2atype.TaskStateCompleted {
125-
t.Fatalf("persisted task = %#v, want completed task %s in context %s", got, task.ID, instance.GetId())
126-
}
279+
return message, request, task
280+
}
127281

128-
listRequest, err := pbconv.ToProtoListTasksRequest(&a2atype.ListTasksRequest{ContextID: instance.GetId()})
129-
if err != nil {
130-
t.Fatalf("build ListTasks request: %v", err)
131-
}
132-
listedProto, err := a2aClient.ListTasks(interactionCtx, listRequest)
133-
if err != nil {
134-
t.Fatalf("list persisted tasks: %v", err)
135-
}
136-
listed, err := pbconv.FromProtoListTasksResponse(listedProto)
282+
func newMessageRequest(t *testing.T, text string) (*a2atype.Message, *a2apb.SendMessageRequest) {
283+
t.Helper()
284+
message := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart(text))
285+
request, err := pbconv.ToProtoSendMessageRequest(&a2atype.SendMessageRequest{Message: message})
137286
if err != nil {
138-
t.Fatalf("decode listed tasks: %v", err)
139-
}
140-
if listed.TotalSize != 1 || len(listed.Tasks) != 1 || listed.Tasks[0].ID != task.ID || listed.Tasks[0].ContextID != instance.GetId() {
141-
t.Fatalf("listed tasks = %#v, want only task %s in context %s", listed, task.ID, instance.GetId())
287+
t.Fatalf("build A2A request: %v", err)
142288
}
289+
return message, request
290+
}
143291

144-
replayedProto, err := a2aClient.SendMessage(interactionCtx, request)
145-
if err != nil {
146-
t.Fatalf("replay A2A message: %v", err)
147-
}
148-
replayed, err := pbconv.FromProtoSendMessageResponse(replayedProto)
149-
if err != nil {
150-
t.Fatalf("decode replayed response: %v", err)
151-
}
152-
replayedTask, ok := replayed.(*a2atype.Task)
153-
if !ok || replayedTask.ID != task.ID {
154-
t.Fatalf("replayed response = %#v, want task %s", replayed, task.ID)
155-
}
292+
type streamReceiver interface {
293+
Recv() (*a2apb.StreamResponse, error)
294+
}
156295

157-
conflictingMessage := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("What is 3+3?"))
158-
conflictingMessage.ID = message.ID
159-
conflictingRequest, err := pbconv.ToProtoSendMessageRequest(&a2atype.SendMessageRequest{Message: conflictingMessage})
160-
if err != nil {
161-
t.Fatalf("build conflicting A2A request: %v", err)
162-
}
163-
if _, err := a2aClient.SendMessage(interactionCtx, conflictingRequest); status.Code(err) != codes.InvalidArgument {
164-
t.Fatalf("conflicting message error = %v, want %s", err, codes.InvalidArgument)
296+
func waitForTaskState(t *testing.T, stream streamReceiver, want a2atype.TaskState) {
297+
t.Helper()
298+
for {
299+
response, err := stream.Recv()
300+
if err != nil {
301+
t.Fatalf("receive task stream: %v", err)
302+
}
303+
event, err := pbconv.FromProtoStreamResponse(response)
304+
if err != nil {
305+
t.Fatalf("decode task stream: %v", err)
306+
}
307+
switch event := event.(type) {
308+
case *a2atype.Task:
309+
if event.Status.State == want {
310+
return
311+
}
312+
case *a2atype.TaskStatusUpdateEvent:
313+
if event.Status.State == want {
314+
return
315+
}
316+
}
165317
}
166318
}
167319

@@ -181,7 +333,45 @@ func startInteractionMock(t *testing.T) string {
181333
t.Errorf("stop mock LLM: %v", err)
182334
}
183335
})
336+
return reachableModelURL(t, baseURL)
337+
}
184338

339+
func startBlockingInteractionMock(t *testing.T) (string, <-chan struct{}) {
340+
t.Helper()
341+
cfg, err := mockllm.LoadConfigFromFile("mocks/invoke_golang_adk_agent.json", interactionMocks)
342+
if err != nil {
343+
t.Fatalf("load mock LLM response: %v", err)
344+
}
345+
started := make(chan struct{})
346+
release := make(chan struct{})
347+
var startedOnce, releaseOnce sync.Once
348+
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
349+
startedOnce.Do(func() { close(started) })
350+
select {
351+
case <-r.Context().Done():
352+
return
353+
case <-release:
354+
}
355+
w.Header().Set("Content-Type", "application/json")
356+
if err := json.NewEncoder(w).Encode(cfg.OpenAI[0].Response); err != nil {
357+
t.Errorf("write mock LLM response: %v", err)
358+
}
359+
}))
360+
_ = server.Listener.Close()
361+
server.Listener, err = net.Listen("tcp", "0.0.0.0:0")
362+
if err != nil {
363+
t.Fatalf("listen for blocking mock LLM: %v", err)
364+
}
365+
server.Start()
366+
t.Cleanup(func() {
367+
releaseOnce.Do(func() { close(release) })
368+
server.Close()
369+
})
370+
return reachableModelURL(t, server.URL), started
371+
}
372+
373+
func reachableModelURL(t *testing.T, baseURL string) string {
374+
t.Helper()
185375
parsed, err := url.Parse(baseURL)
186376
if err != nil {
187377
t.Fatalf("parse mock LLM URL: %v", err)

0 commit comments

Comments
 (0)