Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 43 additions & 0 deletions go/core/internal/database/client_agent_instance_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,49 @@ func TestConcurrentAgentInstanceMessageReplay(t *testing.T) {
}
}

func TestAgentInstanceReplyArchivesStatusMessageAtomically(t *testing.T) {
db := setupTestDB(t)
ctx := context.Background()
if _, err := db.Exec(ctx, `
INSERT INTO a2a_context (id, namespace, user_id)
VALUES ('instance-1', 'team-a', 'alice');
INSERT INTO agent_instance (id, namespace, user_id, request_id, context_id, state, data)
VALUES ('instance-1', 'team-a', 'alice', 'request-1', 'instance-1', 'READY', '\\x00')
`); err != nil {
t.Fatal(err)
}
client := NewClient(db)
asked := &a2a.Message{ID: "message-1", Role: a2a.MessageRoleUser, TaskID: "task-1", ContextID: "instance-1"}
question := &a2a.Message{ID: "question-1", Role: a2a.MessageRoleAgent, TaskID: "task-1", ContextID: "instance-1"}
parked := &a2a.Task{
ID: "task-1", ContextID: "instance-1", History: []*a2a.Message{asked},
Status: a2a.TaskStatus{State: a2a.TaskStateInputRequired, Message: question},
}
if _, _, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("request-1"), parked); err != nil {
t.Fatal(err)
}

answer := &a2a.Message{ID: "answer-1", Role: a2a.MessageRoleUser, TaskID: "task-1", ContextID: "instance-1"}
resumed := *parked
resumed.History = []*a2a.Message{asked, question, answer}
resumed.Status = a2a.TaskStatus{State: a2a.TaskStateSubmitted}
if err := client.StoreAgentInstanceTaskEvent(ctx, "instance-1", &resumed, answer, nil); err != nil {
t.Fatal(err)
}

got, err := client.GetAgentInstanceTask(ctx, "instance-1", "task-1")
if err != nil {
t.Fatal(err)
}
ids := make([]string, 0, len(got.History))
for _, message := range got.History {
ids = append(ids, message.ID)
}
if strings.Join(ids, ",") != "message-1,question-1,answer-1" {
t.Fatalf("history = %v, want the question between the message it answers and its own answer", ids)
}
}

