diff --git a/go/api/database/models.go b/go/api/database/models.go index 9843f4707..147405a6d 100644 --- a/go/api/database/models.go +++ b/go/api/database/models.go @@ -5,6 +5,7 @@ import ( "time" a2a "github.com/a2aproject/a2a-go/v2/a2a" + "github.com/google/uuid" "github.com/kagent-dev/kagent/go/api/adk" "github.com/kagent-dev/kagent/go/api/v1alpha3" "github.com/pgvector/pgvector-go" @@ -289,9 +290,9 @@ type AgentInstanceQuery struct { } type AgentInstanceShare struct { - ID string + ID uuid.UUID Namespace string - InstanceID string + InstanceID uuid.UUID Permission string TokenHash []byte CreatedAt time.Time @@ -314,10 +315,10 @@ type AgentInstanceTaskSnapshot struct { } type AgentInstanceCheckpoint struct { - ID string + ID uuid.UUID Namespace string - SourceInstanceID string - SourceContextID string + SourceInstanceID uuid.UUID + SourceContextID uuid.UUID UserID string RequestID string HeadTaskID string diff --git a/go/core/internal/database/client_agent_instance_test.go b/go/core/internal/database/client_agent_instance_test.go index f03938ad4..c91a59e76 100644 --- a/go/core/internal/database/client_agent_instance_test.go +++ b/go/core/internal/database/client_agent_instance_test.go @@ -9,6 +9,7 @@ import ( "time" a2a "github.com/a2aproject/a2a-go/v2/a2a" + "github.com/google/uuid" dbpkg "github.com/kagent-dev/kagent/go/api/database" apiv1alpha1 "github.com/kagent-dev/kagent/go/api/gen/kagent/api/v1alpha1" dbgen "github.com/kagent-dev/kagent/go/core/internal/database/gen" @@ -17,7 +18,7 @@ import ( func TestToAgentInstanceUsesIndexedLifecycleColumns(t *testing.T) { data, err := proto.Marshal(&apiv1alpha1.AgentInstance{ - Id: "instance-1", + Id: "11111111-1111-4111-8111-111111111111", State: apiv1alpha1.AgentInstanceState_AGENT_INSTANCE_STATE_READY, Operation: apiv1alpha1.AgentInstanceOperation_AGENT_INSTANCE_OPERATION_UNSPECIFIED, }) @@ -26,7 +27,7 @@ func TestToAgentInstanceUsesIndexedLifecycleColumns(t *testing.T) { } instance, err := toAgentInstance(dbgen.AgentInstance{ - ID: "instance-1", Data: data, State: "SUSPENDED", Operation: "RESUME", Name: "Renamed later", + ID: uuid.MustParse("11111111-1111-4111-8111-111111111111"), Data: data, State: "SUSPENDED", Operation: "RESUME", Name: "Renamed later", }) if err != nil { t.Fatal(err) @@ -53,7 +54,7 @@ func TestToAgentInstanceLeavesAnEmptyNameEmpty(t *testing.T) { if err != nil { t.Fatal(err) } - instance, err := toAgentInstance(dbgen.AgentInstance{ID: "instance-1", Data: data, State: "READY", Operation: "NONE"}) + instance, err := toAgentInstance(dbgen.AgentInstance{ID: uuid.MustParse("11111111-1111-4111-8111-111111111111"), Data: data, State: "READY", Operation: "NONE"}) if err != nil { t.Fatal(err) } @@ -67,30 +68,30 @@ func TestAgentInstanceTasksAreDurableAndExclusive(t *testing.T) { ctx := context.Background() if _, err := db.Exec(ctx, ` INSERT INTO a2a_context (id, namespace, user_id) - VALUES ('instance-1', 'team-a', 'alice'); + VALUES ('11111111-1111-4111-8111-111111111111', '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') + VALUES ('11111111-1111-4111-8111-111111111111', 'team-a', 'alice', 'request-1', '11111111-1111-4111-8111-111111111111', 'READY', '\x00') `); err != nil { t.Fatal(err) } client := NewClient(db) now := time.Now() first := &a2a.Task{ - ID: "task-1", ContextID: "instance-1", + ID: "task-1", ContextID: "11111111-1111-4111-8111-111111111111", Status: a2a.TaskStatus{State: a2a.TaskStateSubmitted, Timestamp: &now}, History: []*a2a.Message{{ID: "message-1", Role: a2a.MessageRoleUser}}, } - stored, created, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("request-1"), first) + stored, created, err := client.CreateAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", []byte("request-1"), first) if err != nil || !created || stored.ID != first.ID { t.Fatalf("CreateAgentInstanceTask() = %#v, created %v, error %v", stored, created, err) } - replayed, created, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("request-1"), - &a2a.Task{ID: "ignored", ContextID: "instance-1", Status: a2a.TaskStatus{State: a2a.TaskStateSubmitted}, History: first.History}) + replayed, created, err := client.CreateAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", []byte("request-1"), + &a2a.Task{ID: "ignored", ContextID: "11111111-1111-4111-8111-111111111111", Status: a2a.TaskStatus{State: a2a.TaskStateSubmitted}, History: first.History}) if err != nil || created || replayed.ID != first.ID { t.Fatalf("replayed CreateAgentInstanceTask() = %#v, created %v, error %v", replayed, created, err) } - if _, _, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("different"), first); !errors.Is(err, dbpkg.ErrIdempotencyConflict) { + if _, _, err := client.CreateAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", []byte("different"), first); !errors.Is(err, dbpkg.ErrIdempotencyConflict) { t.Fatalf("conflicting message error = %v", err) } if events := countRows(t, db, "SELECT COUNT(*) FROM agent_instance_task_event"); events != 1 { @@ -100,33 +101,33 @@ func TestAgentInstanceTasksAreDurableAndExclusive(t *testing.T) { if err := db.QueryRow(ctx, "SELECT task_id FROM agent_instance_task_event").Scan(&eventTaskID); err != nil || eventTaskID != string(first.ID) { t.Fatalf("initial event task ID = %q, want %q: %v", eventTaskID, first.ID, err) } - got, err := client.GetAgentInstanceTask(ctx, "instance-1", "task-1") + got, err := client.GetAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "task-1") if err != nil || got.ID != first.ID || got.Status.State != first.Status.State || len(got.History) != 1 { t.Fatalf("GetAgentInstanceTask() = %#v, %v", got, err) } var projectionData []byte - if err := db.QueryRow(ctx, `SELECT data FROM agent_instance_task WHERE context_id = 'instance-1' AND id = 'task-1'`).Scan(&projectionData); err != nil { + if err := db.QueryRow(ctx, `SELECT data FROM agent_instance_task WHERE context_id = '11111111-1111-4111-8111-111111111111' AND id = 'task-1'`).Scan(&projectionData); err != nil { t.Fatal(err) } projection, err := unmarshalAgentInstanceTask(projectionData) if err != nil || len(projection.History) != 0 { t.Fatalf("stored task projection history = %#v, error %v", projection.History, err) } - second := &a2a.Task{ID: "task-2", ContextID: "instance-1", Status: a2a.TaskStatus{State: a2a.TaskStateSubmitted}} - if err := client.StoreAgentInstanceTaskEvent(ctx, "instance-1", second, second, nil); !errors.Is(err, dbpkg.ErrAgentInstanceTaskConflict) { + second := &a2a.Task{ID: "task-2", ContextID: "11111111-1111-4111-8111-111111111111", Status: a2a.TaskStatus{State: a2a.TaskStateSubmitted}} + if err := client.StoreAgentInstanceTaskEvent(ctx, "11111111-1111-4111-8111-111111111111", second, second, nil); !errors.Is(err, dbpkg.ErrAgentInstanceTaskConflict) { t.Fatalf("second active task error = %v", err) } first.History = append(first.History, a2a.NewMessageForTask(a2a.MessageRoleAgent, first, a2a.NewTextPart("done"))) first.Status.State = a2a.TaskStateCompleted snapshot := &dbpkg.AgentInstanceTaskSnapshot{Atespace: "team-a", Name: "snapshot-1", UID: "snapshot-uid"} - if err := client.StoreAgentInstanceTaskEvent(ctx, "instance-1", first, first, snapshot); err != nil { + if err := client.StoreAgentInstanceTaskEvent(ctx, "11111111-1111-4111-8111-111111111111", first, first, snapshot); err != nil { t.Fatal(err) } var snapshotAtespace, snapshotName, snapshotUID string var historySequence, latestSequence int64 if err := db.QueryRow(ctx, ` SELECT snapshot_atespace, snapshot_name, snapshot_uid, history_sequence - FROM agent_instance_task WHERE context_id = 'instance-1' AND id = 'task-1' + FROM agent_instance_task WHERE context_id = '11111111-1111-4111-8111-111111111111' AND id = 'task-1' `).Scan(&snapshotAtespace, &snapshotName, &snapshotUID, &historySequence); err != nil { t.Fatal(err) } @@ -136,22 +137,22 @@ func TestAgentInstanceTasksAreDurableAndExclusive(t *testing.T) { if snapshotAtespace != snapshot.Atespace || snapshotName != snapshot.Name || snapshotUID != snapshot.UID || historySequence != latestSequence { t.Fatalf("stored boundary = %s/%s uid %s sequence %d", snapshotAtespace, snapshotName, snapshotUID, historySequence) } - got, err = client.GetAgentInstanceTask(ctx, "instance-1", "task-1") + got, err = client.GetAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "task-1") if err != nil || len(got.History) != 2 || got.History[1].Role != a2a.MessageRoleAgent { t.Fatalf("reconstructed task history = %#v, error %v", got, err) } - if err := client.StoreAgentInstanceTaskEvent(ctx, "instance-1", second, second, nil); err != nil { + if err := client.StoreAgentInstanceTaskEvent(ctx, "11111111-1111-4111-8111-111111111111", second, second, nil); err != nil { t.Fatal(err) } if events := countRows(t, db, "SELECT COUNT(*) FROM agent_instance_task_event"); events != 4 { t.Fatalf("event count = %d, want 4", events) } - tasks, total, err := client.ListAgentInstanceTasks(ctx, "instance-1", "", a2a.TaskStateUnspecified, nil, 1) + tasks, total, err := client.ListAgentInstanceTasks(ctx, "11111111-1111-4111-8111-111111111111", "", a2a.TaskStateUnspecified, nil, 1) if err != nil || total != 2 || len(tasks) != 1 || tasks[0].ID != first.ID { t.Fatalf("first page = %#v, total %d, error %v", tasks, total, err) } - tasks, total, err = client.ListAgentInstanceTasks(ctx, "instance-1", string(first.ID), a2a.TaskStateSubmitted, nil, 2) + tasks, total, err = client.ListAgentInstanceTasks(ctx, "11111111-1111-4111-8111-111111111111", string(first.ID), a2a.TaskStateSubmitted, nil, 2) if err != nil || total != 1 || len(tasks) != 1 || tasks[0].ID != second.ID { t.Fatalf("filtered page = %#v, total %d, error %v", tasks, total, err) } @@ -162,9 +163,9 @@ func TestConcurrentAgentInstanceMessageReplay(t *testing.T) { ctx := context.Background() if _, err := db.Exec(ctx, ` INSERT INTO a2a_context (id, namespace, user_id) - VALUES ('instance-1', 'team-a', 'alice'); + VALUES ('11111111-1111-4111-8111-111111111111', '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') + VALUES ('11111111-1111-4111-8111-111111111111', 'team-a', 'alice', 'request-1', '11111111-1111-4111-8111-111111111111', 'READY', '\x00') `); err != nil { t.Fatal(err) } @@ -179,9 +180,9 @@ func TestConcurrentAgentInstanceMessageReplay(t *testing.T) { for _, taskID := range []a2a.TaskID{"task-1", "task-2"} { go func() { <-start - message := &a2a.Message{ID: "message-1", Role: a2a.MessageRoleUser, TaskID: taskID, ContextID: "instance-1"} + message := &a2a.Message{ID: "message-1", Role: a2a.MessageRoleUser, TaskID: taskID, ContextID: "11111111-1111-4111-8111-111111111111"} task := a2a.NewSubmittedTask(message, message) - stored, created, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("request-1"), task) + stored, created, err := client.CreateAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", []byte("request-1"), task) results <- result{stored, created, err} }() } @@ -211,34 +212,35 @@ func TestConcurrentAgentInstanceMessageReplay(t *testing.T) { func TestAgentInstanceReplyArchivesStatusMessageAtomically(t *testing.T) { db := setupTestDB(t) ctx := context.Background() + instanceID := "11111111-1111-4111-8111-111111111111" if _, err := db.Exec(ctx, ` INSERT INTO a2a_context (id, namespace, user_id) - VALUES ('instance-1', 'team-a', 'alice'); + VALUES ('11111111-1111-4111-8111-111111111111', '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') + VALUES ('11111111-1111-4111-8111-111111111111', 'team-a', 'alice', 'request-1', '11111111-1111-4111-8111-111111111111', '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"} + asked := &a2a.Message{ID: "message-1", Role: a2a.MessageRoleUser, TaskID: "task-1", ContextID: instanceID} + question := &a2a.Message{ID: "question-1", Role: a2a.MessageRoleAgent, TaskID: "task-1", ContextID: instanceID} parked := &a2a.Task{ - ID: "task-1", ContextID: "instance-1", History: []*a2a.Message{asked}, + ID: "task-1", ContextID: instanceID, 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 { + if _, _, err := client.CreateAgentInstanceTask(ctx, instanceID, []byte("request-1"), parked); err != nil { t.Fatal(err) } - answer := &a2a.Message{ID: "answer-1", Role: a2a.MessageRoleUser, TaskID: "task-1", ContextID: "instance-1"} + answer := &a2a.Message{ID: "answer-1", Role: a2a.MessageRoleUser, TaskID: "task-1", ContextID: instanceID} 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 { + if err := client.StoreAgentInstanceTaskEvent(ctx, instanceID, &resumed, answer, nil); err != nil { t.Fatal(err) } - got, err := client.GetAgentInstanceTask(ctx, "instance-1", "task-1") + got, err := client.GetAgentInstanceTask(ctx, instanceID, "task-1") if err != nil { t.Fatal(err) } @@ -254,8 +256,9 @@ func TestAgentInstanceReplyArchivesStatusMessageAtomically(t *testing.T) { func TestAgentInstanceCheckpointRetainsRecordedBoundary(t *testing.T) { db := setupTestDB(t) ctx := context.Background() + instanceID := "11111111-1111-4111-8111-111111111111" instance := &apiv1alpha1.AgentInstance{ - Id: "instance-1", Namespace: "team-a", Creator: "alice", + Id: instanceID, Namespace: "team-a", Creator: "alice", State: apiv1alpha1.AgentInstanceState_AGENT_INSTANCE_STATE_READY, } instanceData, err := proto.Marshal(instance) @@ -264,29 +267,29 @@ func TestAgentInstanceCheckpointRetainsRecordedBoundary(t *testing.T) { } if _, err := db.Exec(ctx, ` INSERT INTO a2a_context (id, namespace, user_id) - VALUES ('instance-1', 'team-a', 'alice') - `); err != nil { + VALUES ($1, 'team-a', 'alice') + `, instanceID); err != nil { t.Fatal(err) } if _, err := db.Exec(ctx, ` INSERT INTO agent_instance (id, namespace, user_id, request_id, context_id, state, data) - VALUES ('instance-1', 'team-a', 'alice', 'instance-request', 'instance-1', 'READY', $1) - `, instanceData); err != nil { + VALUES ($1, 'team-a', 'alice', 'instance-request', $1, 'READY', $2) + `, instanceID, instanceData); err != nil { t.Fatal(err) } client := NewClient(db) task := newAgentInstanceTask("task-1", "message-1") - if _, _, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("message-request"), task); err != nil { + if _, _, err := client.CreateAgentInstanceTask(ctx, instanceID, []byte("message-request"), task); err != nil { t.Fatal(err) } task.Status.State = a2a.TaskStateCompleted - if err := client.StoreAgentInstanceTaskEvent(ctx, "instance-1", task, task, + if err := client.StoreAgentInstanceTaskEvent(ctx, instanceID, task, task, &dbpkg.AgentInstanceTaskSnapshot{Atespace: "team-a", Name: "snapshot-1", UID: "snapshot-uid", ContentScope: "DATA"}); err != nil { t.Fatal(err) } checkpoint, err := client.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: "checkpoint-1", Namespace: "team-a", SourceInstanceID: "instance-1", UserID: "alice", + ID: uuid.MustParse("22222222-2222-4222-8222-222222222222"), Namespace: "team-a", SourceInstanceID: uuid.MustParse(instanceID), UserID: "alice", RequestID: "checkpoint-request", }) if err != nil { @@ -296,7 +299,7 @@ func TestAgentInstanceCheckpointRetainsRecordedBoundary(t *testing.T) { checkpoint.SnapshotContentScope != "DATA" || checkpoint.HistorySequence == 0 { t.Fatalf("checkpoint boundary = %+v", checkpoint) } - if _, _, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("blocked-request"), newAgentInstanceTask("task-2", "message-2")); !errors.Is(err, dbpkg.ErrAgentInstanceTaskConflict) { + if _, _, err := client.CreateAgentInstanceTask(ctx, instanceID, []byte("blocked-request"), newAgentInstanceTask("task-2", "message-2")); !errors.Is(err, dbpkg.ErrAgentInstanceTaskConflict) { t.Fatalf("CreateAgentInstanceTask() during checkpoint = %v, want %v", err, dbpkg.ErrAgentInstanceTaskConflict) } suspending := proto.Clone(instance).(*apiv1alpha1.AgentInstance) @@ -306,48 +309,48 @@ func TestAgentInstanceCheckpointRetainsRecordedBoundary(t *testing.T) { t.Fatalf("lifecycle transition during checkpoint = %+v, error %v", current, err) } replayed, err := client.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: "ignored", Namespace: "team-a", SourceInstanceID: "instance-1", UserID: "alice", + ID: uuid.MustParse("33333333-3333-4333-8333-333333333333"), Namespace: "team-a", SourceInstanceID: uuid.MustParse(instanceID), UserID: "alice", RequestID: "checkpoint-request", }) if err != nil || replayed.ID != checkpoint.ID { t.Fatalf("replayed checkpoint = %+v, error %v", replayed, err) } - ready, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID, "tag-uid", "") + ready, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID.String(), "tag-uid", "") if err != nil || ready.State != "READY" || ready.TagUID != "tag-uid" { t.Fatalf("ready checkpoint = %+v, error %v", ready, err) } - if replayed, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID, "tag-uid", ""); err != nil || replayed.State != "READY" { + if replayed, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID.String(), "tag-uid", ""); err != nil || replayed.State != "READY" { t.Fatalf("replayed ready checkpoint = %+v, error %v", replayed, err) } failed, err := client.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: "checkpoint-2", Namespace: "team-a", SourceInstanceID: "instance-1", UserID: "alice", + ID: uuid.MustParse("44444444-4444-4444-8444-444444444444"), Namespace: "team-a", SourceInstanceID: uuid.MustParse(instanceID), UserID: "alice", RequestID: "failed-checkpoint-request", }) if err != nil { t.Fatal(err) } - failed, err = client.FinalizeAgentInstanceCheckpoint(ctx, failed.ID, "", "tag creation failed") + failed, err = client.FinalizeAgentInstanceCheckpoint(ctx, failed.ID.String(), "", "tag creation failed") if err != nil || failed.State != "FAILED" || failed.Failure != "tag creation failed" { t.Fatalf("failed checkpoint = %+v, error %v", failed, err) } - if err := client.DeleteAgentInstance(ctx, "instance-1"); err != nil { + if err := client.DeleteAgentInstance(ctx, instanceID); err != nil { t.Fatal(err) } replayed, err = client.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: "ignored-again", Namespace: "team-a", SourceInstanceID: "instance-1", UserID: "alice", + ID: uuid.MustParse("55555555-5555-4555-8555-555555555555"), Namespace: "team-a", SourceInstanceID: uuid.MustParse(instanceID), UserID: "alice", RequestID: "checkpoint-request", }) if err != nil || replayed.ID != checkpoint.ID { t.Fatalf("checkpoint replay after source deletion = %+v, error %v", replayed, err) } - listed, err := client.ListAgentInstanceCheckpoints(ctx, "team-a", "instance-1", "alice", "", 10) + listed, err := client.ListAgentInstanceCheckpoints(ctx, "team-a", instanceID, "alice", "", 10) if err != nil || len(listed) != 1 || listed[0].ID != checkpoint.ID { t.Fatalf("listed checkpoints = %+v, error %v", listed, err) } - if _, err := client.BeginDeleteAgentInstanceCheckpoint(ctx, "team-a", checkpoint.ID, "alice"); err != nil { + if _, err := client.BeginDeleteAgentInstanceCheckpoint(ctx, "team-a", checkpoint.ID.String(), "alice"); err != nil { t.Fatal(err) } - if err := client.DeleteAgentInstanceCheckpoint(ctx, "team-a", checkpoint.ID, "alice"); err != nil { + if err := client.DeleteAgentInstanceCheckpoint(ctx, "team-a", checkpoint.ID.String(), "alice"); err != nil { t.Fatal(err) } } @@ -356,6 +359,9 @@ func TestForkAgentInstanceCopiesBoundedHistory(t *testing.T) { db := setupTestDB(t) client := NewClient(db) ctx := context.Background() + sourceID := "66666666-6666-4666-8666-666666666666" + forkID := "77777777-7777-4777-8777-777777777777" + fork2ID := "88888888-8888-4888-8888-888888888888" revision := dbpkg.RuntimeRevision{ Revision: "revision-1", Namespace: "team-a", AgentTemplateName: "assistant", AgentTemplateUID: "template-uid", @@ -380,7 +386,7 @@ func TestForkAgentInstanceCopiesBoundedHistory(t *testing.T) { } source, _, err := client.CreateAgentInstance(ctx, &apiv1alpha1.AgentInstance{ - Id: "instance-1", Namespace: "team-a", Creator: "alice", + Id: sourceID, Namespace: "team-a", Creator: "alice", Harness: &apiv1alpha1.ResourceReference{Namespace: "team-a", Name: "kagent"}, AgentTemplate: &apiv1alpha1.ResourceReference{Namespace: "team-a", Name: "assistant"}, }, "source-request") @@ -404,12 +410,12 @@ func TestForkAgentInstanceCopiesBoundedHistory(t *testing.T) { t.Fatal(err) } checkpoint, err := client.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: "checkpoint-1", Namespace: "team-a", SourceInstanceID: source.GetId(), UserID: "alice", RequestID: "checkpoint-request-1", + ID: uuid.MustParse("99999999-9999-4999-8999-999999999999"), Namespace: "team-a", SourceInstanceID: uuid.MustParse(source.GetId()), UserID: "alice", RequestID: "checkpoint-request-1", }) if err != nil { t.Fatal(err) } - if _, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID, "tag-uid-1", ""); err != nil { + if _, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID.String(), "tag-uid-1", ""); err != nil { t.Fatal(err) } @@ -426,11 +432,11 @@ func TestForkAgentInstanceCopiesBoundedHistory(t *testing.T) { t.Fatal(err) } - fork, created, err := client.ForkAgentInstance(ctx, "team-a", checkpoint.ID, "alice", "fork-request-1", "fork-1") + fork, created, err := client.ForkAgentInstance(ctx, "team-a", checkpoint.ID.String(), "alice", "fork-request-1", forkID) if err != nil || !created { t.Fatalf("ForkAgentInstance() = %+v, created %v, error %v", fork, created, err) } - if fork.GetId() != "fork-1" || fork.GetPreparedRevision() != revision.Revision || fork.GetA2AAuthority() != "" || + if fork.GetId() != forkID || fork.GetPreparedRevision() != revision.Revision || fork.GetA2AAuthority() != "" || fork.GetState() != apiv1alpha1.AgentInstanceState_AGENT_INSTANCE_STATE_CREATING || fork.GetHarness().GetName() != "kagent" || fork.GetAgentTemplate().GetName() != "assistant" || fork.GetLabels()["app"] != "assistant" { @@ -465,14 +471,14 @@ func TestForkAgentInstanceCopiesBoundedHistory(t *testing.T) { if initialMessageID != nil || requestHash != nil || snapshotUID != "snapshot-uid-1" { t.Fatalf("copied persistence metadata = message %v hash %v snapshot %q", initialMessageID, requestHash, snapshotUID) } - replayed, created, err := client.ForkAgentInstance(ctx, "team-a", checkpoint.ID, "alice", "fork-request-1", "ignored") + replayed, created, err := client.ForkAgentInstance(ctx, "team-a", checkpoint.ID.String(), "alice", "fork-request-1", "ignored") if err != nil || created || replayed.GetId() != fork.GetId() { t.Fatalf("replayed fork = %+v, created %v, error %v", replayed, created, err) } - if _, _, err := client.ForkAgentInstance(ctx, "team-a", "other-checkpoint", "alice", "fork-request-1", "ignored"); !errors.Is(err, dbpkg.ErrIdempotencyConflict) { + if _, _, err := client.ForkAgentInstance(ctx, "team-a", "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa", "alice", "fork-request-1", "ignored"); !errors.Is(err, dbpkg.ErrIdempotencyConflict) { t.Fatalf("conflicting fork request error = %v", err) } - if _, err := client.BeginDeleteAgentInstanceCheckpoint(ctx, "team-a", checkpoint.ID, "alice"); !errors.Is(err, dbpkg.ErrNotFound) { + if _, err := client.BeginDeleteAgentInstanceCheckpoint(ctx, "team-a", checkpoint.ID.String(), "alice"); !errors.Is(err, dbpkg.ErrNotFound) { t.Fatalf("delete referenced checkpoint error = %v", err) } @@ -480,19 +486,19 @@ func TestForkAgentInstanceCopiesBoundedHistory(t *testing.T) { t.Fatal(err) } checkpoint2, err := client.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: "checkpoint-2", Namespace: "team-a", SourceInstanceID: fork.GetId(), UserID: "alice", RequestID: "checkpoint-request-2", + ID: uuid.MustParse("bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb"), Namespace: "team-a", SourceInstanceID: uuid.MustParse(fork.GetId()), UserID: "alice", RequestID: "checkpoint-request-2", }) if err != nil { t.Fatal(err) } - if _, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint2.ID, "tag-uid-2", ""); err != nil { + if _, err := client.FinalizeAgentInstanceCheckpoint(ctx, checkpoint2.ID.String(), "tag-uid-2", ""); err != nil { t.Fatal(err) } if err := client.DeleteAgentInstance(ctx, fork.GetId()); err != nil { t.Fatal(err) } - fork2, created, err := client.ForkAgentInstance(ctx, "team-a", checkpoint2.ID, "alice", "fork-request-2", "fork-2") - if err != nil || !created || fork2.GetId() != "fork-2" { + fork2, created, err := client.ForkAgentInstance(ctx, "team-a", checkpoint2.ID.String(), "alice", "fork-request-2", fork2ID) + if err != nil || !created || fork2.GetId() != fork2ID { t.Fatalf("fork of fork = %+v, created %v, error %v", fork2, created, err) } } @@ -531,7 +537,7 @@ func TestAgentInstanceCreateAndTransitions(t *testing.T) { } request := &apiv1alpha1.AgentInstance{ - Id: "instance-1", Namespace: "team-a", Creator: "alice", + Id: "11111111-1111-4111-8111-111111111111", Namespace: "team-a", Creator: "alice", Harness: &apiv1alpha1.ResourceReference{Namespace: "team-a", Name: "kagent"}, AgentTemplate: &apiv1alpha1.ResourceReference{Namespace: "team-a", Name: "assistant"}, } @@ -539,7 +545,7 @@ func TestAgentInstanceCreateAndTransitions(t *testing.T) { if err != nil || !wasCreated { t.Fatalf("first CreateAgentInstance() = created %v, error %v", wasCreated, err) } - request.Id = "instance-2" + request.Id = "22222222-2222-4222-8222-222222222222" replayed, wasCreated, err := client.CreateAgentInstance(ctx, request, "request-1") if err != nil || wasCreated { t.Fatalf("replayed CreateAgentInstance() = created %v, error %v", wasCreated, err) @@ -587,7 +593,7 @@ func TestAgentInstanceCreateAndTransitions(t *testing.T) { func newAgentInstanceTask(id, messageID string) *a2a.Task { now := time.Now() return &a2a.Task{ - ID: a2a.TaskID(id), ContextID: "instance-1", + ID: a2a.TaskID(id), ContextID: "11111111-1111-4111-8111-111111111111", Status: a2a.TaskStatus{State: a2a.TaskStateWorking, Timestamp: &now}, History: []*a2a.Message{{ID: messageID, Role: a2a.MessageRoleUser}}, } @@ -598,47 +604,47 @@ func TestInterruptActiveAgentInstanceTaskRequiresMatchingTaskAndReusesSlot(t *te ctx := context.Background() if _, err := db.Exec(ctx, ` INSERT INTO a2a_context (id, namespace, user_id) - VALUES ('instance-1', 'team-a', 'alice'); + VALUES ('11111111-1111-4111-8111-111111111111', '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') + VALUES ('11111111-1111-4111-8111-111111111111', 'team-a', 'alice', 'request-1', '11111111-1111-4111-8111-111111111111', 'READY', '\x00') `); err != nil { t.Fatal(err) } client := NewClient(db) interrupted := newAgentInstanceTask("task-1", "message-1") - if _, _, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("request-1"), interrupted); err != nil { + if _, _, err := client.CreateAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", []byte("request-1"), interrupted); err != nil { t.Fatal(err) } - active, err := client.GetActiveAgentInstanceTask(ctx, "instance-1") + active, err := client.GetActiveAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111") if err != nil || active.ID != interrupted.ID { t.Fatalf("GetActiveAgentInstanceTask() = %#v, %v", active, err) } - if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "instance-1", "different-task"); err != nil || interruptedTask { + if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "different-task"); err != nil || interruptedTask { t.Fatalf("InterruptActiveAgentInstanceTask(wrong task) = %v, %v", interruptedTask, err) } - if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "instance-1", "task-1"); err != nil || !interruptedTask { + if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "task-1"); err != nil || !interruptedTask { t.Fatalf("InterruptActiveAgentInstanceTask() = %v, %v", interruptedTask, err) } replacement := newAgentInstanceTask("task-2", "message-2") - stored, created, err := client.CreateAgentInstanceTask(ctx, "instance-1", []byte("request-2"), replacement) + stored, created, err := client.CreateAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", []byte("request-2"), replacement) if err != nil || !created || stored.ID != "task-2" { t.Fatalf("send after interruption = %#v, created %v, error %v", stored, created, err) } - if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "instance-1", "task-1"); err != nil || interruptedTask { + if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "task-1"); err != nil || interruptedTask { t.Fatalf("InterruptActiveAgentInstanceTask(replaced task) = %v, %v", interruptedTask, err) } replacement.Status.State = a2a.TaskStateCompleted - if err := client.StoreAgentInstanceTaskEvent(ctx, "instance-1", replacement, replacement, nil); err != nil { + if err := client.StoreAgentInstanceTaskEvent(ctx, "11111111-1111-4111-8111-111111111111", replacement, replacement, nil); err != nil { t.Fatal(err) } - if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "instance-1", "task-2"); err != nil || interruptedTask { + if interruptedTask, err := client.InterruptActiveAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "task-2"); err != nil || interruptedTask { t.Fatalf("InterruptActiveAgentInstanceTask(terminal task) = %v, %v", interruptedTask, err) } - terminated, err := client.GetAgentInstanceTask(ctx, "instance-1", "task-1") + terminated, err := client.GetAgentInstanceTask(ctx, "11111111-1111-4111-8111-111111111111", "task-1") if err != nil { t.Fatal(err) } @@ -700,6 +706,8 @@ func TestAgentInstanceNameRoundTripsAndRenames(t *testing.T) { client := NewClient(setupTestDB(t)) ctx := context.Background() agentInstanceFixture(t, client, ctx, "revision-1", "assistant", "kagent") + namedID := "11111111-1111-4111-8111-111111111111" + unnamedID := "22222222-2222-4222-8222-222222222222" for _, test := range []struct { name string @@ -707,10 +715,10 @@ func TestAgentInstanceNameRoundTripsAndRenames(t *testing.T) { given string wantName string }{ - {name: "a name round-trips", id: "instance-named", given: "Debugging the ingress", wantName: "Debugging the ingress"}, + {name: "a name round-trips", id: namedID, given: "Debugging the ingress", wantName: "Debugging the ingress"}, // An instance created without a name must read back empty, which is how // every row written before the column existed reads. - {name: "an omitted name stays empty", id: "instance-unnamed", given: "", wantName: ""}, + {name: "an omitted name stays empty", id: unnamedID, given: "", wantName: ""}, } { t.Run(test.name, func(t *testing.T) { created, wasCreated, err := client.CreateAgentInstance(ctx, newAgentInstanceRequest(test.id, "assistant", "kagent", test.given), test.id) @@ -727,27 +735,27 @@ func TestAgentInstanceNameRoundTripsAndRenames(t *testing.T) { }) } - renamed, err := client.UpdateAgentInstanceName(ctx, "team-a", "instance-unnamed", "alice", "Named afterwards") + renamed, err := client.UpdateAgentInstanceName(ctx, "team-a", unnamedID, "alice", "Named afterwards") if err != nil || renamed.GetName() != "Named afterwards" { t.Fatalf("UpdateAgentInstanceName() = %+v, error %v", renamed, err) } // The rename has to survive a re-read, not just be echoed back: the name lives // in a column while the rest of the message lives in a blob the rename does not // rewrite, so an echoed value proves nothing about what was stored. - read, err := client.GetAgentInstance(ctx, "team-a", "instance-unnamed", "alice") + read, err := client.GetAgentInstance(ctx, "team-a", unnamedID, "alice") if err != nil || read.GetName() != "Named afterwards" { t.Fatalf("re-read after rename = %+v, error %v", read, err) } // Renaming back to empty must be possible, or a name can never be undone. - cleared, err := client.UpdateAgentInstanceName(ctx, "team-a", "instance-unnamed", "alice", "") + cleared, err := client.UpdateAgentInstanceName(ctx, "team-a", unnamedID, "alice", "") if err != nil || cleared.GetName() != "" { t.Fatalf("UpdateAgentInstanceName(\"\") = %+v, error %v", cleared, err) } // A rename is scoped to the owner, so it cannot reach another reader's row. - if _, err := client.UpdateAgentInstanceName(ctx, "team-a", "instance-named", "bob", "Stolen"); !errors.Is(err, dbpkg.ErrNotFound) { + if _, err := client.UpdateAgentInstanceName(ctx, "team-a", namedID, "bob", "Stolen"); !errors.Is(err, dbpkg.ErrNotFound) { t.Fatalf("UpdateAgentInstanceName() as another user error = %v, want %v", err, dbpkg.ErrNotFound) } - if _, err := client.UpdateAgentInstanceName(ctx, "team-a", "missing", "alice", "Nothing"); !errors.Is(err, dbpkg.ErrNotFound) { + if _, err := client.UpdateAgentInstanceName(ctx, "team-a", "33333333-3333-4333-8333-333333333333", "alice", "Nothing"); !errors.Is(err, dbpkg.ErrNotFound) { t.Fatalf("UpdateAgentInstanceName() of a missing instance error = %v, want %v", err, dbpkg.ErrNotFound) } } @@ -765,9 +773,9 @@ func TestListAgentInstancesFiltersByAgentPair(t *testing.T) { agentInstanceFixture(t, client, ctx, "revision-3", "researcher", "kagent") for id, pair := range map[string][2]string{ - "instance-1": {"assistant", "kagent"}, - "instance-2": {"assistant", "claude"}, - "instance-3": {"researcher", "kagent"}, + "11111111-1111-4111-8111-111111111111": {"assistant", "kagent"}, + "22222222-2222-4222-8222-222222222222": {"assistant", "claude"}, + "33333333-3333-4333-8333-333333333333": {"researcher", "kagent"}, } { if _, _, err := client.CreateAgentInstance(ctx, newAgentInstanceRequest(id, pair[0], pair[1], ""), id); err != nil { t.Fatalf("CreateAgentInstance(%s) error %v", id, err) @@ -782,29 +790,29 @@ func TestListAgentInstancesFiltersByAgentPair(t *testing.T) { { name: "no filter lists every conversation", query: dbpkg.AgentInstanceQuery{}, - want: []string{"instance-1", "instance-2", "instance-3"}, + want: []string{"11111111-1111-4111-8111-111111111111", "22222222-2222-4222-8222-222222222222", "33333333-3333-4333-8333-333333333333"}, }, { name: "one agent, which is one pair", query: dbpkg.AgentInstanceQuery{AgentTemplate: "assistant", Harness: "kagent"}, - want: []string{"instance-1"}, + want: []string{"11111111-1111-4111-8111-111111111111"}, }, { // The case labels could never serve: one template, two harnesses, two // agents, and identical labels on both instances. name: "the same template on a different harness is a different agent", query: dbpkg.AgentInstanceQuery{AgentTemplate: "assistant", Harness: "claude"}, - want: []string{"instance-2"}, + want: []string{"22222222-2222-4222-8222-222222222222"}, }, { name: "template alone spans its harnesses", query: dbpkg.AgentInstanceQuery{AgentTemplate: "assistant"}, - want: []string{"instance-1", "instance-2"}, + want: []string{"11111111-1111-4111-8111-111111111111", "22222222-2222-4222-8222-222222222222"}, }, { name: "harness alone spans its templates", query: dbpkg.AgentInstanceQuery{Harness: "kagent"}, - want: []string{"instance-1", "instance-3"}, + want: []string{"11111111-1111-4111-8111-111111111111", "33333333-3333-4333-8333-333333333333"}, }, { name: "an unknown agent matches nothing rather than everything", diff --git a/go/core/internal/database/client_postgres.go b/go/core/internal/database/client_postgres.go index 2903d8e07..5f1fed410 100644 --- a/go/core/internal/database/client_postgres.go +++ b/go/core/internal/database/client_postgres.go @@ -212,6 +212,7 @@ func toAgentInstance(row dbgen.AgentInstance) (*apiv1alpha1.AgentInstance, error // migrations can backfill them without rewriting the protobuf blob. instance.State = apiv1alpha1.AgentInstanceState(state) instance.Operation = apiv1alpha1.AgentInstanceOperation(operationValue) + instance.Id = row.ID.String() // The name is a column for the same reason, and because a rename writes only // the column: reading it from the blob would serve the name the row was // created with for the rest of the instance's life. @@ -270,16 +271,17 @@ func (c *postgresClient) CreateAgentInstance(ctx context.Context, request *apiv1 if err != nil { return nil, false, err } + instanceID := uuid.MustParse(request.GetId()) var row dbgen.AgentInstance err = c.withTx(ctx, func(q *dbgen.Queries) error { if err := q.InsertA2AContext(ctx, dbgen.InsertA2AContextParams{ - ID: request.GetId(), Namespace: request.GetNamespace(), UserID: request.GetCreator(), + ID: instanceID, Namespace: request.GetNamespace(), UserID: request.GetCreator(), }); err != nil { return fmt.Errorf("insert A2A context: %w", err) } row, err = q.InsertAgentInstance(ctx, dbgen.InsertAgentInstanceParams{ - ID: request.GetId(), Namespace: request.GetNamespace(), UserID: request.GetCreator(), RequestID: requestID, - ContextID: request.GetId(), PreparedRevision: &revision.Revision, Labels: revision.AgentTemplateLabels, + ID: instanceID, Namespace: request.GetNamespace(), UserID: request.GetCreator(), RequestID: requestID, + ContextID: instanceID, PreparedRevision: &revision.Revision, Labels: revision.AgentTemplateLabels, Name: request.GetName(), Data: data, }) return err @@ -303,10 +305,11 @@ func (c *postgresClient) CreateAgentInstance(ctx context.Context, request *apiv1 } func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, checkpointID, userID, requestID, instanceID string) (*apiv1alpha1.AgentInstance, bool, error) { + checkpointUUID := uuid.MustParse(checkpointID) requestKey := dbgen.GetAgentInstanceByRequestParams{UserID: userID, Namespace: namespace, RequestID: requestID} existing, err := c.q.GetAgentInstanceByRequest(ctx, requestKey) if err == nil { - if derefStr(existing.SourceCheckpointID) != checkpointID { + if existing.SourceCheckpointID == nil || *existing.SourceCheckpointID != checkpointUUID { return nil, false, dbpkg.ErrIdempotencyConflict } instance, err := toAgentInstance(existing) @@ -316,10 +319,11 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check return nil, false, fmt.Errorf("get fork request: %w", err) } + instanceUUID := uuid.MustParse(instanceID) var row dbgen.AgentInstance err = c.withTx(ctx, func(q *dbgen.Queries) error { checkpoint, err := q.LockReadyAgentInstanceCheckpoint(ctx, dbgen.LockReadyAgentInstanceCheckpointParams{ - Namespace: namespace, ID: checkpointID, UserID: userID, + Namespace: namespace, ID: checkpointUUID, UserID: userID, }) if errors.Is(err, pgx.ErrNoRows) { return dbpkg.ErrNotFound @@ -357,12 +361,12 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check if err != nil { return fmt.Errorf("encode fork labels: %w", err) } - if err := q.InsertA2AContext(ctx, dbgen.InsertA2AContextParams{ID: instanceID, Namespace: namespace, UserID: userID}); err != nil { + if err := q.InsertA2AContext(ctx, dbgen.InsertA2AContextParams{ID: instanceUUID, Namespace: namespace, UserID: userID}); err != nil { return fmt.Errorf("insert fork A2A context: %w", err) } row, err = q.InsertForkedAgentInstance(ctx, dbgen.InsertForkedAgentInstanceParams{ - ID: instanceID, Namespace: namespace, UserID: userID, RequestID: requestID, - ContextID: instanceID, PreparedRevision: checkpoint.PreparedRevision, + ID: instanceUUID, Namespace: namespace, UserID: userID, RequestID: requestID, + ContextID: instanceUUID, PreparedRevision: checkpoint.PreparedRevision, SourceCheckpointID: &checkpoint.ID, Labels: encodedLabels, Data: data, }) if err != nil { @@ -382,7 +386,7 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check } reidentifyForkTask(task, instanceID, ids) if len(task.History) > 0 { - copiedHistorySequence, err = storeAgentInstanceTaskMessages(ctx, q, instanceID, string(task.ID), task.History) + copiedHistorySequence, err = storeAgentInstanceTaskMessages(ctx, q, instanceUUID, string(task.ID), task.History) if err != nil { return fmt.Errorf("copy checkpoint task %s history: %w", source.ID, err) } @@ -392,7 +396,7 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check return fmt.Errorf("encode fork task %s: %w", task.ID, err) } copy := dbgen.InsertCopiedAgentInstanceTaskParams{ - ContextID: instanceID, ID: string(task.ID), State: string(task.Status.State), + ContextID: instanceUUID, ID: string(task.ID), State: string(task.Status.State), StatusTimestamp: task.Status.Timestamp, Data: data, CreatedAt: source.CreatedAt, UpdatedAt: source.UpdatedAt, } @@ -420,7 +424,7 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check messageID = &mapped } copiedHistorySequence, err = q.InsertAgentInstanceTaskEvent(ctx, dbgen.InsertAgentInstanceTaskEventParams{ - ContextID: instanceID, TaskID: strPtrIfNotEmpty(string(event.TaskInfo().TaskID)), + ContextID: instanceUUID, TaskID: strPtrIfNotEmpty(string(event.TaskInfo().TaskID)), MessageID: messageID, Data: data, }) if err != nil { @@ -432,7 +436,7 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check } headID := string(ids.task(a2a.TaskID(checkpoint.HeadTaskID))) if err := q.SetAgentInstanceTaskSnapshot(ctx, dbgen.SetAgentInstanceTaskSnapshotParams{ - ContextID: instanceID, ID: headID, SnapshotAtespace: &checkpoint.SnapshotAtespace, + ContextID: instanceUUID, ID: headID, SnapshotAtespace: &checkpoint.SnapshotAtespace, SnapshotName: &checkpoint.SnapshotName, SnapshotUid: &checkpoint.SnapshotUid, SnapshotContentScope: &checkpoint.SnapshotContentScope, HistorySequence: &copiedHistorySequence, }); err != nil { @@ -445,7 +449,7 @@ func (c *postgresClient) ForkAgentInstance(ctx context.Context, namespace, check if err != nil { return nil, false, fmt.Errorf("get concurrent fork request: %w", err) } - if derefStr(existing.SourceCheckpointID) != checkpointID { + if existing.SourceCheckpointID == nil || *existing.SourceCheckpointID != checkpointUUID { return nil, false, dbpkg.ErrIdempotencyConflict } instance, err := toAgentInstance(existing) @@ -537,7 +541,7 @@ func reidentifyForkEvent(event a2a.Event, sourceTaskID, contextID string, ids *f } func (c *postgresClient) GetAgentInstance(ctx context.Context, namespace, id, userID string) (*apiv1alpha1.AgentInstance, error) { - row, err := c.q.GetAgentInstanceForUser(ctx, dbgen.GetAgentInstanceForUserParams{Namespace: namespace, ID: id, UserID: userID}) + row, err := c.q.GetAgentInstanceForUser(ctx, dbgen.GetAgentInstanceForUserParams{Namespace: namespace, ID: uuid.MustParse(id), UserID: userID}) if err != nil { return nil, fmt.Errorf("get AgentInstance %s/%s: %w", namespace, id, notFoundOr(err)) } @@ -577,7 +581,7 @@ func (c *postgresClient) ListAgentInstances(ctx context.Context, query dbpkg.Age // owner so a rename cannot reach another reader's conversation. func (c *postgresClient) UpdateAgentInstanceName(ctx context.Context, namespace, id, userID, name string) (*apiv1alpha1.AgentInstance, error) { row, err := c.q.UpdateAgentInstanceName(ctx, dbgen.UpdateAgentInstanceNameParams{ - Namespace: namespace, ID: id, UserID: userID, Name: name, + Namespace: namespace, ID: uuid.MustParse(id), UserID: userID, Name: name, }) if err != nil { return nil, fmt.Errorf("rename AgentInstance %s/%s: %w", namespace, id, notFoundOr(err)) @@ -586,7 +590,8 @@ func (c *postgresClient) UpdateAgentInstanceName(ctx context.Context, namespace, } func (c *postgresClient) MarkAgentInstanceReady(ctx context.Context, id, authority string) (*apiv1alpha1.AgentInstance, error) { - row, err := c.q.GetAgentInstanceByID(ctx, id) + instanceID := uuid.MustParse(id) + row, err := c.q.GetAgentInstanceByID(ctx, instanceID) if err != nil { return nil, fmt.Errorf("get AgentInstance %s before ready: %w", id, notFoundOr(err)) } @@ -603,9 +608,9 @@ func (c *postgresClient) MarkAgentInstanceReady(ctx context.Context, id, authori if err != nil { return nil, err } - row, err = c.q.MarkAgentInstanceReady(ctx, dbgen.MarkAgentInstanceReadyParams{ID: id, Data: data}) + row, err = c.q.MarkAgentInstanceReady(ctx, dbgen.MarkAgentInstanceReadyParams{ID: instanceID, Data: data}) if errors.Is(err, pgx.ErrNoRows) { - row, err = c.q.GetAgentInstanceByID(ctx, id) + row, err = c.q.GetAgentInstanceByID(ctx, instanceID) } if err != nil { return nil, fmt.Errorf("mark AgentInstance %s ready: %w", id, notFoundOr(err)) @@ -624,12 +629,12 @@ func (c *postgresClient) TransitionAgentInstance( return nil, err } row, err := c.q.TransitionAgentInstance(ctx, dbgen.TransitionAgentInstanceParams{ - ID: instance.GetId(), Data: data, + ID: uuid.MustParse(instance.GetId()), Data: data, ExpectedState: agentInstanceStateName(expectedState), ExpectedOperation: agentInstanceOperationName(expectedOperation), NextState: agentInstanceStateName(instance.GetState()), NextOperation: agentInstanceOperationName(instance.GetOperation()), }) if errors.Is(err, pgx.ErrNoRows) { - row, err = c.q.GetAgentInstanceByID(ctx, instance.GetId()) + row, err = c.q.GetAgentInstanceByID(ctx, uuid.MustParse(instance.GetId())) if errors.Is(err, pgx.ErrNoRows) { return nil, fmt.Errorf("transition AgentInstance %s: %w", instance.GetId(), dbpkg.ErrNotFound) } @@ -660,7 +665,7 @@ func agentInstanceOperationName(operation apiv1alpha1.AgentInstanceOperation) st } func (c *postgresClient) DeleteAgentInstance(ctx context.Context, id string) error { - if err := c.q.DeleteAgentInstance(ctx, id); err != nil { + if err := c.q.DeleteAgentInstance(ctx, uuid.MustParse(id)); err != nil { return fmt.Errorf("delete AgentInstance %s: %w", id, err) } return nil @@ -703,7 +708,7 @@ func (c *postgresClient) GetAgentInstanceShareByTokenHash(ctx context.Context, t func (c *postgresClient) ListAgentInstanceShares(ctx context.Context, namespace, instanceID, userID, afterID string, limit int) ([]dbpkg.AgentInstanceShare, error) { rows, err := c.q.ListAgentInstanceShares(ctx, dbgen.ListAgentInstanceSharesParams{ - Namespace: namespace, InstanceID: instanceID, UserID: userID, + Namespace: namespace, InstanceID: uuid.MustParse(instanceID), UserID: userID, AfterID: afterID, PageSize: int32(limit), }) if err != nil { @@ -717,7 +722,7 @@ func (c *postgresClient) ListAgentInstanceShares(ctx context.Context, namespace, } func (c *postgresClient) DeleteAgentInstanceShare(ctx context.Context, namespace, id, userID string) error { - count, err := c.q.DeleteAgentInstanceShare(ctx, dbgen.DeleteAgentInstanceShareParams{Namespace: namespace, ID: id, UserID: userID}) + count, err := c.q.DeleteAgentInstanceShare(ctx, dbgen.DeleteAgentInstanceShareParams{Namespace: namespace, ID: uuid.MustParse(id), UserID: userID}) if err != nil { return fmt.Errorf("delete AgentInstance share %s/%s: %w", namespace, id, err) } @@ -739,8 +744,9 @@ func (c *postgresClient) CreateAgentInstanceTask(ctx context.Context, instanceID result := task created := false + contextID := uuid.MustParse(instanceID) err = c.withTx(ctx, func(q *dbgen.Queries) error { - instance, err := q.LockAgentInstance(ctx, instanceID) + instance, err := q.LockAgentInstance(ctx, contextID) if err != nil { return fmt.Errorf("lock AgentInstance %s: %w", instanceID, err) } @@ -748,7 +754,7 @@ func (c *postgresClient) CreateAgentInstanceTask(ctx context.Context, instanceID return dbpkg.ErrAgentInstanceTaskConflict } rows, err := q.CreateAgentInstanceTask(ctx, dbgen.CreateAgentInstanceTaskParams{ - ContextID: instanceID, ID: string(task.ID), State: string(task.Status.State), + ContextID: contextID, ID: string(task.ID), State: string(task.Status.State), StatusTimestamp: task.Status.Timestamp, Data: taskData, InitialMessageID: &message.ID, RequestHash: requestHash, }) @@ -760,7 +766,7 @@ func (c *postgresClient) CreateAgentInstanceTask(ctx context.Context, instanceID } if rows == 0 { row, err := q.GetAgentInstanceTaskByMessageID(ctx, dbgen.GetAgentInstanceTaskByMessageIDParams{ - ContextID: instanceID, InitialMessageID: &message.ID, + ContextID: contextID, InitialMessageID: &message.ID, }) if errors.Is(err, pgx.ErrNoRows) { return dbpkg.ErrAgentInstanceTaskConflict @@ -775,10 +781,10 @@ func (c *postgresClient) CreateAgentInstanceTask(ctx context.Context, instanceID if err != nil { return err } - return loadAgentInstanceTaskHistories(ctx, q, instanceID, []*a2a.Task{result}) + return loadAgentInstanceTaskHistories(ctx, q, contextID, []*a2a.Task{result}) } created = true - _, err = storeAgentInstanceTaskMessages(ctx, q, instanceID, string(task.ID), task.History) + _, err = storeAgentInstanceTaskMessages(ctx, q, contextID, string(task.ID), task.History) return err }) if err != nil { @@ -792,13 +798,14 @@ func (c *postgresClient) CreateAgentInstanceTask(ctx context.Context, instanceID const taskInterruptedMessage = "The turn was interrupted before it completed, and the process running it is no longer reporting progress." func (c *postgresClient) GetActiveAgentInstanceTask(ctx context.Context, instanceID string) (*a2a.Task, error) { - row, err := c.q.GetActiveAgentInstanceTask(ctx, instanceID) + contextID := uuid.MustParse(instanceID) + row, err := c.q.GetActiveAgentInstanceTask(ctx, contextID) if err != nil { return nil, fmt.Errorf("get active AgentInstance task: %w", notFoundOr(err)) } task, err := unmarshalAgentInstanceTask(row.Data) if err == nil { - err = loadAgentInstanceTaskHistories(ctx, c.q, instanceID, []*a2a.Task{task}) + err = loadAgentInstanceTaskHistories(ctx, c.q, contextID, []*a2a.Task{task}) } return task, err } @@ -807,8 +814,9 @@ func (c *postgresClient) GetActiveAgentInstanceTask(ctx context.Context, instanc // instance's active task. func (c *postgresClient) InterruptActiveAgentInstanceTask(ctx context.Context, instanceID, taskID string) (bool, error) { interruptedTask := false + contextID := uuid.MustParse(instanceID) err := c.withTx(ctx, func(q *dbgen.Queries) error { - row, err := q.LockActiveAgentInstanceTask(ctx, instanceID) + row, err := q.LockActiveAgentInstanceTask(ctx, contextID) if errors.Is(err, pgx.ErrNoRows) { return nil } @@ -822,7 +830,7 @@ func (c *postgresClient) InterruptActiveAgentInstanceTask(ctx context.Context, i if err != nil { return err } - if err := loadAgentInstanceTaskHistories(ctx, q, instanceID, []*a2a.Task{task}); err != nil { + if err := loadAgentInstanceTaskHistories(ctx, q, contextID, []*a2a.Task{task}); err != nil { return err } interrupted := a2a.NewMessageForTask(a2a.MessageRoleAgent, task, a2a.NewTextPart(taskInterruptedMessage)) @@ -834,12 +842,12 @@ func (c *postgresClient) InterruptActiveAgentInstanceTask(ctx context.Context, i return err } if err := q.UpsertAgentInstanceTask(ctx, dbgen.UpsertAgentInstanceTaskParams{ - ContextID: instanceID, ID: string(task.ID), State: string(task.Status.State), + ContextID: contextID, ID: string(task.ID), State: string(task.Status.State), StatusTimestamp: task.Status.Timestamp, Data: data, }); err != nil { return fmt.Errorf("interrupt AgentInstance task %s: %w", task.ID, err) } - if _, err := storeAgentInstanceTaskMessages(ctx, q, instanceID, string(task.ID), task.History); err != nil { + if _, err := storeAgentInstanceTaskMessages(ctx, q, contextID, string(task.ID), task.History); err != nil { return fmt.Errorf("record AgentInstance task interruption: %w", err) } interruptedTask = true @@ -852,11 +860,12 @@ 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 { + contextID := uuid.MustParse(instanceID) 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 { + if row, err := q.GetAgentInstanceTask(ctx, dbgen.GetAgentInstanceTaskParams{ContextID: contextID, ID: string(task.ID)}); err == nil { previous, err := unmarshalAgentInstanceTask(row.Data) if err != nil { return err @@ -868,7 +877,7 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan replacedStatusMessage = &message } if len(previous.History) > 0 { - sequence, err = storeAgentInstanceTaskMessages(ctx, q, instanceID, string(task.ID), previous.History) + sequence, err = storeAgentInstanceTaskMessages(ctx, q, contextID, string(task.ID), previous.History) if err != nil { return fmt.Errorf("normalize legacy AgentInstance task history: %w", err) } @@ -881,7 +890,7 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan return err } if err := q.UpsertAgentInstanceTask(ctx, dbgen.UpsertAgentInstanceTaskParams{ - ContextID: instanceID, ID: string(task.ID), State: string(task.Status.State), + ContextID: contextID, ID: string(task.ID), State: string(task.Status.State), StatusTimestamp: task.Status.Timestamp, Data: data, }); err != nil { if isActiveTaskConflict(err) { @@ -896,7 +905,7 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan } if len(messages) > 0 { var err error - sequence, err = storeAgentInstanceTaskMessages(ctx, q, instanceID, string(event.TaskInfo().TaskID), messages) + sequence, err = storeAgentInstanceTaskMessages(ctx, q, contextID, string(event.TaskInfo().TaskID), messages) if err != nil { return fmt.Errorf("store AgentInstance task history: %w", err) } @@ -907,7 +916,7 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan return err } sequence, err = q.InsertAgentInstanceTaskEvent(ctx, dbgen.InsertAgentInstanceTaskEventParams{ - ContextID: instanceID, TaskID: strPtrIfNotEmpty(string(event.TaskInfo().TaskID)), Data: eventData, + ContextID: contextID, TaskID: strPtrIfNotEmpty(string(event.TaskInfo().TaskID)), Data: eventData, }) if err != nil { return fmt.Errorf("store AgentInstance task event: %w", err) @@ -918,7 +927,7 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan return fmt.Errorf("snapshot has no history boundary") } if err := q.SetAgentInstanceTaskSnapshot(ctx, dbgen.SetAgentInstanceTaskSnapshotParams{ - ContextID: instanceID, ID: string(task.ID), SnapshotAtespace: &snapshot.Atespace, + ContextID: contextID, ID: string(task.ID), SnapshotAtespace: &snapshot.Atespace, SnapshotName: &snapshot.Name, SnapshotUid: &snapshot.UID, SnapshotContentScope: &snapshot.ContentScope, HistorySequence: &sequence, }); err != nil { @@ -934,27 +943,29 @@ func (c *postgresClient) StoreAgentInstanceTaskEvent(ctx context.Context, instan } func (c *postgresClient) GetAgentInstanceTask(ctx context.Context, instanceID, taskID string) (*a2a.Task, error) { - row, err := c.q.GetAgentInstanceTask(ctx, dbgen.GetAgentInstanceTaskParams{ContextID: instanceID, ID: taskID}) + contextID := uuid.MustParse(instanceID) + row, err := c.q.GetAgentInstanceTask(ctx, dbgen.GetAgentInstanceTaskParams{ContextID: contextID, ID: taskID}) if err != nil { return nil, fmt.Errorf("get AgentInstance task %s: %w", taskID, notFoundOr(err)) } task, err := unmarshalAgentInstanceTask(row.Data) if err == nil { - err = loadAgentInstanceTaskHistories(ctx, c.q, instanceID, []*a2a.Task{task}) + err = loadAgentInstanceTaskHistories(ctx, c.q, contextID, []*a2a.Task{task}) } return task, err } func (c *postgresClient) ListAgentInstanceTasks(ctx context.Context, instanceID, afterID string, state a2a.TaskState, statusTimestampAfter *time.Time, limit int) ([]*a2a.Task, int, error) { + contextID := uuid.MustParse(instanceID) params := dbgen.CountAgentInstanceTasksParams{ - ContextID: instanceID, State: string(state), StatusTimestampAfter: statusTimestampAfter, + ContextID: contextID, State: string(state), StatusTimestampAfter: statusTimestampAfter, } total, err := c.q.CountAgentInstanceTasks(ctx, params) if err != nil { return nil, 0, fmt.Errorf("count AgentInstance tasks: %w", err) } rows, err := c.q.ListAgentInstanceTasks(ctx, dbgen.ListAgentInstanceTasksParams{ - ContextID: instanceID, AfterID: afterID, State: params.State, + ContextID: contextID, AfterID: afterID, State: params.State, StatusTimestampAfter: statusTimestampAfter, PageSize: int32(limit), }) if err != nil { @@ -968,7 +979,7 @@ func (c *postgresClient) ListAgentInstanceTasks(ctx context.Context, instanceID, } tasks = append(tasks, task) } - if err := loadAgentInstanceTaskHistories(ctx, c.q, instanceID, tasks); err != nil { + if err := loadAgentInstanceTaskHistories(ctx, c.q, contextID, tasks); err != nil { return nil, 0, err } return tasks, int(total), nil @@ -1055,7 +1066,7 @@ func (c *postgresClient) FinalizeAgentInstanceCheckpoint(ctx context.Context, id return nil, fmt.Errorf("finalize AgentInstance checkpoint requires exactly one of tag UID or failure") } row, err := c.q.FinalizeAgentInstanceCheckpoint(ctx, dbgen.FinalizeAgentInstanceCheckpointParams{ - ID: id, TagUid: tagUID, Failure: failure, + ID: uuid.MustParse(id), TagUid: tagUID, Failure: failure, }) if err != nil { return nil, fmt.Errorf("finalize AgentInstance checkpoint: %w", notFoundOr(err)) @@ -1064,7 +1075,7 @@ func (c *postgresClient) FinalizeAgentInstanceCheckpoint(ctx context.Context, id } func (c *postgresClient) GetAgentInstanceCheckpoint(ctx context.Context, namespace, id, userID string) (*dbpkg.AgentInstanceCheckpoint, error) { - row, err := c.q.GetAgentInstanceCheckpoint(ctx, dbgen.GetAgentInstanceCheckpointParams{Namespace: namespace, ID: id, UserID: userID}) + row, err := c.q.GetAgentInstanceCheckpoint(ctx, dbgen.GetAgentInstanceCheckpointParams{Namespace: namespace, ID: uuid.MustParse(id), UserID: userID}) if err != nil { return nil, fmt.Errorf("get AgentInstance checkpoint: %w", notFoundOr(err)) } @@ -1073,7 +1084,7 @@ func (c *postgresClient) GetAgentInstanceCheckpoint(ctx context.Context, namespa func (c *postgresClient) ListAgentInstanceCheckpoints(ctx context.Context, namespace, instanceID, userID, afterID string, limit int) ([]dbpkg.AgentInstanceCheckpoint, error) { rows, err := c.q.ListAgentInstanceCheckpoints(ctx, dbgen.ListAgentInstanceCheckpointsParams{ - Namespace: namespace, SourceInstanceID: instanceID, UserID: userID, AfterID: afterID, PageSize: int32(limit), + Namespace: namespace, SourceInstanceID: uuid.MustParse(instanceID), UserID: userID, AfterID: afterID, PageSize: int32(limit), }) if err != nil { return nil, fmt.Errorf("list AgentInstance checkpoints: %w", err) @@ -1087,7 +1098,7 @@ func (c *postgresClient) ListAgentInstanceCheckpoints(ctx context.Context, names func (c *postgresClient) BeginDeleteAgentInstanceCheckpoint(ctx context.Context, namespace, id, userID string) (*dbpkg.AgentInstanceCheckpoint, error) { row, err := c.q.BeginDeleteAgentInstanceCheckpoint(ctx, dbgen.BeginDeleteAgentInstanceCheckpointParams{ - Namespace: namespace, ID: id, UserID: userID, + Namespace: namespace, ID: uuid.MustParse(id), UserID: userID, }) if err != nil { return nil, fmt.Errorf("begin delete AgentInstance checkpoint: %w", notFoundOr(err)) @@ -1096,7 +1107,7 @@ func (c *postgresClient) BeginDeleteAgentInstanceCheckpoint(ctx context.Context, } func (c *postgresClient) DeleteAgentInstanceCheckpoint(ctx context.Context, namespace, id, userID string) error { - _, err := c.q.DeleteAgentInstanceCheckpoint(ctx, dbgen.DeleteAgentInstanceCheckpointParams{Namespace: namespace, ID: id, UserID: userID}) + _, err := c.q.DeleteAgentInstanceCheckpoint(ctx, dbgen.DeleteAgentInstanceCheckpointParams{Namespace: namespace, ID: uuid.MustParse(id), UserID: userID}) if err != nil { return fmt.Errorf("delete AgentInstance checkpoint: %w", err) } @@ -1160,7 +1171,7 @@ func agentInstanceTaskEventMessages(task *a2a.Task, event a2a.Event) []*a2a.Mess return nil } -func storeAgentInstanceTaskMessages(ctx context.Context, q *dbgen.Queries, contextID, taskID string, messages []*a2a.Message) (int64, error) { +func storeAgentInstanceTaskMessages(ctx context.Context, q *dbgen.Queries, contextID uuid.UUID, taskID string, messages []*a2a.Message) (int64, error) { var sequence int64 for _, message := range messages { if message == nil || message.ID == "" { @@ -1180,7 +1191,7 @@ func storeAgentInstanceTaskMessages(ctx context.Context, q *dbgen.Queries, conte return sequence, nil } -func loadAgentInstanceTaskHistories(ctx context.Context, q *dbgen.Queries, contextID string, tasks []*a2a.Task) error { +func loadAgentInstanceTaskHistories(ctx context.Context, q *dbgen.Queries, contextID uuid.UUID, tasks []*a2a.Task) error { if len(tasks) == 0 { return nil } diff --git a/go/core/internal/database/gen/agent_instance_checkpoints.sql.go b/go/core/internal/database/gen/agent_instance_checkpoints.sql.go index 3312d4d84..7a3fb49ab 100644 --- a/go/core/internal/database/gen/agent_instance_checkpoints.sql.go +++ b/go/core/internal/database/gen/agent_instance_checkpoints.sql.go @@ -7,6 +7,8 @@ package dbgen import ( "context" + + "github.com/google/uuid" ) const beginDeleteAgentInstanceCheckpoint = `-- name: BeginDeleteAgentInstanceCheckpoint :one @@ -22,7 +24,7 @@ RETURNING id, namespace, source_instance_id, user_id, request_id, head_task_id, type BeginDeleteAgentInstanceCheckpointParams struct { Namespace string - ID string + ID uuid.UUID UserID string } @@ -59,7 +61,7 @@ WHERE namespace = $1 AND id = $2 AND user_id = $3 AND state = 'DELETING' type DeleteAgentInstanceCheckpointParams struct { Namespace string - ID string + ID uuid.UUID UserID string } @@ -86,7 +88,7 @@ RETURNING id, namespace, source_instance_id, user_id, request_id, head_task_id, ` type FinalizeAgentInstanceCheckpointParams struct { - ID string + ID uuid.UUID TagUid string Failure string } @@ -124,7 +126,7 @@ WHERE namespace = $1 AND id = $2 AND user_id = $3 AND state = 'READY' type GetAgentInstanceCheckpointParams struct { Namespace string - ID string + ID uuid.UUID UserID string } @@ -213,7 +215,7 @@ WHERE NOT EXISTS ( ) ` -func (q *Queries) GetLatestQuiescentAgentInstanceTask(ctx context.Context, contextID string) (AgentInstanceTask, error) { +func (q *Queries) GetLatestQuiescentAgentInstanceTask(ctx context.Context, contextID uuid.UUID) (AgentInstanceTask, error) { row := q.db.QueryRow(ctx, getLatestQuiescentAgentInstanceTask, contextID) var i AgentInstanceTask err := row.Scan( @@ -246,9 +248,9 @@ RETURNING id, namespace, source_instance_id, user_id, request_id, head_task_id, ` type InsertAgentInstanceCheckpointParams struct { - ID string + ID uuid.UUID Namespace string - SourceInstanceID string + SourceInstanceID uuid.UUID UserID string RequestID string HeadTaskID string @@ -257,7 +259,7 @@ type InsertAgentInstanceCheckpointParams struct { SnapshotName string SnapshotUid string SnapshotContentScope string - SourceContextID string + SourceContextID uuid.UUID PreparedRevision *string SourceLabels []byte } @@ -313,7 +315,7 @@ WHERE c.id = $1 ORDER BY e.sequence ` -func (q *Queries) ListAgentInstanceCheckpointEvents(ctx context.Context, checkpointID string) ([]AgentInstanceTaskEvent, error) { +func (q *Queries) ListAgentInstanceCheckpointEvents(ctx context.Context, checkpointID uuid.UUID) ([]AgentInstanceTaskEvent, error) { rows, err := q.db.Query(ctx, listAgentInstanceCheckpointEvents, checkpointID) if err != nil { return nil, err @@ -352,7 +354,7 @@ WHERE c.id = $1 ORDER BY t.created_at, t.id ` -func (q *Queries) ListAgentInstanceCheckpointTasks(ctx context.Context, checkpointID string) ([]AgentInstanceTask, error) { +func (q *Queries) ListAgentInstanceCheckpointTasks(ctx context.Context, checkpointID uuid.UUID) ([]AgentInstanceTask, error) { rows, err := q.db.Query(ctx, listAgentInstanceCheckpointTasks, checkpointID) if err != nil { return nil, err @@ -393,14 +395,14 @@ WHERE namespace = $1 AND source_instance_id = $2 AND user_id = $3 AND state = 'READY' - AND id > $4 + AND (NULLIF($4::text, '') IS NULL OR id > NULLIF($4::text, '')::uuid) ORDER BY id LIMIT $5 ` type ListAgentInstanceCheckpointsParams struct { Namespace string - SourceInstanceID string + SourceInstanceID uuid.UUID UserID string AfterID string PageSize int32 @@ -459,7 +461,7 @@ FOR UPDATE type LockReadyAgentInstanceCheckpointParams struct { Namespace string - ID string + ID uuid.UUID UserID string } diff --git a/go/core/internal/database/gen/agent_instance_tasks.sql.go b/go/core/internal/database/gen/agent_instance_tasks.sql.go index 7068230ae..3da9142fe 100644 --- a/go/core/internal/database/gen/agent_instance_tasks.sql.go +++ b/go/core/internal/database/gen/agent_instance_tasks.sql.go @@ -8,6 +8,8 @@ package dbgen import ( "context" "time" + + "github.com/google/uuid" ) const countAgentInstanceTasks = `-- name: CountAgentInstanceTasks :one @@ -19,7 +21,7 @@ WHERE context_id = $1 ` type CountAgentInstanceTasksParams struct { - ContextID string + ContextID uuid.UUID State string StatusTimestampAfter *time.Time } @@ -46,7 +48,7 @@ DO NOTHING ` type CreateAgentInstanceTaskParams struct { - ContextID string + ContextID uuid.UUID ID string State string StatusTimestamp *time.Time @@ -84,7 +86,7 @@ WHERE context_id = $1 ) ` -func (q *Queries) GetActiveAgentInstanceTask(ctx context.Context, contextID string) (AgentInstanceTask, error) { +func (q *Queries) GetActiveAgentInstanceTask(ctx context.Context, contextID uuid.UUID) (AgentInstanceTask, error) { row := q.db.QueryRow(ctx, getActiveAgentInstanceTask, contextID) var i AgentInstanceTask err := row.Scan( @@ -112,7 +114,7 @@ WHERE context_id = $1 AND id = $2 ` type GetAgentInstanceTaskParams struct { - ContextID string + ContextID uuid.UUID ID string } @@ -144,7 +146,7 @@ WHERE context_id = $1 AND initial_message_id = $2 ` type GetAgentInstanceTaskByMessageIDParams struct { - ContextID string + ContextID uuid.UUID InitialMessageID *string } @@ -187,7 +189,7 @@ LIMIT 1 ` type InsertAgentInstanceTaskEventParams struct { - ContextID string + ContextID uuid.UUID TaskID *string MessageID *string Data []byte @@ -214,7 +216,7 @@ INSERT INTO agent_instance_task ( ` type InsertCopiedAgentInstanceTaskParams struct { - ContextID string + ContextID uuid.UUID ID string State string StatusTimestamp *time.Time @@ -260,7 +262,7 @@ ORDER BY sequence ` type ListAgentInstanceTaskHistoryParams struct { - ContextID string + ContextID uuid.UUID TaskIds []string } @@ -301,7 +303,7 @@ LIMIT $5 ` type ListAgentInstanceTasksParams struct { - ContextID string + ContextID uuid.UUID AfterID string State string StatusTimestampAfter *time.Time @@ -365,7 +367,7 @@ FOR UPDATE // LockActiveAgentInstanceTask holds the instance's non-terminal task for the // rest of the transaction so reclamation cannot overwrite concurrent progress. -func (q *Queries) LockActiveAgentInstanceTask(ctx context.Context, contextID string) (AgentInstanceTask, error) { +func (q *Queries) LockActiveAgentInstanceTask(ctx context.Context, contextID uuid.UUID) (AgentInstanceTask, error) { row := q.db.QueryRow(ctx, lockActiveAgentInstanceTask, contextID) var i AgentInstanceTask err := row.Scan( @@ -398,7 +400,7 @@ WHERE context_id = $1 AND id = $2 ` type SetAgentInstanceTaskSnapshotParams struct { - ContextID string + ContextID uuid.UUID ID string SnapshotAtespace *string SnapshotName *string @@ -431,7 +433,7 @@ ON CONFLICT (context_id, id) DO UPDATE SET ` type UpsertAgentInstanceTaskParams struct { - ContextID string + ContextID uuid.UUID ID string State string StatusTimestamp *time.Time diff --git a/go/core/internal/database/gen/agent_instances.sql.go b/go/core/internal/database/gen/agent_instances.sql.go index 5d6473421..5971a8a52 100644 --- a/go/core/internal/database/gen/agent_instances.sql.go +++ b/go/core/internal/database/gen/agent_instances.sql.go @@ -8,6 +8,8 @@ package dbgen import ( "context" "time" + + "github.com/google/uuid" ) const createAgentInstanceShare = `-- name: CreateAgentInstanceShare :one @@ -18,9 +20,9 @@ RETURNING id, namespace, instance_id, permission, token_hash, created_at ` type CreateAgentInstanceShareParams struct { - ID string + ID uuid.UUID Namespace string - InstanceID string + InstanceID uuid.UUID Permission string TokenHash []byte } @@ -49,7 +51,7 @@ const deleteAgentInstance = `-- name: DeleteAgentInstance :exec DELETE FROM agent_instance WHERE id = $1 ` -func (q *Queries) DeleteAgentInstance(ctx context.Context, id string) error { +func (q *Queries) DeleteAgentInstance(ctx context.Context, id uuid.UUID) error { _, err := q.db.Exec(ctx, deleteAgentInstance, id) return err } @@ -63,7 +65,7 @@ WHERE s.namespace = $1 AND s.id = $2 type DeleteAgentInstanceShareParams struct { Namespace string - ID string + ID uuid.UUID UserID string } @@ -79,7 +81,7 @@ const getAgentInstanceByID = `-- name: GetAgentInstanceByID :one SELECT id, namespace, user_id, request_id, prepared_revision, state, labels, data, operation, context_id, source_checkpoint_id, name FROM agent_instance WHERE id = $1 ` -func (q *Queries) GetAgentInstanceByID(ctx context.Context, id string) (AgentInstance, error) { +func (q *Queries) GetAgentInstanceByID(ctx context.Context, id uuid.UUID) (AgentInstance, error) { row := q.db.QueryRow(ctx, getAgentInstanceByID, id) var i AgentInstance err := row.Scan( @@ -136,7 +138,7 @@ SELECT id, namespace, user_id, request_id, prepared_revision, state, labels, dat type GetAgentInstanceForUserParams struct { Namespace string - ID string + ID uuid.UUID UserID string } @@ -168,9 +170,9 @@ WHERE s.token_hash = $1 ` type GetAgentInstanceShareByTokenHashRow struct { - ID string + ID uuid.UUID Namespace string - InstanceID string + InstanceID uuid.UUID Permission string TokenHash []byte CreatedAt time.Time @@ -265,7 +267,7 @@ VALUES ($1, $2, $3) ` type InsertA2AContextParams struct { - ID string + ID uuid.UUID Namespace string UserID string } @@ -284,11 +286,11 @@ RETURNING id, namespace, user_id, request_id, prepared_revision, state, labels, ` type InsertAgentInstanceParams struct { - ID string + ID uuid.UUID Namespace string UserID string RequestID string - ContextID string + ContextID uuid.UUID PreparedRevision *string Labels []byte Name string @@ -335,13 +337,13 @@ RETURNING id, namespace, user_id, request_id, prepared_revision, state, labels, ` type InsertForkedAgentInstanceParams struct { - ID string + ID uuid.UUID Namespace string UserID string RequestID string - ContextID string + ContextID uuid.UUID PreparedRevision *string - SourceCheckpointID *string + SourceCheckpointID *uuid.UUID Labels []byte Data []byte } @@ -380,14 +382,14 @@ const listAgentInstanceShares = `-- name: ListAgentInstanceShares :many SELECT s.id, s.namespace, s.instance_id, s.permission, s.token_hash, s.created_at FROM agent_instance_share s JOIN agent_instance i ON i.id = s.instance_id WHERE s.namespace = $1 AND s.instance_id = $2 AND i.user_id = $3 - AND s.id > $4 + AND (NULLIF($4::text, '') IS NULL OR s.id > NULLIF($4::text, '')::uuid) ORDER BY s.id LIMIT $5 ` type ListAgentInstanceSharesParams struct { Namespace string - InstanceID string + InstanceID uuid.UUID UserID string AfterID string PageSize int32 @@ -431,7 +433,7 @@ SELECT i.id, i.namespace, i.user_id, i.request_id, i.prepared_revision, i.state, LEFT JOIN runtime_revision r ON r.revision = i.prepared_revision WHERE i.namespace = $1 AND ($2::boolean OR i.user_id = $3) - AND i.id > $4 + AND (NULLIF($4::text, '') IS NULL OR i.id > NULLIF($4::text, '')::uuid) AND i.labels @> $5::jsonb AND ($6::text = '' OR r.agent_template_name = $6) AND ($7::text = '' OR r.harness_name = $7) @@ -504,7 +506,7 @@ const lockAgentInstance = `-- name: LockAgentInstance :one SELECT id, namespace, user_id, request_id, prepared_revision, state, labels, data, operation, context_id, source_checkpoint_id, name FROM agent_instance WHERE id = $1 FOR UPDATE ` -func (q *Queries) LockAgentInstance(ctx context.Context, id string) (AgentInstance, error) { +func (q *Queries) LockAgentInstance(ctx context.Context, id uuid.UUID) (AgentInstance, error) { row := q.db.QueryRow(ctx, lockAgentInstance, id) var i AgentInstance err := row.Scan( @@ -532,7 +534,7 @@ RETURNING id, namespace, user_id, request_id, prepared_revision, state, labels, ` type MarkAgentInstanceReadyParams struct { - ID string + ID uuid.UUID Data []byte } @@ -576,7 +578,7 @@ type TransitionAgentInstanceParams struct { NextState string NextOperation string Data []byte - ID string + ID uuid.UUID ExpectedState string ExpectedOperation string } @@ -618,7 +620,7 @@ RETURNING id, namespace, user_id, request_id, prepared_revision, state, labels, type UpdateAgentInstanceNameParams struct { Name string Namespace string - ID string + ID uuid.UUID UserID string } diff --git a/go/core/internal/database/gen/models.go b/go/core/internal/database/gen/models.go index 51592d398..a1135fb56 100644 --- a/go/core/internal/database/gen/models.go +++ b/go/core/internal/database/gen/models.go @@ -7,13 +7,14 @@ package dbgen import ( "time" + "github.com/google/uuid" "github.com/kagent-dev/kagent/go/api/adk" "github.com/kagent-dev/kagent/go/api/database" pgvector_go "github.com/pgvector/pgvector-go" ) type A2aContext struct { - ID string + ID uuid.UUID Namespace string UserID string CreatedAt time.Time @@ -30,7 +31,7 @@ type Agent struct { } type AgentInstance struct { - ID string + ID uuid.UUID Namespace string UserID string RequestID string @@ -39,15 +40,15 @@ type AgentInstance struct { Labels []byte Data []byte Operation string - ContextID string - SourceCheckpointID *string + ContextID uuid.UUID + SourceCheckpointID *uuid.UUID Name string } type AgentInstanceCheckpoint struct { - ID string + ID uuid.UUID Namespace string - SourceInstanceID string + SourceInstanceID uuid.UUID UserID string RequestID string HeadTaskID string @@ -60,22 +61,22 @@ type AgentInstanceCheckpoint struct { State string Failure string CreatedAt time.Time - SourceContextID string + SourceContextID uuid.UUID PreparedRevision *string SourceLabels []byte } type AgentInstanceShare struct { - ID string + ID uuid.UUID Namespace string - InstanceID string + InstanceID uuid.UUID Permission string TokenHash []byte CreatedAt time.Time } type AgentInstanceTask struct { - ContextID string + ContextID uuid.UUID ID string State string StatusTimestamp *time.Time @@ -93,7 +94,7 @@ type AgentInstanceTask struct { type AgentInstanceTaskEvent struct { Sequence int64 - ContextID string + ContextID uuid.UUID TaskID *string Data []byte CreatedAt time.Time diff --git a/go/core/internal/database/gen/querier.go b/go/core/internal/database/gen/querier.go index 2b471f13a..7612f8444 100644 --- a/go/core/internal/database/gen/querier.go +++ b/go/core/internal/database/gen/querier.go @@ -6,6 +6,8 @@ package dbgen import ( "context" + + "github.com/google/uuid" ) type Querier interface { @@ -13,7 +15,7 @@ type Querier interface { CountAgentInstanceTasks(ctx context.Context, arg CountAgentInstanceTasksParams) (int64, error) CreateAgentInstanceShare(ctx context.Context, arg CreateAgentInstanceShareParams) (AgentInstanceShare, error) CreateAgentInstanceTask(ctx context.Context, arg CreateAgentInstanceTaskParams) (int64, error) - DeleteAgentInstance(ctx context.Context, id string) error + DeleteAgentInstance(ctx context.Context, id uuid.UUID) error DeleteAgentInstanceCheckpoint(ctx context.Context, arg DeleteAgentInstanceCheckpointParams) (int64, error) DeleteAgentInstanceShare(ctx context.Context, arg DeleteAgentInstanceShareParams) (int64, error) DeleteAgentMemory(ctx context.Context, arg DeleteAgentMemoryParams) error @@ -21,9 +23,9 @@ type Querier interface { DeleteUnreferencedRuntimeRevision(ctx context.Context, revision string) error ExtendMemoryTTL(ctx context.Context) error FinalizeAgentInstanceCheckpoint(ctx context.Context, arg FinalizeAgentInstanceCheckpointParams) (AgentInstanceCheckpoint, error) - GetActiveAgentInstanceTask(ctx context.Context, contextID string) (AgentInstanceTask, error) + GetActiveAgentInstanceTask(ctx context.Context, contextID uuid.UUID) (AgentInstanceTask, error) GetAgent(ctx context.Context, id string) (Agent, error) - GetAgentInstanceByID(ctx context.Context, id string) (AgentInstance, error) + GetAgentInstanceByID(ctx context.Context, id uuid.UUID) (AgentInstance, error) GetAgentInstanceByRequest(ctx context.Context, arg GetAgentInstanceByRequestParams) (AgentInstance, error) GetAgentInstanceCheckpoint(ctx context.Context, arg GetAgentInstanceCheckpointParams) (AgentInstanceCheckpoint, error) GetAgentInstanceCheckpointByRequest(ctx context.Context, arg GetAgentInstanceCheckpointByRequestParams) (AgentInstanceCheckpoint, error) @@ -39,7 +41,7 @@ type Querier interface { GetAgentInstanceTaskByMessageID(ctx context.Context, arg GetAgentInstanceTaskByMessageIDParams) (AgentInstanceTask, error) GetCheckpoint(ctx context.Context, arg GetCheckpointParams) (LgCheckpoint, error) GetLatestCrewAIFlowState(ctx context.Context, arg GetLatestCrewAIFlowStateParams) (CrewaiFlowState, error) - GetLatestQuiescentAgentInstanceTask(ctx context.Context, contextID string) (AgentInstanceTask, error) + GetLatestQuiescentAgentInstanceTask(ctx context.Context, contextID uuid.UUID) (AgentInstanceTask, error) GetLatestRuntimeRevisionForInstance(ctx context.Context, arg GetLatestRuntimeRevisionForInstanceParams) (GetLatestRuntimeRevisionForInstanceRow, error) GetRuntimeRevision(ctx context.Context, revision string) (RuntimeRevision, error) GetTool(ctx context.Context, id string) (Tool, error) @@ -55,8 +57,8 @@ type Querier interface { InsertFeedback(ctx context.Context, arg InsertFeedbackParams) error InsertForkedAgentInstance(ctx context.Context, arg InsertForkedAgentInstanceParams) (AgentInstance, error) InsertMemory(ctx context.Context, arg InsertMemoryParams) (string, error) - ListAgentInstanceCheckpointEvents(ctx context.Context, checkpointID string) ([]AgentInstanceTaskEvent, error) - ListAgentInstanceCheckpointTasks(ctx context.Context, checkpointID string) ([]AgentInstanceTask, error) + ListAgentInstanceCheckpointEvents(ctx context.Context, checkpointID uuid.UUID) ([]AgentInstanceTaskEvent, error) + ListAgentInstanceCheckpointTasks(ctx context.Context, checkpointID uuid.UUID) ([]AgentInstanceTask, error) ListAgentInstanceCheckpoints(ctx context.Context, arg ListAgentInstanceCheckpointsParams) ([]AgentInstanceCheckpoint, error) ListAgentInstanceShares(ctx context.Context, arg ListAgentInstanceSharesParams) ([]AgentInstanceShare, error) ListAgentInstanceTaskHistory(ctx context.Context, arg ListAgentInstanceTaskHistoryParams) ([]ListAgentInstanceTaskHistoryRow, error) @@ -82,8 +84,8 @@ type Querier interface { ListUnreferencedRuntimeRevisions(ctx context.Context) ([]RuntimeRevision, error) // LockActiveAgentInstanceTask holds the instance's non-terminal task for the // rest of the transaction so reclamation cannot overwrite concurrent progress. - LockActiveAgentInstanceTask(ctx context.Context, contextID string) (AgentInstanceTask, error) - LockAgentInstance(ctx context.Context, id string) (AgentInstance, error) + LockActiveAgentInstanceTask(ctx context.Context, contextID uuid.UUID) (AgentInstanceTask, error) + LockAgentInstance(ctx context.Context, id uuid.UUID) (AgentInstance, error) LockReadyAgentInstanceCheckpoint(ctx context.Context, arg LockReadyAgentInstanceCheckpointParams) (AgentInstanceCheckpoint, error) MarkAgentInstanceReady(ctx context.Context, arg MarkAgentInstanceReadyParams) (AgentInstance, error) MarkRuntimeRevisionSuccessful(ctx context.Context, arg MarkRuntimeRevisionSuccessfulParams) error diff --git a/go/core/internal/database/queries/agent_instance_checkpoints.sql b/go/core/internal/database/queries/agent_instance_checkpoints.sql index 02d61b8ff..06a294e3c 100644 --- a/go/core/internal/database/queries/agent_instance_checkpoints.sql +++ b/go/core/internal/database/queries/agent_instance_checkpoints.sql @@ -75,7 +75,7 @@ WHERE namespace = sqlc.arg(namespace) AND source_instance_id = sqlc.arg(source_instance_id) AND user_id = sqlc.arg(user_id) AND state = 'READY' - AND id > sqlc.arg(after_id) + AND (NULLIF(sqlc.arg(after_id)::text, '') IS NULL OR id > NULLIF(sqlc.arg(after_id)::text, '')::uuid) ORDER BY id LIMIT sqlc.arg(page_size); diff --git a/go/core/internal/database/queries/agent_instances.sql b/go/core/internal/database/queries/agent_instances.sql index 8c783b041..12021c5df 100644 --- a/go/core/internal/database/queries/agent_instances.sql +++ b/go/core/internal/database/queries/agent_instances.sql @@ -52,7 +52,7 @@ SELECT i.* FROM agent_instance i LEFT JOIN runtime_revision r ON r.revision = i.prepared_revision WHERE i.namespace = sqlc.arg(namespace) AND (sqlc.arg(all_users)::boolean OR i.user_id = sqlc.arg(user_id)) - AND i.id > sqlc.arg(after_id) + AND (NULLIF(sqlc.arg(after_id)::text, '') IS NULL OR i.id > NULLIF(sqlc.arg(after_id)::text, '')::uuid) AND i.labels @> sqlc.arg(match_labels)::jsonb AND (sqlc.arg(agent_template)::text = '' OR r.agent_template_name = sqlc.arg(agent_template)) AND (sqlc.arg(harness)::text = '' OR r.harness_name = sqlc.arg(harness)) @@ -115,7 +115,7 @@ WHERE s.token_hash = $1; SELECT s.* FROM agent_instance_share s JOIN agent_instance i ON i.id = s.instance_id WHERE s.namespace = $1 AND s.instance_id = $2 AND i.user_id = $3 - AND s.id > sqlc.arg(after_id) + AND (NULLIF(sqlc.arg(after_id)::text, '') IS NULL OR s.id > NULLIF(sqlc.arg(after_id)::text, '')::uuid) ORDER BY s.id LIMIT sqlc.arg(page_size); diff --git a/go/core/internal/database/sqlc.yaml b/go/core/internal/database/sqlc.yaml index 44c25c9eb..267beb3a6 100644 --- a/go/core/internal/database/sqlc.yaml +++ b/go/core/internal/database/sqlc.yaml @@ -30,6 +30,17 @@ sql: go_type: type: "time.Time" pointer: true + - db_type: "uuid" + nullable: false + go_type: + import: "github.com/google/uuid" + type: "UUID" + - db_type: "uuid" + nullable: true + go_type: + import: "github.com/google/uuid" + type: "UUID" + pointer: true # Use domain types for columns that would otherwise be plain strings/bytes. - column: "agent.config" go_type: diff --git a/go/core/internal/grpcserver/interceptors.go b/go/core/internal/grpcserver/interceptors.go index 82e50c229..c83eb2585 100644 --- a/go/core/internal/grpcserver/interceptors.go +++ b/go/core/internal/grpcserver/interceptors.go @@ -98,7 +98,7 @@ func authenticate(ctx context.Context, fullMethod string, authenticator auth.Aut // to what the owner can see, and the instance read runs as the owner. UserID: instanceShare.OwnerUserID, ReadOnly: readOnly, - AgentInstanceID: instanceShare.InstanceID, + AgentInstanceID: instanceShare.InstanceID.String(), }), nil } diff --git a/go/core/internal/grpcserver/interceptors_test.go b/go/core/internal/grpcserver/interceptors_test.go index 739d65efe..3c15283f2 100644 --- a/go/core/internal/grpcserver/interceptors_test.go +++ b/go/core/internal/grpcserver/interceptors_test.go @@ -7,6 +7,7 @@ import ( "net/url" "testing" + "github.com/google/uuid" dbpkg "github.com/kagent-dev/kagent/go/api/database" "github.com/kagent-dev/kagent/go/core/internal/service/serviceerrors" pkgauth "github.com/kagent-dev/kagent/go/core/pkg/auth" @@ -22,6 +23,8 @@ const ( createMethod = "/test.Service/Create" ) +var testInstanceID = uuid.MustParse("22222222-2222-4222-8222-222222222222") + type testSession struct { principal pkgauth.Principal } @@ -124,7 +127,7 @@ func TestAuthenticationUnaryInterceptor(t *testing.T) { t.Run("an AgentInstance share is attached to a read call", func(t *testing.T) { store := &testShareStore{ instanceShare: &dbpkg.AgentInstanceShare{ - ID: "share-1", Namespace: "kagent", InstanceID: "instance-1", + Namespace: "kagent", InstanceID: testInstanceID, Permission: "READ_ONLY", OwnerUserID: "owner", }, } @@ -136,7 +139,7 @@ func TestAuthenticationUnaryInterceptor(t *testing.T) { if !ok { t.Fatal("no share context") } - if !share.IsForAgentInstance("instance-1") { + if !share.IsForAgentInstance(testInstanceID.String()) { t.Errorf("share is not for instance-1: %#v", share) } // The owner, not the visitor: the instance read runs as the owner or @@ -158,7 +161,7 @@ func TestAuthenticationUnaryInterceptor(t *testing.T) { t.Run("a read-only AgentInstance share cannot send", func(t *testing.T) { store := &testShareStore{ instanceShare: &dbpkg.AgentInstanceShare{ - InstanceID: "instance-1", Permission: "READ_ONLY", OwnerUserID: "owner", + InstanceID: testInstanceID, Permission: "READ_ONLY", OwnerUserID: "owner", }, } ctx := metadata.NewIncomingContext(t.Context(), metadata.Pairs("x-share-token", "share")) @@ -177,7 +180,7 @@ func TestAuthenticationUnaryInterceptor(t *testing.T) { t.Run("a READ_WRITE AgentInstance share may send", func(t *testing.T) { store := &testShareStore{ instanceShare: &dbpkg.AgentInstanceShare{ - InstanceID: "instance-1", Permission: "READ_WRITE", OwnerUserID: "owner", + InstanceID: testInstanceID, Permission: "READ_WRITE", OwnerUserID: "owner", }, } ctx := metadata.NewIncomingContext(t.Context(), metadata.Pairs("x-share-token", "share")) diff --git a/go/core/internal/grpcserver/policy_test.go b/go/core/internal/grpcserver/policy_test.go index 1a9c67e50..2f5584060 100644 --- a/go/core/internal/grpcserver/policy_test.go +++ b/go/core/internal/grpcserver/policy_test.go @@ -74,7 +74,7 @@ func TestReadOnlyShareCannotRenameAConversation(t *testing.T) { t.Run(test.name, func(t *testing.T) { shareStore := &testShareStore{ instanceShare: &dbpkg.AgentInstanceShare{ - ID: "share-1", InstanceID: "instance-1", Permission: test.permission, OwnerUserID: "owner", + InstanceID: testInstanceID, Permission: test.permission, OwnerUserID: "owner", }, } ctx := metadata.NewIncomingContext(t.Context(), metadata.Pairs("x-share-token", "token")) diff --git a/go/core/pkg/migrations/core/000018_agent_instance_uuid_ids.down.sql b/go/core/pkg/migrations/core/000018_agent_instance_uuid_ids.down.sql new file mode 100644 index 000000000..b0a46cbb0 --- /dev/null +++ b/go/core/pkg/migrations/core/000018_agent_instance_uuid_ids.down.sql @@ -0,0 +1,34 @@ +ALTER TABLE agent_instance_share DROP CONSTRAINT agent_instance_share_instance_id_fkey; +ALTER TABLE agent_instance_task DROP CONSTRAINT agent_instance_task_context_id_fkey; +ALTER TABLE agent_instance_task_event DROP CONSTRAINT agent_instance_task_event_context_id_fkey; +ALTER TABLE agent_instance_checkpoint DROP CONSTRAINT agent_instance_checkpoint_source_context_id_fkey; +ALTER TABLE agent_instance DROP CONSTRAINT agent_instance_context_id_fkey; +ALTER TABLE agent_instance DROP CONSTRAINT agent_instance_source_checkpoint_id_fkey; + +ALTER TABLE a2a_context ALTER COLUMN id TYPE TEXT USING id::text; +ALTER TABLE agent_instance + ALTER COLUMN id TYPE TEXT USING id::text, + ALTER COLUMN context_id TYPE TEXT USING context_id::text, + ALTER COLUMN source_checkpoint_id TYPE TEXT USING source_checkpoint_id::text; +ALTER TABLE agent_instance_share + ALTER COLUMN id TYPE TEXT USING id::text, + ALTER COLUMN instance_id TYPE TEXT USING instance_id::text; +ALTER TABLE agent_instance_task ALTER COLUMN context_id TYPE TEXT USING context_id::text; +ALTER TABLE agent_instance_task_event ALTER COLUMN context_id TYPE TEXT USING context_id::text; +ALTER TABLE agent_instance_checkpoint + ALTER COLUMN id TYPE TEXT USING id::text, + ALTER COLUMN source_instance_id TYPE TEXT USING source_instance_id::text, + ALTER COLUMN source_context_id TYPE TEXT USING source_context_id::text; + +ALTER TABLE agent_instance_share ADD CONSTRAINT agent_instance_share_instance_id_fkey + FOREIGN KEY (instance_id) REFERENCES agent_instance(id) ON DELETE CASCADE; +ALTER TABLE agent_instance_task ADD CONSTRAINT agent_instance_task_context_id_fkey + FOREIGN KEY (context_id) REFERENCES a2a_context(id) ON DELETE CASCADE; +ALTER TABLE agent_instance_task_event ADD CONSTRAINT agent_instance_task_event_context_id_fkey + FOREIGN KEY (context_id) REFERENCES a2a_context(id) ON DELETE CASCADE; +ALTER TABLE agent_instance_checkpoint ADD CONSTRAINT agent_instance_checkpoint_source_context_id_fkey + FOREIGN KEY (source_context_id) REFERENCES a2a_context(id) ON DELETE RESTRICT; +ALTER TABLE agent_instance ADD CONSTRAINT agent_instance_context_id_fkey + FOREIGN KEY (context_id) REFERENCES a2a_context(id) ON DELETE RESTRICT; +ALTER TABLE agent_instance ADD CONSTRAINT agent_instance_source_checkpoint_id_fkey + FOREIGN KEY (source_checkpoint_id) REFERENCES agent_instance_checkpoint(id) ON DELETE RESTRICT; diff --git a/go/core/pkg/migrations/core/000018_agent_instance_uuid_ids.up.sql b/go/core/pkg/migrations/core/000018_agent_instance_uuid_ids.up.sql new file mode 100644 index 000000000..50de0fc58 --- /dev/null +++ b/go/core/pkg/migrations/core/000018_agent_instance_uuid_ids.up.sql @@ -0,0 +1,34 @@ +ALTER TABLE agent_instance_share DROP CONSTRAINT agent_instance_share_instance_id_fkey; +ALTER TABLE agent_instance_task DROP CONSTRAINT agent_instance_task_context_id_fkey; +ALTER TABLE agent_instance_task_event DROP CONSTRAINT agent_instance_task_event_context_id_fkey; +ALTER TABLE agent_instance_checkpoint DROP CONSTRAINT agent_instance_checkpoint_source_context_id_fkey; +ALTER TABLE agent_instance DROP CONSTRAINT agent_instance_context_id_fkey; +ALTER TABLE agent_instance DROP CONSTRAINT agent_instance_source_checkpoint_id_fkey; + +ALTER TABLE a2a_context ALTER COLUMN id TYPE UUID USING id::uuid; +ALTER TABLE agent_instance + ALTER COLUMN id TYPE UUID USING id::uuid, + ALTER COLUMN context_id TYPE UUID USING context_id::uuid, + ALTER COLUMN source_checkpoint_id TYPE UUID USING source_checkpoint_id::uuid; +ALTER TABLE agent_instance_share + ALTER COLUMN id TYPE UUID USING id::uuid, + ALTER COLUMN instance_id TYPE UUID USING instance_id::uuid; +ALTER TABLE agent_instance_task ALTER COLUMN context_id TYPE UUID USING context_id::uuid; +ALTER TABLE agent_instance_task_event ALTER COLUMN context_id TYPE UUID USING context_id::uuid; +ALTER TABLE agent_instance_checkpoint + ALTER COLUMN id TYPE UUID USING id::uuid, + ALTER COLUMN source_instance_id TYPE UUID USING source_instance_id::uuid, + ALTER COLUMN source_context_id TYPE UUID USING source_context_id::uuid; + +ALTER TABLE agent_instance_share ADD CONSTRAINT agent_instance_share_instance_id_fkey + FOREIGN KEY (instance_id) REFERENCES agent_instance(id) ON DELETE CASCADE; +ALTER TABLE agent_instance_task ADD CONSTRAINT agent_instance_task_context_id_fkey + FOREIGN KEY (context_id) REFERENCES a2a_context(id) ON DELETE CASCADE; +ALTER TABLE agent_instance_task_event ADD CONSTRAINT agent_instance_task_event_context_id_fkey + FOREIGN KEY (context_id) REFERENCES a2a_context(id) ON DELETE CASCADE; +ALTER TABLE agent_instance_checkpoint ADD CONSTRAINT agent_instance_checkpoint_source_context_id_fkey + FOREIGN KEY (source_context_id) REFERENCES a2a_context(id) ON DELETE RESTRICT; +ALTER TABLE agent_instance ADD CONSTRAINT agent_instance_context_id_fkey + FOREIGN KEY (context_id) REFERENCES a2a_context(id) ON DELETE RESTRICT; +ALTER TABLE agent_instance ADD CONSTRAINT agent_instance_source_checkpoint_id_fkey + FOREIGN KEY (source_checkpoint_id) REFERENCES agent_instance_checkpoint(id) ON DELETE RESTRICT; diff --git a/go/core/v2/agentinstance/grpc.go b/go/core/v2/agentinstance/grpc.go index 4ba04e2fe..8245bfacf 100644 --- a/go/core/v2/agentinstance/grpc.go +++ b/go/core/v2/agentinstance/grpc.go @@ -112,7 +112,7 @@ func (s *grpcServer) RevokeAgentInstanceShare(ctx context.Context, request *apiv func agentInstanceShareProto(share *dbpkg.AgentInstanceShare) *apiv1alpha1.AgentInstanceShare { return &apiv1alpha1.AgentInstanceShare{ - Id: share.ID, Namespace: share.Namespace, AgentInstanceId: share.InstanceID, + Id: share.ID.String(), Namespace: share.Namespace, AgentInstanceId: share.InstanceID.String(), Permission: agentInstanceSharePermission(share.Permission), CreatedAt: timestamppb.New(share.CreatedAt), } } diff --git a/go/core/v2/agentinstance/service.go b/go/core/v2/agentinstance/service.go index 9712d7f60..621097079 100644 --- a/go/core/v2/agentinstance/service.go +++ b/go/core/v2/agentinstance/service.go @@ -278,7 +278,7 @@ func (s *Service) CreateShare(ctx context.Context, namespace, instanceID, permis return nil, "", serviceerrors.NewInternal("Failed to create share token", err) } share, err := s.store.CreateAgentInstanceShare(ctx, dbpkg.AgentInstanceShare{ - ID: uuid.NewString(), Namespace: namespace, InstanceID: instanceID, + ID: uuid.New(), Namespace: namespace, InstanceID: uuid.MustParse(instanceID), Permission: permission, TokenHash: tokenHash, }) if err != nil { @@ -311,7 +311,7 @@ func (s *Service) ListShares(ctx context.Context, namespace, instanceID string, } result := ShareListResult{Shares: shares} if len(result.Shares) > pageSize { - result.NextPageToken = encodePageToken(result.Shares[pageSize-1].ID) + result.NextPageToken = encodePageToken(result.Shares[pageSize-1].ID.String()) result.Shares = result.Shares[:pageSize] } return result, nil diff --git a/go/core/v2/agentinstance/service_test.go b/go/core/v2/agentinstance/service_test.go index cc1feee6d..933157399 100644 --- a/go/core/v2/agentinstance/service_test.go +++ b/go/core/v2/agentinstance/service_test.go @@ -222,7 +222,7 @@ func TestServiceCreateShareGeneratesTokenAndUUID(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err := uuid.Parse(share.ID); err != nil { + if share.ID == uuid.Nil { t.Fatalf("generated share id %q is not a UUID: %v", share.ID, err) } digest := sha256.Sum256([]byte(token)) @@ -238,7 +238,9 @@ func TestServiceListSharesPaginatesInStore(t *testing.T) { "33333333-3333-4333-8333-333333333333", "44444444-4444-4444-8444-444444444444", } - store := &serviceTestStore{shares: []dbpkg.AgentInstanceShare{{ID: ids[1]}, {ID: ids[2]}, {ID: ids[3]}}} + store := &serviceTestStore{shares: []dbpkg.AgentInstanceShare{ + {ID: uuid.MustParse(ids[1])}, {ID: uuid.MustParse(ids[2])}, {ID: uuid.MustParse(ids[3])}, + }} service := NewService(store, serviceTestAuthorizer{}, serviceTestWorkflow{}) result, err := service.ListShares(serviceTestContext("alice"), "team-a", ids[0], 2, encodePageToken(ids[0])) if err != nil { diff --git a/go/core/v2/agentinstance/workflow.go b/go/core/v2/agentinstance/workflow.go index 153b38892..e1544a220 100644 --- a/go/core/v2/agentinstance/workflow.go +++ b/go/core/v2/agentinstance/workflow.go @@ -134,7 +134,7 @@ func (w *ActorWorkflow) Fork(ctx context.Context, instance *apiv1alpha1.AgentIns if err := w.actors.EnsureAtespace(ctx, atespace); err != nil { return nil, fmt.Errorf("ensure Atespace %s: %w", atespace, err) } - tag := &ateapipb.ObjectRef{Atespace: checkpoint.SnapshotAtespace, Name: "checkpoint-" + checkpoint.ID} + tag := &ateapipb.ObjectRef{Atespace: checkpoint.SnapshotAtespace, Name: "checkpoint-" + checkpoint.ID.String()} actor, err := w.actors.GetActor(ctx, atespace, name) if status.Code(err) == codes.NotFound { actor, err = w.actors.CreateActorFromSnapshotTag(ctx, atespace, name, diff --git a/go/core/v2/agentinstance/workflow_test.go b/go/core/v2/agentinstance/workflow_test.go index f9c9f1da0..afe3b38d7 100644 --- a/go/core/v2/agentinstance/workflow_test.go +++ b/go/core/v2/agentinstance/workflow_test.go @@ -5,6 +5,7 @@ import ( "testing" "github.com/agent-substrate/substrate/pkg/proto/ateapipb" + "github.com/google/uuid" dbpkg "github.com/kagent-dev/kagent/go/api/database" apiv1alpha1 "github.com/kagent-dev/kagent/go/api/gen/kagent/api/v1alpha1" "google.golang.org/grpc/codes" @@ -88,7 +89,7 @@ func TestActorWorkflowForkCreatesSuspendedActorFromCheckpoint(t *testing.T) { } actors := &lifecycleTestActors{actors: map[string]*ateapipb.Actor{}} checkpoint := &dbpkg.AgentInstanceCheckpoint{ - ID: "checkpoint-1", SnapshotAtespace: "team-a", SnapshotName: "snapshot-1", SnapshotUID: "snapshot-uid", + ID: uuid.MustParse("018f47a2-4efb-7c21-a848-123456789abc"), SnapshotAtespace: "team-a", SnapshotName: "snapshot-1", SnapshotUID: "snapshot-uid", } fork, err := NewActorWorkflow(store, actors).Fork(context.Background(), instance, checkpoint) if err != nil { @@ -97,7 +98,7 @@ func TestActorWorkflowForkCreatesSuspendedActorFromCheckpoint(t *testing.T) { actor := actors.actors[actorKey("team-a", actorName(instance.GetId()))] if fork.GetState() != apiv1alpha1.AgentInstanceState_AGENT_INSTANCE_STATE_READY || actor.GetStatus().GetState() != ateapipb.ActorState_ACTOR_STATE_SUSPENDED || - actor.GetSourceSnapshotTag().GetName() != "checkpoint-checkpoint-1" { + actor.GetSourceSnapshotTag().GetName() != "checkpoint-018f47a2-4efb-7c21-a848-123456789abc" { t.Fatalf("fork = %+v, actor = %+v", fork, actor) } instance.State = apiv1alpha1.AgentInstanceState_AGENT_INSTANCE_STATE_CREATING diff --git a/go/core/v2/checkpoint/service.go b/go/core/v2/checkpoint/service.go index cb582408f..b87175962 100644 --- a/go/core/v2/checkpoint/service.go +++ b/go/core/v2/checkpoint/service.go @@ -80,8 +80,9 @@ func (s *Service) Create(ctx context.Context, namespace, instanceID, requestID s if err != nil { return nil, serviceerrors.NewInternal("Failed to generate checkpoint identifier", err) } + instanceUUID := uuid.MustParse(instanceID) checkpoint, err := s.store.ReserveAgentInstanceCheckpoint(ctx, dbpkg.AgentInstanceCheckpoint{ - ID: id.String(), Namespace: namespace, SourceInstanceID: instanceID, UserID: userID, + ID: id, Namespace: namespace, SourceInstanceID: instanceUUID, UserID: userID, RequestID: requestID, }) if errors.Is(err, dbpkg.ErrIdempotencyConflict) { @@ -102,13 +103,13 @@ func (s *Service) Create(ctx context.Context, namespace, instanceID, requestID s tag, err := s.ensureTag(ctx, checkpoint) if err != nil { - cleanupErr := s.tags.DeleteActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, tagName(checkpoint.ID)) + cleanupErr := s.tags.DeleteActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, tagName(checkpoint.ID.String())) if cleanupErr == nil || status.Code(cleanupErr) == codes.NotFound { - _, _ = s.store.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID, "", err.Error()) + _, _ = s.store.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID.String(), "", err.Error()) } return nil, serviceerrors.NewUnavailable("Failed to retain checkpoint snapshot", err) } - checkpoint, err = s.store.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID, tag.GetMetadata().GetUid(), "") + checkpoint, err = s.store.FinalizeAgentInstanceCheckpoint(ctx, checkpoint.ID.String(), tag.GetMetadata().GetUid(), "") if err != nil { return nil, serviceerrors.NewInternal("Failed to publish checkpoint", err) } @@ -119,7 +120,7 @@ func (s *Service) ensureTag(ctx context.Context, checkpoint *dbpkg.AgentInstance if err := s.verifySnapshot(ctx, checkpoint); err != nil { return nil, err } - name := tagName(checkpoint.ID) + name := tagName(checkpoint.ID.String()) tag, err := s.tags.CreateActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, name, checkpoint.SnapshotName) if err != nil { tag, err = s.tags.GetActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, name) @@ -200,7 +201,7 @@ func (s *Service) List(ctx context.Context, request ListRequest) (ListResult, er result.Checkpoints[i] = checkpointProto(&rows[i]) } if len(rows) > pageSize { - result.NextPageToken = encodePageToken(rows[pageSize-1].ID) + result.NextPageToken = encodePageToken(rows[pageSize-1].ID.String()) } return result, nil } @@ -220,7 +221,7 @@ func (s *Service) Delete(ctx context.Context, namespace, checkpointID string) er if err != nil { return serviceerrors.NewInternal("Failed to begin checkpoint deletion", err) } - tag, err := s.tags.GetActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, tagName(checkpoint.ID)) + tag, err := s.tags.GetActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, tagName(checkpoint.ID.String())) if err != nil && status.Code(err) != codes.NotFound { return serviceerrors.NewUnavailable("Failed to get checkpoint snapshot tag", err) } @@ -228,7 +229,7 @@ func (s *Service) Delete(ctx context.Context, namespace, checkpointID string) er tag.GetSnapshot().GetAtespace() != checkpoint.SnapshotAtespace || tag.GetSnapshot().GetName() != checkpoint.SnapshotName) { return serviceerrors.NewFailedPrecondition("Checkpoint snapshot tag identity changed", nil) } - if err := s.tags.DeleteActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, tagName(checkpoint.ID)); err != nil && status.Code(err) != codes.NotFound { + if err := s.tags.DeleteActorSnapshotTag(ctx, checkpoint.SnapshotAtespace, tagName(checkpoint.ID.String())); err != nil && status.Code(err) != codes.NotFound { return serviceerrors.NewUnavailable("Failed to delete checkpoint snapshot tag", err) } if err := s.store.DeleteAgentInstanceCheckpoint(ctx, namespace, checkpointID, userID); err != nil { @@ -290,7 +291,7 @@ func (s *Service) authorize(ctx context.Context, verb auth.Verb, resourceType, n func checkpointProto(checkpoint *dbpkg.AgentInstanceCheckpoint) *apiv1alpha1.Checkpoint { result := &apiv1alpha1.Checkpoint{ - Id: checkpoint.ID, Namespace: checkpoint.Namespace, AgentInstanceId: checkpoint.SourceInstanceID, + Id: checkpoint.ID.String(), Namespace: checkpoint.Namespace, AgentInstanceId: checkpoint.SourceInstanceID.String(), HeadTaskId: checkpoint.HeadTaskID, HistorySequence: uint64(checkpoint.HistorySequence), State: checkpointState(checkpoint.State), CreatedAt: timestamppb.New(checkpoint.CreatedAt), } diff --git a/go/core/v2/checkpoint/service_test.go b/go/core/v2/checkpoint/service_test.go index 1fd962638..8fd6c276e 100644 --- a/go/core/v2/checkpoint/service_test.go +++ b/go/core/v2/checkpoint/service_test.go @@ -6,6 +6,7 @@ import ( "time" "github.com/agent-substrate/substrate/pkg/proto/ateapipb" + "github.com/google/uuid" dbpkg "github.com/kagent-dev/kagent/go/api/database" apiv1alpha1 "github.com/kagent-dev/kagent/go/api/gen/kagent/api/v1alpha1" "github.com/kagent-dev/kagent/go/core/pkg/auth" @@ -165,28 +166,27 @@ func TestCreateCleansTagBeforeFailing(t *testing.T) { func TestDeleteHidesCheckpointBeforeDeletingTag(t *testing.T) { checkpoint := &dbpkg.AgentInstanceCheckpoint{ - ID: "018f47a2-4efb-7c21-a848-123456789abc", Namespace: "team-a", UserID: "alice", + ID: uuid.MustParse("018f47a2-4efb-7c21-a848-123456789abc"), Namespace: "team-a", UserID: "alice", SnapshotAtespace: "team-a", SnapshotName: "snapshot-1", TagUID: "tag-uid", State: "READY", } store := &testStore{prepared: checkpoint} tags := &testTags{created: &ateapipb.ActorSnapshotTag{ - Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: tagName(checkpoint.ID), Uid: "tag-uid"}, + Metadata: &ateapipb.ResourceMetadata{Atespace: "team-a", Name: tagName(checkpoint.ID.String()), Uid: "tag-uid"}, Snapshot: &ateapipb.ObjectRef{Atespace: "team-a", Name: "snapshot-1"}, }} service := NewService(store, testAuthorizer{}, tags, nil) ctx := auth.AuthSessionTo(context.Background(), testSession{userID: "alice"}) - if err := service.Delete(ctx, "team-a", checkpoint.ID); err != nil { + if err := service.Delete(ctx, "team-a", checkpoint.ID.String()); err != nil { t.Fatal(err) } if checkpoint.State != "DELETING" || tags.deleteCalls != 1 || !store.deleted { t.Fatalf("checkpoint state = %s, tag deletes = %d, row deleted = %v", checkpoint.State, tags.deleteCalls, store.deleted) } } - func TestForkCreatesAgentInstanceFromCheckpoint(t *testing.T) { checkpoint := &dbpkg.AgentInstanceCheckpoint{ - ID: "018f47a2-4efb-7c21-a848-123456789abc", Namespace: "team-a", UserID: "alice", + ID: uuid.MustParse("018f47a2-4efb-7c21-a848-123456789abc"), Namespace: "team-a", UserID: "alice", SnapshotAtespace: "team-a", SnapshotName: "snapshot-1", SnapshotUID: "snapshot-uid", SnapshotContentScope: "DATA", State: "READY", } store := &testStore{prepared: checkpoint} @@ -194,7 +194,7 @@ func TestForkCreatesAgentInstanceFromCheckpoint(t *testing.T) { service := NewService(store, testAuthorizer{}, &testTags{}, workflow) ctx := auth.AuthSessionTo(context.Background(), testSession{userID: "alice"}) - instance, err := service.Fork(ctx, "team-a", checkpoint.ID, "fork-request") + instance, err := service.Fork(ctx, "team-a", checkpoint.ID.String(), "fork-request") if err != nil { t.Fatal(err) }