diff --git a/go/core/internal/database/client_agent_instance_test.go b/go/core/internal/database/client_agent_instance_test.go index 4650f2009..f03938ad4 100644 --- a/go/core/internal/database/client_agent_instance_test.go +++ b/go/core/internal/database/client_agent_instance_test.go @@ -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() diff --git a/go/core/internal/database/client_postgres.go b/go/core/internal/database/client_postgres.go index 31240b144..2903d8e07 100644 --- a/go/core/internal/database/client_postgres.go +++ b/go/core/internal/database/client_postgres.go @@ -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 { @@ -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) diff --git a/go/core/v2/a2agateway/gateway.go b/go/core/v2/a2agateway/gateway.go index 1ff615a64..c6a02610f 100644 --- a/go/core/v2/a2agateway/gateway.go +++ b/go/core/v2/a2agateway/gateway.go @@ -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 { diff --git a/go/core/v2/a2agateway/gateway_test.go b/go/core/v2/a2agateway/gateway_test.go index 991a13a71..9a9de157e 100644 --- a/go/core/v2/a2agateway/gateway_test.go +++ b/go/core/v2/a2agateway/gateway_test.go @@ -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{}