func TestAgentInstanceCheckpointRetainsRecordedBoundary(t *testing.T) {
db := setupTestDB(t)
ctx := context.Background()
Expand Down
10 changes: 10 additions & 0 deletions go/core/internal/database/client_postgres.go
Original file line number Diff line number Diff line change
Expand Up @@ -854,12 +854,19 @@ func (c *postgresClient) InterruptActiveAgentInstanceTask(ctx context.Context, i
func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instanceID string, task *a2a.Task, event a2a.Event, snapshot *dbpkg.AgentInstanceTaskSnapshot) error {
err := c.withTx(ctx, func(q *dbgen.Queries) error {
var sequence int64
var replacedStatusMessage *a2a.Message
if task != nil {
if row, err := q.GetAgentInstanceTask(ctx, dbgen.GetAgentInstanceTaskParams{ContextID: instanceID, ID: string(task.ID)}); err == nil {
previous, err := unmarshalAgentInstanceTask(row.Data)
if err != nil {
return err
}
// A reply replaces the current status message, so archive both atomically.
if _, ok := event.(*a2a.Message); ok && previous.Status.Message != nil {
message := *previous.Status.Message
message.TaskID, message.ContextID = task.ID, task.ContextID
replacedStatusMessage = &message
}
if len(previous.History) > 0 {
sequence, err = storeAgentInstanceTaskMessages(ctx, q, instanceID, string(task.ID), previous.History)
if err != nil {
Expand All @@ -884,6 +891,9 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan
}
}
messages := agentInstanceTaskEventMessages(task, event)
if replacedStatusMessage != nil {
messages = append([]*a2a.Message{replacedStatusMessage}, messages...)
}
if len(messages) > 0 {
var err error
sequence, err = storeAgentInstanceTaskMessages(ctx, q, instanceID, string(event.TaskInfo().TaskID), messages)
Expand Down
10 changes: 9 additions & 1 deletion go/core/v2/a2agateway/gateway.go
Original file line number Diff line number Diff line change
Expand Up @@ -517,7 +517,15 @@ func (g *Gateway) prepareReply(ctx context.Context, instance *apiv1alpha1.AgentI
}
message.ContextID = stored.ContextID
attempt := *stored
attempt.History = append(append([]*a2atype.Message{}, stored.History...), message)
attempt.History = append([]*a2atype.Message{}, stored.History...)
if question := stored.Status.Message; question != nil {
if question.ID == "" {
return nil, a2atype.NewError(a2atype.ErrInternalError, "stored task status message has no ID")
}
question.TaskID, question.ContextID = stored.ID, stored.ContextID
attempt.History = append(attempt.History, question)
}
attempt.History = append(attempt.History, message)
now := time.Now()
attempt.Status = a2atype.TaskStatus{State: a2atype.TaskStateSubmitted, Timestamp: &now}
if err := g.store.StoreAgentInstanceTaskEvent(ctx, instance.GetId(), &attempt, message, nil); err != nil {
Expand Down
44 changes: 44 additions & 0 deletions go/core/v2/a2agateway/gateway_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -392,6 +392,50 @@ func TestGatewayContinuesInputRequiredTask(t *testing.T) {
}
}

func TestGatewayMovesInputRequiredMessageBeforeReply(t *testing.T) {
question := a2atype.NewMessage(a2atype.MessageRoleAgent, a2atype.NewTextPart("Which database?"))
waiting := &a2atype.Task{
ID: "task-1", ContextID: gatewayTestID,
Status: a2atype.TaskStatus{State: a2atype.TaskStateInputRequired, Message: question},
}
reply := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("PostgreSQL"))
reply.TaskID = waiting.ID
store := &gatewayTestStore{task: waiting}
gateway := &Gateway{store: store}

prepared, err := gateway.prepareReply(t.Context(), gatewayTestInstance(), &a2atype.SendMessageRequest{Message: reply})
if err != nil {
t.Fatal(err)
}
if len(prepared.task.History) != 2 || prepared.task.History[0] != question || prepared.task.History[1] != reply {
t.Fatalf("history = %#v, want question followed by reply", prepared.task.History)
}
if len(store.stored) != 1 || store.stored[0] != reply {
t.Fatalf("stored events = %#v, want one atomic reply update", store.stored)
}
if question.TaskID != waiting.ID || question.ContextID != waiting.ContextID {
t.Fatalf("archived question = task %q context %q, want the task it was asked in", question.TaskID, question.ContextID)
}
}

func TestGatewayRejectsInputRequiredMessageWithoutID(t *testing.T) {
waiting := &a2atype.Task{
ID: "task-1", ContextID: gatewayTestID,
Status: a2atype.TaskStatus{State: a2atype.TaskStateInputRequired, Message: &a2atype.Message{}},
}
store := &gatewayTestStore{task: waiting}
gateway := &Gateway{store: store}
reply := a2atype.NewMessage(a2atype.MessageRoleUser, a2atype.NewTextPart("PostgreSQL"))
reply.TaskID = waiting.ID

if _, err := gateway.prepareReply(t.Context(), gatewayTestInstance(), &a2atype.SendMessageRequest{Message: reply}); err == nil {
t.Fatal("prepareReply() succeeded with an unidentifiable status message")
}
if len(store.stored) != 0 {
t.Fatalf("stored events = %#v, want no partial write", store.stored)
}
}

func TestGatewayClosesRuntimeAfterStreaming(t *testing.T) {
instance := gatewayTestInstance()
runtime := &gatewayTestRuntime{}
Expand Down
Loading