diff --git a/hub-server/internal/repository/agent_test.go b/hub-server/internal/repository/agent_test.go index ffd311ec0..5519ed816 100644 --- a/hub-server/internal/repository/agent_test.go +++ b/hub-server/internal/repository/agent_test.go @@ -4,6 +4,7 @@ import ( "testing" "time" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" @@ -109,3 +110,456 @@ func TestBumpRunningTaskExpireAt_UnknownTaskIsNotAnError(t *testing.T) { // (bumpRunningTaskHeartbeat) warns and increments a metric on error only. require.NoError(t, BumpRunningTaskExpireAt(db, "task-missing", 10*time.Minute)) } + +// ============================================================================= +// AgentInstance repository tests +// ============================================================================= + +func TestAgentInstanceRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + ai := &model.AgentInstance{ + AgentType: "code-explorer", + SessionID: "session-ai-1", + InviterUserID: "user-inviter", + DisplayName: "Code Explorer", + } + err := CreateAgentInstance(db, ai) + require.NoError(t, err) + assert.NotEmpty(t, ai.ID) + + // Get by ID + fetched, err := GetAgentInstanceByID(db, ai.ID) + require.NoError(t, err) + assert.Equal(t, "code-explorer", fetched.AgentType) + + // Create a second agent instance in the same session + ai2 := &model.AgentInstance{ + AgentType: "code-reviewer", + SessionID: "session-ai-1", + InviterUserID: "user-inviter", + DisplayName: "Code Reviewer", + } + require.NoError(t, CreateAgentInstance(db, ai2)) + + // List by session + list, err := ListAgentInstancesBySession(db, "session-ai-1") + require.NoError(t, err) + assert.Len(t, list, 2) + + // List by inviter + list, err = ListAgentInstancesByInviter(db, "session-ai-1", "user-inviter") + require.NoError(t, err) + assert.Len(t, list, 2) + + // List by different inviter + list, err = ListAgentInstancesByInviter(db, "session-ai-1", "other-user") + require.NoError(t, err) + assert.Len(t, list, 0) + + // Delete + require.NoError(t, DeleteAgentInstance(db, ai2.ID)) + list, err = ListAgentInstancesBySession(db, "session-ai-1") + require.NoError(t, err) + assert.Len(t, list, 1) +} + +// ============================================================================= +// CustomAgent repository tests +// ============================================================================= + +func TestCustomAgentRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + ca := &model.CustomAgent{ + OwnerUserID: "user-ca", + Name: "My Agent", + AgentType: "code-explorer", + SystemPrompt: "You are a helpful assistant.", + CapabilityTags: `["code"]`, + ToolWhitelist: `["read","write"]`, + ModelParams: `{}`, + } + err := CreateCustomAgent(db, ca) + require.NoError(t, err) + assert.NotEmpty(t, ca.ID) + + // Get by ID + fetched, err := GetCustomAgentByID(db, ca.ID) + require.NoError(t, err) + assert.Equal(t, "My Agent", fetched.Name) + + // List by owner + list, err := ListCustomAgentsByOwner(db, "user-ca") + require.NoError(t, err) + assert.Len(t, list, 1) + + // Create another + ca2 := &model.CustomAgent{ + OwnerUserID: "user-ca", + Name: "Agent 2", + AgentType: "code-reviewer", + SystemPrompt: "Review code.", + } + require.NoError(t, CreateCustomAgent(db, ca2)) + list, err = ListCustomAgentsByOwner(db, "user-ca") + require.NoError(t, err) + assert.Len(t, list, 2) + + // Update + ca.Name = "Renamed Agent" + err = UpdateCustomAgent(db, ca) + require.NoError(t, err) + fetched, err = GetCustomAgentByID(db, ca.ID) + require.NoError(t, err) + assert.Equal(t, "Renamed Agent", fetched.Name) + + // Soft delete + require.NoError(t, SoftDeleteCustomAgent(db, ca2.ID)) + _, err = GetCustomAgentByID(db, ca2.ID) + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + + // But the first is still there + fetched, err = GetCustomAgentByID(db, ca.ID) + require.NoError(t, err) + assert.NotNil(t, fetched) +} + +// ============================================================================= +// PendingAgentTask repository tests +// ============================================================================= + +func TestPendingTaskRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + expireAt := time.Now().Add(time.Hour) + task := &model.PendingAgentTask{ + AgentInstanceID: "agent-inst-1", + TriggeredByUserID: "user-trigger", + TriggerMessageID: "msg-trigger-1", + Status: model.TaskStatusQueued, + ExpireAt: expireAt, + } + err := CreatePendingTask(db, task) + require.NoError(t, err) + assert.NotEmpty(t, task.ID) + + // Get by ID + fetched, err := GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusQueued, fetched.Status) + + // Update status to dispatched + err = UpdatePendingTaskStatus(db, task.ID, model.TaskStatusDispatched, "") + require.NoError(t, err) + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusDispatched, fetched.Status) + assert.NotNil(t, fetched.DispatchedAt) + + err = UpdatePendingTaskDispatched(db, task.ID, "device-edge-1") + require.NoError(t, err) + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusDispatched, fetched.Status) + assert.Equal(t, "device-edge-1", fetched.EdgeDeviceID) + + // Update status to running and persist the Edge run mapping. + err = UpdatePendingTaskStatusWithEdgeRunID(db, task.ID, model.TaskStatusRunning, "", "run-edge-1") + require.NoError(t, err) + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusRunning, fetched.Status) + assert.Equal(t, "run-edge-1", fetched.EdgeRunID) + + // Update status to done + err = UpdatePendingTaskStatus(db, task.ID, model.TaskStatusDone, "") + require.NoError(t, err) + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusDone, fetched.Status) + assert.NotNil(t, fetched.FinishedAt) + + err = UpdatePendingTaskDispatched(db, task.ID, "device-after-done") + require.ErrorIs(t, err, gorm.ErrRecordNotFound) + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusDone, fetched.Status, "terminal task should not be moved back to dispatched") + assert.Equal(t, "device-edge-1", fetched.EdgeDeviceID) + + // Update with error + task2 := &model.PendingAgentTask{ + AgentInstanceID: "agent-inst-2", + TriggeredByUserID: "user-trigger", + TriggerMessageID: "msg-trigger-2", + Status: model.TaskStatusQueued, + ExpireAt: expireAt, + } + require.NoError(t, CreatePendingTask(db, task2)) + err = UpdatePendingTaskStatus(db, task2.ID, model.TaskStatusFailed, "something went wrong") + require.NoError(t, err) + fetched, err = GetPendingTaskByID(db, task2.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusFailed, fetched.Status) + assert.Equal(t, "something went wrong", fetched.ErrorMessage) +} + +func TestPendingTaskRepo_CancelTasksByAgent(t *testing.T) { + db := setupSQLite(t) + + expireAt := time.Now().Add(time.Hour) + task1 := &model.PendingAgentTask{ + AgentInstanceID: "agent-cancel", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-1", + Status: model.TaskStatusQueued, + ExpireAt: expireAt, + } + task2 := &model.PendingAgentTask{ + AgentInstanceID: "agent-cancel", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-2", + Status: model.TaskStatusDispatched, + ExpireAt: expireAt, + } + task3 := &model.PendingAgentTask{ + AgentInstanceID: "agent-other", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-3", + Status: model.TaskStatusQueued, + ExpireAt: expireAt, + } + require.NoError(t, CreatePendingTask(db, task1)) + require.NoError(t, CreatePendingTask(db, task2)) + require.NoError(t, CreatePendingTask(db, task3)) + + err := CancelTasksByAgentInstance(db, "agent-cancel") + require.NoError(t, err) + + // Tasks 1 and 2 should be cancelled + fetched, _ := GetPendingTaskByID(db, task1.ID) + require.NotNil(t, fetched) + assert.Equal(t, model.TaskStatusCancelled, fetched.Status) + + fetched, _ = GetPendingTaskByID(db, task2.ID) + require.NotNil(t, fetched) + assert.Equal(t, model.TaskStatusCancelled, fetched.Status) + + // Task 3 (different agent) should still be queued + fetched, _ = GetPendingTaskByID(db, task3.ID) + require.NotNil(t, fetched) + assert.Equal(t, model.TaskStatusQueued, fetched.Status) +} + +func TestPendingTaskRepo_ScanExpiredTasks(t *testing.T) { + db := setupSQLite(t) + + // Create tasks with different statuses and expire times + // Expired queued + task1 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-1", + Status: model.TaskStatusQueued, + ExpireAt: time.Now().Add(-time.Hour), // expired + } + // Expired dispatched + task2 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-2", + Status: model.TaskStatusDispatched, + ExpireAt: time.Now().Add(-time.Hour), // expired + } + // Expired running (#132: running tasks should also be scanned) + task3 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-3", + Status: model.TaskStatusRunning, + ExpireAt: time.Now().Add(-time.Hour), // expired + } + // Not expired yet + task4 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-4", + Status: model.TaskStatusQueued, + ExpireAt: time.Now().Add(time.Hour), // not expired + } + // Expired but already done (terminal states excluded) + task5 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-5", + Status: model.TaskStatusDone, + ExpireAt: time.Now().Add(-time.Hour), // expired but done + } + // Expired but already failed (terminal states excluded) + task6 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-6", + Status: model.TaskStatusFailed, + ExpireAt: time.Now().Add(-time.Hour), // expired but failed + } + // Expired but already cancelled (terminal states excluded) + task7 := &model.PendingAgentTask{ + AgentInstanceID: "agent-scan", + TriggeredByUserID: "user-scan", + TriggerMessageID: "msg-7", + Status: model.TaskStatusCancelled, + ExpireAt: time.Now().Add(-time.Hour), // expired but cancelled + } + + require.NoError(t, CreatePendingTask(db, task1)) + require.NoError(t, CreatePendingTask(db, task2)) + require.NoError(t, CreatePendingTask(db, task3)) + require.NoError(t, CreatePendingTask(db, task4)) + require.NoError(t, CreatePendingTask(db, task5)) + require.NoError(t, CreatePendingTask(db, task6)) + require.NoError(t, CreatePendingTask(db, task7)) + + tasks, err := ScanExpiredTasks(db) + require.NoError(t, err) + + // Should only return tasks 1, 2, 3 (expired and in non-terminal states) + assert.Len(t, tasks, 3) + + taskIDs := make(map[string]bool) + for _, t := range tasks { + taskIDs[t.ID] = true + } + assert.True(t, taskIDs[task1.ID], "expired queued task should be returned") + assert.True(t, taskIDs[task2.ID], "expired dispatched task should be returned") + assert.True(t, taskIDs[task3.ID], "expired running task should be returned (#132)") + assert.False(t, taskIDs[task4.ID], "non-expired task should not be returned") + assert.False(t, taskIDs[task5.ID], "expired done task should not be returned") + assert.False(t, taskIDs[task6.ID], "expired failed task should not be returned") + assert.False(t, taskIDs[task7.ID], "expired cancelled task should not be returned") +} + +// ============================================================================= +// B2 data-integrity tests — atomic status, password+token, pin limit +// ============================================================================= + +func TestPendingTaskRepo_AtomicStatusUpdate(t *testing.T) { + db := setupSQLite(t) + + expireAt := time.Now().Add(time.Hour) + task := &model.PendingAgentTask{ + AgentInstanceID: "agent-atomic", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-1", + Status: model.TaskStatusDispatched, + ExpireAt: expireAt, + } + require.NoError(t, CreatePendingTask(db, task)) + + // Atomic transition: dispatched → running (should succeed) + rows, err := UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusDispatched, model.TaskStatusRunning, "") + require.NoError(t, err) + assert.Equal(t, int64(1), rows) + + fetched, err := GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusRunning, fetched.Status) + + // Atomic transition with wrong old status (should fail — 0 rows) + rows, err = UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusDispatched, model.TaskStatusDone, "") + require.NoError(t, err) + assert.Equal(t, int64(0), rows, "should not transition from running when oldStatus is dispatched") + + // Verify status unchanged + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusRunning, fetched.Status) + + // Atomic transition: running → done (should succeed) + rows, err = UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusRunning, model.TaskStatusDone, "") + require.NoError(t, err) + assert.Equal(t, int64(1), rows) + + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusDone, fetched.Status) + assert.NotNil(t, fetched.FinishedAt) +} + +func TestPendingTaskRepo_AtomicWithEdgeRunID(t *testing.T) { + db := setupSQLite(t) + + expireAt := time.Now().Add(time.Hour) + task := &model.PendingAgentTask{ + AgentInstanceID: "agent-edge", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-1", + Status: model.TaskStatusDispatched, + ExpireAt: expireAt, + } + require.NoError(t, CreatePendingTask(db, task)) + + // Atomic: dispatched → running with edgeRunID + rows, err := UpdatePendingTaskStatusAtomicWithEdgeRunID(db, task.ID, model.TaskStatusDispatched, model.TaskStatusRunning, "", "run-001") + require.NoError(t, err) + assert.Equal(t, int64(1), rows) + + fetched, err := GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusRunning, fetched.Status) + assert.Equal(t, "run-001", fetched.EdgeRunID) + + // Edge run id backfill on an already populated edgeRunID is a no-op. + rows, err = UpdatePendingTaskEdgeRunID(db, task.ID, "run-002") + // edge_run_id is already "run-001" so WHERE edge_run_id = '' matches 0 rows + require.NoError(t, err) + assert.Equal(t, int64(0), rows) + fetched, err = GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, "run-001", fetched.EdgeRunID, "edgeRunID should not be overwritten") + + backfillTask := &model.PendingAgentTask{ + AgentInstanceID: "agent-edge", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-2", + Status: model.TaskStatusRunning, + ExpireAt: expireAt, + } + require.NoError(t, CreatePendingTask(db, backfillTask)) + + rows, err = UpdatePendingTaskEdgeRunID(db, backfillTask.ID, "run-002") + require.NoError(t, err) + assert.Equal(t, int64(1), rows) + fetched, err = GetPendingTaskByID(db, backfillTask.ID) + require.NoError(t, err) + assert.Equal(t, "run-002", fetched.EdgeRunID) +} + +func TestPendingTaskRepo_AtomicFailClosed(t *testing.T) { + db := setupSQLite(t) + + expireAt := time.Now().Add(time.Hour) + task := &model.PendingAgentTask{ + AgentInstanceID: "agent-failclosed", + TriggeredByUserID: "user-t", + TriggerMessageID: "msg-1", + Status: model.TaskStatusRunning, + ExpireAt: expireAt, + } + require.NoError(t, CreatePendingTask(db, task)) + + // Simulate a race: first writer marks it done + rows, err := UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusRunning, model.TaskStatusDone, "") + require.NoError(t, err) + assert.Equal(t, int64(1), rows) + + // Second writer tries to mark it failed — 0 rows (fail closed) + rows, err = UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusRunning, model.TaskStatusFailed, "boom") + require.NoError(t, err) + assert.Equal(t, int64(0), rows, "second writer should get 0 rows affected") + + // Status should remain done + fetched, err := GetPendingTaskByID(db, task.ID) + require.NoError(t, err) + assert.Equal(t, model.TaskStatusDone, fetched.Status) +} diff --git a/hub-server/internal/repository/attachment_test.go b/hub-server/internal/repository/attachment_test.go new file mode 100644 index 000000000..d925a3c7e --- /dev/null +++ b/hub-server/internal/repository/attachment_test.go @@ -0,0 +1,49 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// ============================================================================= +// Attachment repository tests +// ============================================================================= + +func TestAttachmentRepo_CreateAndGet(t *testing.T) { + db := setupSQLite(t) + + a := &model.Attachment{ + Hash: "abc123hash", + Size: 2048, + MimeType: "image/png", + OriginalName: "screenshot.png", + UploaderUserID: "user-att", + Metadata: `{"origin":"test"}`, + } + err := CreateAttachment(db, a) + require.NoError(t, err) + assert.NotEmpty(t, a.ID) + + // Get by ID + fetched, err := GetAttachmentByID(db, a.ID) + require.NoError(t, err) + assert.Equal(t, "abc123hash", fetched.Hash) + assert.Equal(t, int64(2048), fetched.Size) + assert.JSONEq(t, `{"origin":"test"}`, fetched.Metadata) + + // Get by hash + fetched, err = GetAttachmentByHash(db, "abc123hash") + require.NoError(t, err) + assert.Equal(t, a.ID, fetched.ID) + + // Non-existent + _, err = GetAttachmentByID(db, "nonexistent") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + + _, err = GetAttachmentByHash(db, "nonexistent") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) +} diff --git a/hub-server/internal/repository/db_test.go b/hub-server/internal/repository/db_test.go new file mode 100644 index 000000000..fba122f83 --- /dev/null +++ b/hub-server/internal/repository/db_test.go @@ -0,0 +1,20 @@ +package repository + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "gorm.io/gorm" +) + +func TestWrapNotFound(t *testing.T) { + mappedErr := errors.New("mapped not found") + + assert.NoError(t, WrapNotFound(nil, mappedErr)) + assert.ErrorIs(t, WrapNotFound(gorm.ErrRecordNotFound, mappedErr), mappedErr) + assert.ErrorIs(t, WrapNotFound(errors.Join(errors.New("lookup failed"), gorm.ErrRecordNotFound), mappedErr), mappedErr) + + otherErr := errors.New("other") + assert.ErrorIs(t, WrapNotFound(otherErr, mappedErr), otherErr) +} diff --git a/hub-server/internal/repository/device_test.go b/hub-server/internal/repository/device_test.go new file mode 100644 index 000000000..7033e8d2f --- /dev/null +++ b/hub-server/internal/repository/device_test.go @@ -0,0 +1,188 @@ +package repository + +import ( + "testing" + "time" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func TestListDevicesByUserOrdersMostRecentFirst(t *testing.T) { + db := setupSQLite(t) + now := time.Now() + require.NoError(t, CreateUser(db, &model.User{ID: "user-a", Username: "user-a", Nickname: "User A"})) + require.NoError(t, CreateUser(db, &model.User{ID: "user-b", Username: "user-b", Nickname: "User B"})) + require.NoError(t, db.Create(&model.Device{ + ID: "device-old", + UserID: "user-a", + DeviceType: "desktop", + Capabilities: "[]", + LastActiveAt: now.Add(-time.Hour), + }).Error) + require.NoError(t, db.Create(&model.Device{ + ID: "device-new", + UserID: "user-a", + DeviceType: "desktop", + Capabilities: "[]", + LastActiveAt: now, + }).Error) + require.NoError(t, db.Create(&model.Device{ + ID: "device-other", + UserID: "user-b", + DeviceType: "desktop", + Capabilities: "[]", + LastActiveAt: now.Add(time.Hour), + }).Error) + + devices, err := ListDevicesByUser(db, "user-a") + require.NoError(t, err) + require.Len(t, devices, 2) + assert.Equal(t, "device-new", devices[0].ID) + assert.Equal(t, "device-old", devices[1].ID) +} + +func TestUpdateDeviceAppliesProvidedFields(t *testing.T) { + db := setupSQLite(t) + oldLastActive := time.Now().Add(-time.Hour) + newLastActive := time.Now() + require.NoError(t, db.Create(&model.Device{ + ID: "device-update", + UserID: "user-a", + DeviceType: "desktop", + AppVersion: "1.0.0", + Capabilities: `["shell"]`, + LastActiveAt: oldLastActive, + }).Error) + + err := UpdateDevice(db, "device-update", map[string]interface{}{ + "app_version": "1.1.0", + "capabilities": `["shell","browser"]`, + "last_active_at": newLastActive, + }) + require.NoError(t, err) + + device, err := GetDeviceByID(db, "device-update") + require.NoError(t, err) + assert.Equal(t, "1.1.0", device.AppVersion) + assert.Equal(t, `["shell","browser"]`, device.Capabilities) + assert.WithinDuration(t, newLastActive, device.LastActiveAt, time.Second) +} + +func TestDeleteDeviceRemovesOnlyTargetDevice(t *testing.T) { + db := setupSQLite(t) + now := time.Now() + require.NoError(t, db.Create(&model.Device{ + ID: "device-delete", + UserID: "user-a", + DeviceType: "desktop", + Capabilities: "[]", + LastActiveAt: now, + }).Error) + require.NoError(t, db.Create(&model.Device{ + ID: "device-keep", + UserID: "user-a", + DeviceType: "mobile", + Capabilities: "[]", + LastActiveAt: now, + }).Error) + + require.NoError(t, DeleteDevice(db, "device-delete")) + + _, err := GetDeviceByID(db, "device-delete") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + kept, err := GetDeviceByID(db, "device-keep") + require.NoError(t, err) + assert.Equal(t, "device-keep", kept.ID) +} + +// ============================================================================= +// Device repository tests +// ============================================================================= + +func TestDeviceRepo_Upsert(t *testing.T) { + db := setupSQLite(t) + + device := &model.Device{ + ID: "dev-001", + UserID: "user-001", + DeviceType: "desktop", + AppVersion: "1.0.0", + Capabilities: `["chat","agent"]`, + } + + // First upsert: creates + err := UpsertDevice(db, device) + require.NoError(t, err) + + fetched, err := GetDeviceByID(db, "dev-001") + require.NoError(t, err) + assert.Equal(t, "desktop", fetched.DeviceType) + assert.Equal(t, "1.0.0", fetched.AppVersion) + + // Second upsert: same physical device updates by device ID. + device2 := &model.Device{ + ID: "dev-001", + UserID: "user-001", + DeviceType: "desktop", + AppVersion: "2.0.0", + Capabilities: `["chat","agent","file"]`, + } + err = UpsertDevice(db, device2) + require.NoError(t, err) + + // ON CONFLICT preserves the original row's ID but updates other columns. + // Verify the original row was updated. + fetched, err = GetDeviceByID(db, "dev-001") + require.NoError(t, err) + assert.Equal(t, "2.0.0", fetched.AppVersion) + + // A second desktop for the same user is a distinct device and must keep its + // own row so refresh_tokens.device_id can reference it. + device3 := &model.Device{ + ID: "dev-001-second-desktop", + UserID: "user-001", + DeviceType: "desktop", + AppVersion: "1.0.0", + Capabilities: `["chat"]`, + } + require.NoError(t, UpsertDevice(db, device3)) + fetched, err = GetDeviceByID(db, "dev-001-second-desktop") + require.NoError(t, err) + assert.Equal(t, "user-001", fetched.UserID) + assert.Equal(t, "desktop", fetched.DeviceType) + + stolen := &model.Device{ + ID: "dev-001", + UserID: "user-attacker", + DeviceType: "desktop", + AppVersion: "9.9.9", + Capabilities: `[]`, + } + require.Error(t, UpsertDevice(db, stolen)) + fetched, err = GetDeviceByID(db, "dev-001") + require.NoError(t, err) + assert.Equal(t, "user-001", fetched.UserID) + assert.Equal(t, "2.0.0", fetched.AppVersion) +} + +func TestDeviceRepo_GetByID(t *testing.T) { + db := setupSQLite(t) + + device := &model.Device{ + ID: "dev-002", + UserID: "user-002", + DeviceType: "mobile", + } + require.NoError(t, UpsertDevice(db, device)) + + fetched, err := GetDeviceByID(db, "dev-002") + require.NoError(t, err) + assert.Equal(t, "user-002", fetched.UserID) + assert.Equal(t, "mobile", fetched.DeviceType) + + _, err = GetDeviceByID(db, "nonexistent") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) +} diff --git a/hub-server/internal/repository/friendship_test.go b/hub-server/internal/repository/friendship_test.go new file mode 100644 index 000000000..d6df807de --- /dev/null +++ b/hub-server/internal/repository/friendship_test.go @@ -0,0 +1,199 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// ============================================================================= +// Friendship repository tests +// ============================================================================= + +func TestFriendshipRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + f := &model.Friendship{ + UserID: "user-a", + FriendID: "user-b", + Status: model.StatusPending, + RequestMessage: "Please add me", + } + + // Create + err := CreateFriendship(db, f) + require.NoError(t, err) + assert.NotEmpty(t, f.ID) + + // Find between + found, err := FindFriendshipBetween(db, "user-a", "user-b") + require.NoError(t, err) + require.NotNil(t, found) + assert.Equal(t, model.StatusPending, found.Status) + + // Also find reversed + found, err = FindFriendshipBetween(db, "user-b", "user-a") + require.NoError(t, err) + require.NotNil(t, found) +} + +func TestFriendshipRepo_StatusTransitions(t *testing.T) { + db := setupSQLite(t) + + f := &model.Friendship{ + UserID: "user-1", + FriendID: "user-2", + Status: model.StatusPending, + } + require.NoError(t, CreateFriendship(db, f)) + + // Accept + err := UpdateFriendshipByID(db, f.ID, model.StatusAccepted) + require.NoError(t, err) + + fetched, err := GetFriendshipByID(db, f.ID) + require.NoError(t, err) + assert.Equal(t, model.StatusAccepted, fetched.Status) + + // Update remark + err = UpdateFriendshipRemark(db, "user-1", "user-2", "Bestie") + require.NoError(t, err) + + fetched, err = GetFriendshipByID(db, f.ID) + require.NoError(t, err) + assert.Equal(t, "Bestie", fetched.Remark) +} + +func TestFriendshipRepo_UpdateRemarkNoRows(t *testing.T) { + db := setupSQLite(t) + + // No friendship rows exist, so UpdateFriendshipRemark should return ErrRecordNotFound. + err := UpdateFriendshipRemark(db, "user-x", "user-y", "remark") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + + // Create a friendship with pending status - remark should NOT be updatable + f := &model.Friendship{UserID: "user-x", FriendID: "user-y", Status: model.StatusPending} + require.NoError(t, CreateFriendship(db, f)) + + err = UpdateFriendshipRemark(db, "user-x", "user-y", "remark") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound, "pending friendship should not allow remark update") + + // Update to accepted - remark SHOULD be updatable + require.NoError(t, UpdateFriendshipByID(db, f.ID, model.StatusAccepted)) + err = UpdateFriendshipRemark(db, "user-x", "user-y", "cleared") + require.NoError(t, err) + + fetched, err := GetFriendshipByID(db, f.ID) + require.NoError(t, err) + assert.Equal(t, "cleared", fetched.Remark) + + // Clearing remark to empty string + err = UpdateFriendshipRemark(db, "user-x", "user-y", "") + require.NoError(t, err) + fetched, err = GetFriendshipByID(db, f.ID) + require.NoError(t, err) + assert.Equal(t, "", fetched.Remark) +} + +func TestFriendshipRepo_Lists(t *testing.T) { + db := setupSQLite(t) + + // Pending incoming + f1 := &model.Friendship{UserID: "alice", FriendID: "bob", Status: model.StatusPending} + f2 := &model.Friendship{UserID: "carol", FriendID: "bob", Status: model.StatusPending} + // Accepted + f3 := &model.Friendship{UserID: "bob", FriendID: "dave", Status: model.StatusAccepted} + + require.NoError(t, CreateFriendship(db, f1)) + require.NoError(t, CreateFriendship(db, f2)) + require.NoError(t, CreateFriendship(db, f3)) + + // Pending (received for bob) + received, err := ListReceivedRequests(db, "bob") + require.NoError(t, err) + assert.Len(t, received, 2) + + // Pending (sent by alice) + sent, err := ListSentRequests(db, "alice") + require.NoError(t, err) + assert.Len(t, sent, 1) + + // Accepted friends (bob's connections) + accepted, err := ListAcceptedFriends(db, "bob") + require.NoError(t, err) + assert.Len(t, accepted, 1) + assert.Equal(t, "dave", accepted[0].FriendID) + + // Friend IDs + ids, err := GetFriendIDs(db, "bob") + require.NoError(t, err) + assert.Len(t, ids, 1) + assert.Equal(t, "dave", ids[0]) +} + +func TestFriendshipRepo_BlockAndDelete(t *testing.T) { + db := setupSQLite(t) + + f := &model.Friendship{UserID: "u1", FriendID: "u2", Status: model.StatusAccepted} + require.NoError(t, CreateFriendship(db, f)) + + // Block + err := UpdateFriendshipByID(db, f.ID, model.StatusBlocked) + require.NoError(t, err) + + blocked, err := IsBlockedBy(db, "u1", "u2") + require.NoError(t, err) + assert.True(t, blocked) + + // Delete pair + err = DeleteFriendshipPair(db, "u1", "u2") + require.NoError(t, err) + + found, err := FindFriendshipBetween(db, "u1", "u2") + require.NoError(t, err) + assert.Nil(t, found) +} + +func TestFriendshipRepo_DeleteFriendshipDeletesLoadedRow(t *testing.T) { + db := setupSQLite(t) + + target := &model.Friendship{UserID: "delete-u1", FriendID: "delete-u2", Status: model.StatusPending} + other := &model.Friendship{UserID: "keep-u1", FriendID: "keep-u2", Status: model.StatusBlocked} + require.NoError(t, CreateFriendship(db, target)) + require.NoError(t, CreateFriendship(db, other)) + + require.NoError(t, DeleteFriendship(db, target)) + + _, err := GetFriendshipByID(db, target.ID) + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + kept, err := GetFriendshipByID(db, other.ID) + require.NoError(t, err) + assert.Equal(t, other.ID, kept.ID) +} + +func TestFriendshipRepo_Upsert(t *testing.T) { + db := setupSQLite(t) + + f := &model.Friendship{ + UserID: "upsert-a", + FriendID: "upsert-b", + Status: model.StatusPending, + } + require.NoError(t, UpsertFriendship(db, f)) + + // Upsert with same user_id+friend_id updates + f2 := &model.Friendship{ + UserID: "upsert-a", + FriendID: "upsert-b", + Status: model.StatusAccepted, + } + require.NoError(t, UpsertFriendship(db, f2)) + + found, err := FindFriendshipBetween(db, "upsert-a", "upsert-b") + require.NoError(t, err) + require.NotNil(t, found) + assert.Equal(t, model.StatusAccepted, found.Status) +} diff --git a/hub-server/internal/repository/helpers_test.go b/hub-server/internal/repository/helpers_test.go new file mode 100644 index 000000000..eb9edfca6 --- /dev/null +++ b/hub-server/internal/repository/helpers_test.go @@ -0,0 +1,349 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/require" + "gorm.io/gorm" + gormlogger "gorm.io/gorm/logger" +) + +// setupSQLite creates an in-memory SQLite database with tables matching the +// production PostgreSQL schema. Raw SQL is used instead of AutoMigrate because +// GORM's SQLite driver mishandles PostgreSQL-specific GORM tags (jsonb with +// default:'[]' produces SQLite-invalid DEFAULT "[]"). +func setupSQLite(t *testing.T) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + Logger: gormlogger.Default.LogMode(gormlogger.Silent), + }) + require.NoError(t, err) + + tables := []string{ + `CREATE TABLE users ( + id TEXT PRIMARY KEY, + username TEXT NOT NULL UNIQUE, + password_hash TEXT, + nickname TEXT NOT NULL, + avatar_url TEXT DEFAULT '', + tokendance_sub TEXT DEFAULT NULL, + tokendance_sub_linked_at DATETIME DEFAULT NULL, + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE UNIQUE INDEX idx_users_tokendance_sub ON users(tokendance_sub) + WHERE tokendance_sub IS NOT NULL AND tokendance_sub != ''`, + `CREATE TABLE devices ( + id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + device_type TEXT NOT NULL, + app_version TEXT DEFAULT '', + capabilities TEXT DEFAULT '[]', + last_active_at DATETIME NOT NULL DEFAULT (datetime('now')), + created_at DATETIME + )`, + `CREATE INDEX idx_devices_user_type ON devices(user_id, device_type)`, + `CREATE TABLE sessions ( + id TEXT PRIMARY KEY, + type TEXT NOT NULL, + name TEXT DEFAULT '', + avatar_url TEXT DEFAULT '', + tokendance_sub TEXT DEFAULT NULL, + tokendance_sub_linked_at DATETIME DEFAULT NULL, + announcement TEXT DEFAULT '', + owner_user_id TEXT, + workspace_id TEXT, + next_seq INTEGER NOT NULL DEFAULT 0, + last_message_at DATETIME, + dissolved INTEGER NOT NULL DEFAULT 0, + created_at DATETIME + )`, + `CREATE TABLE session_members ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + member_type TEXT NOT NULL, + member_id TEXT NOT NULL, + role TEXT NOT NULL, + pinned INTEGER NOT NULL DEFAULT 0, + archived INTEGER NOT NULL DEFAULT 0, + muted INTEGER NOT NULL DEFAULT 0, + last_read_seq INTEGER NOT NULL DEFAULT 0, + joined_at DATETIME, + left_at DATETIME + )`, + `CREATE TABLE messages ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + seq_id INTEGER NOT NULL, + client_msg_id TEXT NOT NULL, + sender_type TEXT NOT NULL, + sender_id TEXT NOT NULL, + content_type TEXT NOT NULL, + content TEXT NOT NULL DEFAULT '', + reply_to_message_id TEXT, + recalled INTEGER NOT NULL DEFAULT 0, + edited BOOLEAN NOT NULL DEFAULT FALSE, + edited_at DATETIME, + created_at DATETIME + )`, + `CREATE TABLE message_pins ( + session_id TEXT NOT NULL, + message_id TEXT NOT NULL, + pinned_by_user_id TEXT NOT NULL, + pinned_at DATETIME, + PRIMARY KEY (session_id, message_id) + )`, + `CREATE TABLE message_attachments ( + session_id TEXT NOT NULL, + message_id TEXT NOT NULL, + attachment_id TEXT NOT NULL, + created_at DATETIME, + PRIMARY KEY (message_id, attachment_id) + )`, + `CREATE TABLE message_reactions ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + message_id TEXT NOT NULL, + user_id TEXT NOT NULL, + emoji TEXT NOT NULL, + created_at DATETIME, + UNIQUE (session_id, message_id, user_id, emoji) + )`, + `CREATE TABLE friendships ( + id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + friend_id TEXT NOT NULL, + status TEXT NOT NULL, + remark TEXT DEFAULT '', + request_message TEXT DEFAULT '', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE UNIQUE INDEX idx_friendships_user_friend ON friendships(user_id, friend_id)`, + // Additional models for full coverage + `CREATE TABLE notifications ( + id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + type TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '', + read INTEGER NOT NULL DEFAULT 0, + created_at DATETIME + )`, + `CREATE TABLE attachments ( + id TEXT PRIMARY KEY, + hash TEXT NOT NULL UNIQUE, + size INTEGER NOT NULL, + mime_type TEXT NOT NULL, + original_name TEXT DEFAULT '', + uploader_user_id TEXT NOT NULL, + metadata TEXT NOT NULL DEFAULT '{}', + created_at DATETIME + )`, + `CREATE TABLE agent_instances ( + id TEXT PRIMARY KEY, + agent_type TEXT NOT NULL, + custom_agent_id TEXT, + session_id TEXT NOT NULL, + inviter_user_id TEXT NOT NULL, + workspace_id TEXT, + display_name TEXT NOT NULL, + created_at DATETIME + )`, + `CREATE TABLE custom_agents ( + id TEXT PRIMARY KEY, + owner_user_id TEXT NOT NULL, + name TEXT NOT NULL, + avatar_url TEXT DEFAULT '', + tokendance_sub TEXT DEFAULT NULL, + tokendance_sub_linked_at DATETIME DEFAULT NULL, + agent_type TEXT NOT NULL, + system_prompt TEXT NOT NULL DEFAULT '', + capability_tags TEXT DEFAULT '[]', + tool_whitelist TEXT DEFAULT '[]', + model_params TEXT DEFAULT '{}', + output_schema TEXT DEFAULT NULL, + deleted_at DATETIME, + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE pending_agent_tasks ( + id TEXT PRIMARY KEY, + agent_instance_id TEXT NOT NULL, + triggered_by_user_id TEXT NOT NULL, + trigger_message_id TEXT NOT NULL, + target_id TEXT, + status TEXT NOT NULL, + edge_run_id TEXT DEFAULT '', + edge_device_id TEXT DEFAULT '', + error_message TEXT DEFAULT '', + model_params TEXT DEFAULT '{}', + created_at DATETIME, + dispatched_at DATETIME, + finished_at DATETIME, + expire_at DATETIME NOT NULL + )`, + `CREATE TABLE refresh_tokens ( + id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + device_type TEXT NOT NULL DEFAULT '', + device_id TEXT NOT NULL DEFAULT '', + token_hash TEXT NOT NULL UNIQUE, + expires_at DATETIME NOT NULL, + revoked INTEGER NOT NULL DEFAULT 0, + created_at DATETIME + )`, + `CREATE UNIQUE INDEX idx_rt_user_device ON refresh_tokens(user_id, device_type, device_id)`, + `CREATE TABLE agent_teams ( + id TEXT PRIMARY KEY, + owner_id TEXT NOT NULL, + name TEXT NOT NULL, + description TEXT DEFAULT '', + avatar_url TEXT DEFAULT '', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_members ( + id TEXT PRIMARY KEY, + team_id TEXT NOT NULL, + agent_profile_id TEXT, + role TEXT NOT NULL DEFAULT 'executor', + position INTEGER NOT NULL DEFAULT 0, + created_at DATETIME, + FOREIGN KEY (team_id) REFERENCES agent_teams(id) ON DELETE CASCADE + )`, + `CREATE TABLE agent_team_runs ( + id TEXT PRIMARY KEY, + team_id TEXT NOT NULL, + session_id TEXT, + trigger_user_id TEXT NOT NULL, + trigger_message TEXT DEFAULT '', + target_id TEXT, + mode TEXT NOT NULL DEFAULT 'supervisor', + status TEXT NOT NULL DEFAULT 'queued', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_assignments ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + from_member_id TEXT NOT NULL, + to_member_id TEXT NOT NULL, + type TEXT NOT NULL DEFAULT 'delegate', + task_prompt TEXT NOT NULL, + context TEXT DEFAULT '', + status TEXT NOT NULL DEFAULT 'pending', + run_id TEXT, + result TEXT DEFAULT '', + depth INTEGER NOT NULL DEFAULT 0, + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_tasks ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + assignment_id TEXT, + assignee_member_id TEXT NOT NULL, + parent_task_id TEXT, + status TEXT NOT NULL DEFAULT 'pending', + objective TEXT NOT NULL, + input_refs TEXT NOT NULL DEFAULT '{}', + run_id TEXT, + attempt INTEGER NOT NULL DEFAULT 1, + risk_level TEXT NOT NULL DEFAULT 'normal', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_events ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + seq INTEGER NOT NULL, + type TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '{}', + created_at DATETIME + )`, + // Mirrors migration 0056: seq must be unique per run so concurrent + // appends surface as unique violations instead of silent duplicates. + `CREATE UNIQUE INDEX uq_agent_team_events_run_seq ON agent_team_events(team_run_id, seq)`, + `CREATE TABLE agent_profiles ( + id TEXT PRIMARY KEY, + owner_id TEXT NOT NULL, + name TEXT NOT NULL, + description TEXT DEFAULT '', + runtime_id TEXT NOT NULL DEFAULT '', + model TEXT DEFAULT '', + provider TEXT DEFAULT '', + reasoning_effort TEXT DEFAULT 'medium', + model_mapping TEXT DEFAULT '{}', + skills TEXT DEFAULT '[]', + mcp_servers TEXT DEFAULT '[]', + tool_allowlist TEXT DEFAULT '[]', + approval_policy TEXT DEFAULT '{}', + permission_mode TEXT DEFAULT 'default', + target_preferences TEXT DEFAULT '{}', + context_budget_max_tokens INTEGER DEFAULT 200000, + is_public INTEGER DEFAULT 0, + install_count INTEGER DEFAULT 0, + rating_avg REAL DEFAULT 0, + rating_count INTEGER DEFAULT 0, + version INTEGER DEFAULT 1, + created_at DATETIME, + updated_at DATETIME, + deleted_at DATETIME + )`, + `CREATE TABLE agent_run_events ( + id TEXT PRIMARY KEY, + task_id TEXT NOT NULL, + edge_run_id TEXT DEFAULT '', + session_id TEXT NOT NULL, + agent_instance_id TEXT NOT NULL, + event_seq INTEGER NOT NULL, + event_type TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '', + created_at DATETIME + )`, + `CREATE TABLE agent_team_artifacts ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + team_task_id TEXT, + assignment_id TEXT, + member_id TEXT, + agent_task_id TEXT, + edge_run_id TEXT DEFAULT '', + source_event_id TEXT, + event_seq INTEGER NOT NULL DEFAULT 0, + path TEXT NOT NULL, + normalized_path TEXT NOT NULL, + action TEXT DEFAULT '', + tool_name TEXT DEFAULT '', + status TEXT DEFAULT '', + conflict_id TEXT, + created_at DATETIME, + updated_at DATETIME + )`, + } + for _, ddl := range tables { + require.NoError(t, db.Exec(ddl).Error, "DDL: %s", ddl[:60]) + } + return db +} + +// ============================================================================= +// Message repository tests +// ============================================================================= + +func createTestSession(t *testing.T, db *gorm.DB) *model.Session { + t.Helper() + s := &model.Session{Type: model.SessionTypeGroup, Name: "MsgTest"} + require.NoError(t, CreateSession(db, s)) + return s +} + +// ============================================================================= +// Helpers +// ============================================================================= + +func strPtr(s string) *string { + return &s +} diff --git a/hub-server/internal/repository/message_attachment_test.go b/hub-server/internal/repository/message_attachment_test.go new file mode 100644 index 000000000..c533d2179 --- /dev/null +++ b/hub-server/internal/repository/message_attachment_test.go @@ -0,0 +1,107 @@ +package repository + +import ( + "testing" + "time" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMessageAttachmentRepo_CreateAndAccess(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + member := &model.SessionMember{ + SessionID: s.ID, + MemberType: model.MemberTypeUser, + MemberID: "viewer-1", + Role: model.MemberRoleMember, + } + require.NoError(t, CreateSessionMember(db, member)) + + msg := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "attach-client-1", + SenderType: model.SenderTypeUser, + SenderID: "owner-1", + ContentType: model.ContentTypeFile, + Content: `{"attachment_id":"att-1"}`, + } + require.NoError(t, InsertMessage(db, msg)) + + refs := []model.MessageAttachment{{ + SessionID: s.ID, + MessageID: msg.ID, + AttachmentID: "att-1", + }} + require.NoError(t, CreateMessageAttachmentReferences(db, refs)) + require.NoError(t, CreateMessageAttachmentReferences(db, refs)) + + allowed, err := CanUserAccessReferencedAttachment(db, "viewer-1", "att-1") + require.NoError(t, err) + assert.True(t, allowed) + + allowed, err = CanUserAccessReferencedAttachment(db, "outsider-1", "att-1") + require.NoError(t, err) + assert.False(t, allowed) + + require.NoError(t, SoftDeleteMember(db, s.ID, model.MemberTypeUser, "viewer-1")) + allowed, err = CanUserAccessReferencedAttachment(db, "viewer-1", "att-1") + require.NoError(t, err) + assert.False(t, allowed) +} + +func TestMessageAttachmentRepo_ListAttachmentsByMessageIDs(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + msgWithAttachment := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "attach-client-1", + SenderType: model.SenderTypeUser, + SenderID: "owner-1", + ContentType: model.ContentTypeFile, + Content: `{"attachment_id":"att-1"}`, + } + msgWithoutAttachment := &model.Message{ + SessionID: s.ID, + SeqID: 2, + ClientMsgID: "attach-client-2", + SenderType: model.SenderTypeUser, + SenderID: "owner-1", + ContentType: model.ContentTypeText, + Content: `{"text":"plain"}`, + } + require.NoError(t, InsertMessage(db, msgWithAttachment)) + require.NoError(t, InsertMessage(db, msgWithoutAttachment)) + + require.NoError(t, db.Exec( + `INSERT INTO attachments (id, hash, size, mime_type, original_name, uploader_user_id, metadata, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, + "att-1", "hash-1", 42, "text/plain", "notes.txt", "owner-1", `{"height":3,"width":2}`, time.Now(), + ).Error) + require.NoError(t, CreateMessageAttachmentReferences(db, []model.MessageAttachment{{ + SessionID: s.ID, + MessageID: msgWithAttachment.ID, + AttachmentID: "att-1", + }})) + + attachmentsByMessage, err := ListAttachmentsByMessageIDs(db, []string{msgWithAttachment.ID, msgWithoutAttachment.ID}) + require.NoError(t, err) + + require.Len(t, attachmentsByMessage[msgWithAttachment.ID], 1) + assert.Equal(t, "att-1", attachmentsByMessage[msgWithAttachment.ID][0].ID) + assert.Equal(t, "hash-1", attachmentsByMessage[msgWithAttachment.ID][0].Hash) + assert.Equal(t, int64(42), attachmentsByMessage[msgWithAttachment.ID][0].Size) + assert.Equal(t, "text/plain", attachmentsByMessage[msgWithAttachment.ID][0].MimeType) + assert.Equal(t, "notes.txt", attachmentsByMessage[msgWithAttachment.ID][0].OriginalName) + assert.JSONEq(t, `{"height":3,"width":2}`, attachmentsByMessage[msgWithAttachment.ID][0].Metadata) + assert.Empty(t, attachmentsByMessage[msgWithoutAttachment.ID]) + + empty, err := ListAttachmentsByMessageIDs(db, nil) + require.NoError(t, err) + assert.Empty(t, empty) +} diff --git a/hub-server/internal/repository/message_reaction_test.go b/hub-server/internal/repository/message_reaction_test.go new file mode 100644 index 000000000..81c0d256b --- /dev/null +++ b/hub-server/internal/repository/message_reaction_test.go @@ -0,0 +1,117 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMessageReactionRepo_AddListCountAndRemove(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + msg := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "reaction-client-1", + SenderType: model.SenderTypeUser, + SenderID: "sender-1", + ContentType: model.ContentTypeText, + Content: `{"text":"hello"}`, + } + require.NoError(t, InsertMessage(db, msg)) + + reaction := &model.MessageReaction{ + SessionID: s.ID, + MessageID: msg.ID, + UserID: "user-1", + Reaction: "thumbs_up", + } + require.NoError(t, AddReaction(db, reaction)) + require.NotEmpty(t, reaction.ID) + + reactions, err := ListReactionsByMessage(db, s.ID, msg.ID) + require.NoError(t, err) + require.Len(t, reactions, 1) + assert.Equal(t, "thumbs_up", reactions[0].Reaction) + assert.Equal(t, "user-1", reactions[0].UserID) + + byMessage, err := ListReactionsByMessages(db, s.ID, []string{msg.ID, "missing-message"}) + require.NoError(t, err) + require.Len(t, byMessage[msg.ID], 1) + assert.Empty(t, byMessage["missing-message"]) + + counts, err := ReactionCountsByMessage(db, s.ID, []string{msg.ID, "missing-message"}) + require.NoError(t, err) + assert.Equal(t, int64(1), counts[msg.ID]) + assert.Zero(t, counts["missing-message"]) + + summaries, err := ReactionSummariesByMessage(db, s.ID, msg.ID) + require.NoError(t, err) + require.Len(t, summaries, 1) + assert.Equal(t, "thumbs_up", summaries[0].Reaction) + assert.Equal(t, 1, summaries[0].Count) + assert.Equal(t, []string{"user-1"}, summaries[0].UserIDs) + + require.NoError(t, RemoveReaction(db, s.ID, msg.ID, "user-1", "thumbs_up")) + + reactions, err = ListReactionsByMessage(db, s.ID, msg.ID) + require.NoError(t, err) + assert.Empty(t, reactions) +} + +func TestMessageReactionRepo_AddReactionIsIdempotentForDuplicateUserReaction(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + msg := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "reaction-client-duplicate", + SenderType: model.SenderTypeUser, + SenderID: "sender-1", + ContentType: model.ContentTypeText, + Content: `{"text":"hello"}`, + } + require.NoError(t, InsertMessage(db, msg)) + + reaction := model.MessageReaction{ + SessionID: s.ID, + MessageID: msg.ID, + UserID: "user-1", + Reaction: "heart", + } + require.NoError(t, AddReaction(db, &reaction)) + require.NoError(t, AddReaction(db, &model.MessageReaction{ + SessionID: s.ID, + MessageID: msg.ID, + UserID: "user-1", + Reaction: "heart", + })) + + counts, err := ReactionCountsByMessage(db, s.ID, []string{msg.ID}) + require.NoError(t, err) + assert.Equal(t, int64(1), counts[msg.ID]) +} + +func TestMessageReactionRepo_RemoveReactionIsIdempotentForMissingRows(t *testing.T) { + db := setupSQLite(t) + + require.NoError(t, RemoveReaction(db, "missing-session", "missing-message", "missing-user", "heart")) +} + +func TestMessageReactionRepo_EmptyInputsReturnEmptyMaps(t *testing.T) { + db := setupSQLite(t) + + byMessage, err := ListReactionsByMessages(db, "session-1", nil) + require.NoError(t, err) + assert.Empty(t, byMessage) + + counts, err := ReactionCountsByMessage(db, "session-1", nil) + require.NoError(t, err) + assert.Empty(t, counts) + + summaries, err := ReactionSummariesByMessage(db, "session-1", "message-1") + require.NoError(t, err) + assert.Empty(t, summaries) +} diff --git a/hub-server/internal/repository/message_test.go b/hub-server/internal/repository/message_test.go new file mode 100644 index 000000000..7c12a368c --- /dev/null +++ b/hub-server/internal/repository/message_test.go @@ -0,0 +1,247 @@ +package repository + +import ( + "testing" + "time" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func TestUpdateMessageContentMarksEditedAndStoresTimestamp(t *testing.T) { + db := setupSQLite(t) + require.NoError(t, db.Exec(`INSERT INTO messages ( + id, session_id, seq_id, client_msg_id, sender_type, sender_id, content_type, content, recalled, created_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + "msg-edit", "sess-edit", 1, "client-edit", "user", "user-1", "text", `{"text":"before"}`, false, time.Now(), + ).Error) + + require.NoError(t, UpdateMessageContent(db, "msg-edit", "text", `{"text":"after"}`)) + + var stored struct { + ContentType string + Content string + Edited bool + EditedAt *time.Time + } + require.NoError(t, db.Table("messages").Where("id = ?", "msg-edit").First(&stored).Error) + assert.Equal(t, "text", stored.ContentType) + assert.JSONEq(t, `{"text":"after"}`, stored.Content) + assert.True(t, stored.Edited) + if stored.EditedAt == nil { + t.Fatal("edited_at was not set") + } +} + +func TestUpdateMessageContentReturnsNotFoundWhenNoRowsChange(t *testing.T) { + db := setupSQLite(t) + require.ErrorIs(t, UpdateMessageContent(db, "missing", "text", `{"text":"after"}`), gorm.ErrRecordNotFound) +} + +func TestMessageRepo_InsertAndGet(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + msg := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "client-001", + SenderType: model.SenderTypeUser, + SenderID: "user-1", + ContentType: model.ContentTypeText, + Content: `{"text":"Hello"}`, + } + + err := InsertMessage(db, msg) + require.NoError(t, err) + assert.NotEmpty(t, msg.ID) + + fetched, err := GetMessageByID(db, msg.ID) + require.NoError(t, err) + assert.Equal(t, `{"text":"Hello"}`, fetched.Content) +} + +func TestMessageRepo_GetBySession(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + for i := 1; i <= 5; i++ { + msg := &model.Message{ + SessionID: s.ID, + SeqID: int64(i), + ClientMsgID: "client-" + string(rune('0'+i)), + SenderType: model.SenderTypeUser, + SenderID: "user-1", + ContentType: model.ContentTypeText, + Content: `{"text":"Message ` + string(rune('0'+i)) + `"}`, + } + require.NoError(t, InsertMessage(db, msg)) + } + + msgs, err := GetMessagesBySession(db, s.ID, 0, 10) + require.NoError(t, err) + assert.Len(t, msgs, 5) + + // Get with beforeSeq + msgs, err = GetMessagesBySession(db, s.ID, 4, 10) + require.NoError(t, err) + assert.Len(t, msgs, 3) // seq 1,2,3 (before seq 4) + + // Get with small limit + msgs, err = GetMessagesBySession(db, s.ID, 0, 2) + require.NoError(t, err) + assert.Len(t, msgs, 2) +} + +func TestMessageRepo_Increment(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + for i := 1; i <= 5; i++ { + msg := &model.Message{ + SessionID: s.ID, + SeqID: int64(i), + ClientMsgID: "inc-client-" + string(rune('0'+i)), + SenderType: model.SenderTypeUser, + SenderID: "user-1", + ContentType: model.ContentTypeText, + Content: `{"text":"Inc ` + string(rune('0'+i)) + `"}`, + } + require.NoError(t, InsertMessage(db, msg)) + } + + msgs, err := GetMessagesIncrement(db, s.ID, 2, 10) + require.NoError(t, err) + assert.Len(t, msgs, 3) // seq 3,4,5 (after seq 2) + assert.Equal(t, int64(3), msgs[0].SeqID) +} + +func TestMessageRepo_Recall(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + msg := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "recall-001", + SenderType: model.SenderTypeUser, + SenderID: "user-1", + ContentType: model.ContentTypeText, + Content: `{"text":"Recall me"}`, + } + require.NoError(t, InsertMessage(db, msg)) + + err := UpdateMessageRecalled(db, msg.ID) + require.NoError(t, err) + + fetched, err := GetMessageByID(db, msg.ID) + require.NoError(t, err) + assert.True(t, fetched.Recalled) +} + +func TestMessageRepo_DuplicateClientMsgID(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + msg := &model.Message{ + SessionID: s.ID, + SeqID: 1, + ClientMsgID: "dup-client", + SenderType: model.SenderTypeUser, + SenderID: "user-1", + ContentType: model.ContentTypeText, + Content: `{"text":"First"}`, + } + require.NoError(t, InsertMessage(db, msg)) + + fetched, err := GetMessageByClientMsgID(db, s.ID, "dup-client") + require.NoError(t, err) + require.NotNil(t, fetched) + assert.Equal(t, `{"text":"First"}`, fetched.Content) + + // Non-existent returns nil, nil + fetched, err = GetMessageByClientMsgID(db, s.ID, "no-such") + require.NoError(t, err) + assert.Nil(t, fetched) +} + +func TestMessageRepo_Pins(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + pin := &model.MessagePin{ + SessionID: s.ID, + MessageID: "msg-001", + PinnedByUserID: "user-1", + } + + err := InsertPin(db, pin) + require.NoError(t, err) + + count, err := CountPinsBySession(db, s.ID) + require.NoError(t, err) + assert.Equal(t, int64(1), count) + + pins, err := ListPinsBySession(db, s.ID) + require.NoError(t, err) + assert.Len(t, pins, 1) + assert.Equal(t, "msg-001", pins[0].MessageID) + + // Delete pin + err = DeletePin(db, s.ID, "msg-001") + require.NoError(t, err) + + count, err = CountPinsBySession(db, s.ID) + require.NoError(t, err) + assert.Equal(t, int64(0), count) +} + +func TestMessageRepo_GetByIDs(t *testing.T) { + db := setupSQLite(t) + s := createTestSession(t, db) + + msg1 := &model.Message{SessionID: s.ID, SeqID: 1, ClientMsgID: "c1", SenderType: model.SenderTypeUser, SenderID: "u1", ContentType: model.ContentTypeText, Content: `{}`} + msg2 := &model.Message{SessionID: s.ID, SeqID: 2, ClientMsgID: "c2", SenderType: model.SenderTypeUser, SenderID: "u1", ContentType: model.ContentTypeText, Content: `{}`} + require.NoError(t, InsertMessage(db, msg1)) + require.NoError(t, InsertMessage(db, msg2)) + + msgs, err := GetMessagesByIDs(db, []string{msg1.ID, msg2.ID}) + require.NoError(t, err) + assert.Len(t, msgs, 2) + + // Empty list + msgs, err = GetMessagesByIDs(db, []string{}) + require.NoError(t, err) + assert.Empty(t, msgs) +} + +func TestMessageRepo_PinMessageAtomic(t *testing.T) { + db := setupSQLite(t) + + s := &model.Session{Type: model.SessionTypeGroup, Name: "pin-test"} + require.NoError(t, CreateSession(db, s)) + + pin := func(msgID, userID string) error { + return PinMessageAtomic(db, &model.MessagePin{ + SessionID: s.ID, + MessageID: msgID, + PinnedByUserID: userID, + }, 3) // low limit for testing + } + + // Pin 3 messages — all should succeed + require.NoError(t, pin("msg-1", "user-1")) + require.NoError(t, pin("msg-2", "user-1")) + require.NoError(t, pin("msg-3", "user-1")) + + // 4th pin should fail with limit exceeded + err := pin("msg-4", "user-1") + assert.ErrorIs(t, err, ErrPinLimitExceeded) + + // Verify count + count, err := CountPinsBySession(db, s.ID) + require.NoError(t, err) + assert.Equal(t, int64(3), count) +} diff --git a/hub-server/internal/repository/notification_test.go b/hub-server/internal/repository/notification_test.go new file mode 100644 index 000000000..5e9590eee --- /dev/null +++ b/hub-server/internal/repository/notification_test.go @@ -0,0 +1,73 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// ============================================================================= +// Notification repository tests +// ============================================================================= + +func TestNotificationRepo_CreateAndList(t *testing.T) { + db := setupSQLite(t) + + n1 := &model.Notification{UserID: "user-n1", Type: model.TypeMention, Payload: `{"key":"1"}`} + n2 := &model.Notification{UserID: "user-n1", Type: model.TypeSystem, Payload: `{"key":"2"}`} + n3 := &model.Notification{UserID: "user-n2", Type: model.TypeFriendRequest, Payload: `{"key":"3"}`} + require.NoError(t, CreateNotification(db, n1)) + require.NoError(t, CreateNotification(db, n2)) + require.NoError(t, CreateNotification(db, n3)) + + // List all for user-n1 + result, err := ListNotifications(db, "user-n1", false, 10, 0) + require.NoError(t, err) + assert.Len(t, result, 2) + + // Mark first as read + require.NoError(t, MarkNotificationRead(db, "user-n1", n1.ID)) + + // List unread only + result, err = ListNotifications(db, "user-n1", true, 10, 0) + require.NoError(t, err) + assert.Len(t, result, 1) + assert.Equal(t, n2.ID, result[0].ID) +} + +func TestNotificationRepo_MarkReadRequiresOwningUser(t *testing.T) { + db := setupSQLite(t) + + n := &model.Notification{UserID: "user-owner", Type: model.TypeMention, Payload: `{}`} + require.NoError(t, CreateNotification(db, n)) + + err := MarkNotificationRead(db, "user-other", n.ID) + require.ErrorIs(t, err, gorm.ErrRecordNotFound) + + unread, err := ListNotifications(db, "user-owner", true, 10, 0) + require.NoError(t, err) + require.Len(t, unread, 1) + assert.Equal(t, n.ID, unread[0].ID) +} + +func TestNotificationRepo_MarkAllRead(t *testing.T) { + db := setupSQLite(t) + + n1 := &model.Notification{UserID: "user-all", Type: model.TypeMention, Payload: `{}`} + n2 := &model.Notification{UserID: "user-all", Type: model.TypeSystem, Payload: `{}`} + require.NoError(t, CreateNotification(db, n1)) + require.NoError(t, CreateNotification(db, n2)) + + unread, err := ListNotifications(db, "user-all", true, 10, 0) + require.NoError(t, err) + assert.Len(t, unread, 2) + + require.NoError(t, MarkAllNotificationsRead(db, "user-all")) + + unread, err = ListNotifications(db, "user-all", true, 10, 0) + require.NoError(t, err) + assert.Len(t, unread, 0) +} diff --git a/hub-server/internal/repository/refresh_token_test.go b/hub-server/internal/repository/refresh_token_test.go new file mode 100644 index 000000000..be798d532 --- /dev/null +++ b/hub-server/internal/repository/refresh_token_test.go @@ -0,0 +1,93 @@ +package repository + +import ( + "testing" + "time" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// ============================================================================= +// RefreshToken repository tests +// ============================================================================= + +func TestRefreshTokenRepo_UpsertAndGet(t *testing.T) { + db := setupSQLite(t) + + expiresAt := time.Now().Add(24 * time.Hour) + rt := &model.RefreshToken{ + UserID: "user-rt", + DeviceType: "desktop", + DeviceID: "dev-rt-1", + TokenHash: "hash123abc", + ExpiresAt: expiresAt, + } + err := UpsertRefreshToken(db, rt) + require.NoError(t, err) + assert.NotEmpty(t, rt.ID) + + // Find by hash + found, err := FindRefreshTokenByHash(db, "hash123abc") + require.NoError(t, err) + assert.Equal(t, "user-rt", found.UserID) + + // Non-existent hash + _, err = FindRefreshTokenByHash(db, "nonexistent") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + + // Upsert again (same user_id + device_type + device_id) + rt2 := &model.RefreshToken{ + UserID: "user-rt", + DeviceType: "desktop", + DeviceID: "dev-rt-1", + TokenHash: "hash456def", + ExpiresAt: expiresAt.Add(time.Hour), + } + err = UpsertRefreshToken(db, rt2) + require.NoError(t, err) + // Should have same ID (updated existing) + assert.Equal(t, rt.ID, rt2.ID) + + // Find by new hash + found, err = FindRefreshTokenByHash(db, "hash456def") + require.NoError(t, err) + assert.Equal(t, rt.ID, found.ID) + + // Old hash no longer exists (overwritten) + _, err = FindRefreshTokenByHash(db, "hash123abc") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) +} + +func TestRefreshTokenRepo_Revoke(t *testing.T) { + db := setupSQLite(t) + + expiresAt := time.Now().Add(24 * time.Hour) + rt1 := &model.RefreshToken{UserID: "user-rev", DeviceType: "desktop", DeviceID: "dev-1", TokenHash: "h1", ExpiresAt: expiresAt} + rt2 := &model.RefreshToken{UserID: "user-rev", DeviceType: "desktop", DeviceID: "dev-2", TokenHash: "h2", ExpiresAt: expiresAt} + require.NoError(t, UpsertRefreshToken(db, rt1)) + require.NoError(t, UpsertRefreshToken(db, rt2)) + + // Revoke by device + require.NoError(t, RevokeRefreshTokensByUserDevice(db, "user-rev", "dev-1")) + + found, err := FindRefreshTokenByHash(db, "h1") + require.NoError(t, err) + assert.True(t, found.Revoked) + + // rt2 still not revoked + found, err = FindRefreshTokenByHash(db, "h2") + require.NoError(t, err) + assert.False(t, found.Revoked) + + // Revoke all for user + rt3 := &model.RefreshToken{UserID: "user-all-rev", DeviceType: "mobile", DeviceID: "dev-3", TokenHash: "h3", ExpiresAt: expiresAt} + require.NoError(t, UpsertRefreshToken(db, rt3)) + + require.NoError(t, RevokeAllUserTokens(db, "user-all-rev")) + found, err = FindRefreshTokenByHash(db, "h3") + require.NoError(t, err) + assert.True(t, found.Revoked) +} diff --git a/hub-server/internal/repository/repository_test.go b/hub-server/internal/repository/repository_test.go deleted file mode 100644 index 9ee77d73e..000000000 --- a/hub-server/internal/repository/repository_test.go +++ /dev/null @@ -1,2221 +0,0 @@ -package repository - -import ( - "errors" - "strings" - "testing" - "time" - - "github.com/agenthub/hub-server/internal/model" - "github.com/glebarez/sqlite" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "gorm.io/gorm" - gormlogger "gorm.io/gorm/logger" -) - -func TestWrapNotFound(t *testing.T) { - mappedErr := errors.New("mapped not found") - - assert.NoError(t, WrapNotFound(nil, mappedErr)) - assert.ErrorIs(t, WrapNotFound(gorm.ErrRecordNotFound, mappedErr), mappedErr) - assert.ErrorIs(t, WrapNotFound(errors.Join(errors.New("lookup failed"), gorm.ErrRecordNotFound), mappedErr), mappedErr) - - otherErr := errors.New("other") - assert.ErrorIs(t, WrapNotFound(otherErr, mappedErr), otherErr) -} - -func TestListDevicesByUserOrdersMostRecentFirst(t *testing.T) { - db := setupSQLite(t) - now := time.Now() - require.NoError(t, CreateUser(db, &model.User{ID: "user-a", Username: "user-a", Nickname: "User A"})) - require.NoError(t, CreateUser(db, &model.User{ID: "user-b", Username: "user-b", Nickname: "User B"})) - require.NoError(t, db.Create(&model.Device{ - ID: "device-old", - UserID: "user-a", - DeviceType: "desktop", - Capabilities: "[]", - LastActiveAt: now.Add(-time.Hour), - }).Error) - require.NoError(t, db.Create(&model.Device{ - ID: "device-new", - UserID: "user-a", - DeviceType: "desktop", - Capabilities: "[]", - LastActiveAt: now, - }).Error) - require.NoError(t, db.Create(&model.Device{ - ID: "device-other", - UserID: "user-b", - DeviceType: "desktop", - Capabilities: "[]", - LastActiveAt: now.Add(time.Hour), - }).Error) - - devices, err := ListDevicesByUser(db, "user-a") - require.NoError(t, err) - require.Len(t, devices, 2) - assert.Equal(t, "device-new", devices[0].ID) - assert.Equal(t, "device-old", devices[1].ID) -} - -func TestUpdateDeviceAppliesProvidedFields(t *testing.T) { - db := setupSQLite(t) - oldLastActive := time.Now().Add(-time.Hour) - newLastActive := time.Now() - require.NoError(t, db.Create(&model.Device{ - ID: "device-update", - UserID: "user-a", - DeviceType: "desktop", - AppVersion: "1.0.0", - Capabilities: `["shell"]`, - LastActiveAt: oldLastActive, - }).Error) - - err := UpdateDevice(db, "device-update", map[string]interface{}{ - "app_version": "1.1.0", - "capabilities": `["shell","browser"]`, - "last_active_at": newLastActive, - }) - require.NoError(t, err) - - device, err := GetDeviceByID(db, "device-update") - require.NoError(t, err) - assert.Equal(t, "1.1.0", device.AppVersion) - assert.Equal(t, `["shell","browser"]`, device.Capabilities) - assert.WithinDuration(t, newLastActive, device.LastActiveAt, time.Second) -} - -func TestDeleteDeviceRemovesOnlyTargetDevice(t *testing.T) { - db := setupSQLite(t) - now := time.Now() - require.NoError(t, db.Create(&model.Device{ - ID: "device-delete", - UserID: "user-a", - DeviceType: "desktop", - Capabilities: "[]", - LastActiveAt: now, - }).Error) - require.NoError(t, db.Create(&model.Device{ - ID: "device-keep", - UserID: "user-a", - DeviceType: "mobile", - Capabilities: "[]", - LastActiveAt: now, - }).Error) - - require.NoError(t, DeleteDevice(db, "device-delete")) - - _, err := GetDeviceByID(db, "device-delete") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - kept, err := GetDeviceByID(db, "device-keep") - require.NoError(t, err) - assert.Equal(t, "device-keep", kept.ID) -} - -// setupSQLite creates an in-memory SQLite database with tables matching the -// production PostgreSQL schema. Raw SQL is used instead of AutoMigrate because -// GORM's SQLite driver mishandles PostgreSQL-specific GORM tags (jsonb with -// default:'[]' produces SQLite-invalid DEFAULT "[]"). -func setupSQLite(t *testing.T) *gorm.DB { - t.Helper() - db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ - Logger: gormlogger.Default.LogMode(gormlogger.Silent), - }) - require.NoError(t, err) - - tables := []string{ - `CREATE TABLE users ( - id TEXT PRIMARY KEY, - username TEXT NOT NULL UNIQUE, - password_hash TEXT, - nickname TEXT NOT NULL, - avatar_url TEXT DEFAULT '', - tokendance_sub TEXT DEFAULT NULL, - tokendance_sub_linked_at DATETIME DEFAULT NULL, - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE UNIQUE INDEX idx_users_tokendance_sub ON users(tokendance_sub) - WHERE tokendance_sub IS NOT NULL AND tokendance_sub != ''`, - `CREATE TABLE devices ( - id TEXT PRIMARY KEY, - user_id TEXT NOT NULL, - device_type TEXT NOT NULL, - app_version TEXT DEFAULT '', - capabilities TEXT DEFAULT '[]', - last_active_at DATETIME NOT NULL DEFAULT (datetime('now')), - created_at DATETIME - )`, - `CREATE INDEX idx_devices_user_type ON devices(user_id, device_type)`, - `CREATE TABLE sessions ( - id TEXT PRIMARY KEY, - type TEXT NOT NULL, - name TEXT DEFAULT '', - avatar_url TEXT DEFAULT '', - tokendance_sub TEXT DEFAULT NULL, - tokendance_sub_linked_at DATETIME DEFAULT NULL, - announcement TEXT DEFAULT '', - owner_user_id TEXT, - workspace_id TEXT, - next_seq INTEGER NOT NULL DEFAULT 0, - last_message_at DATETIME, - dissolved INTEGER NOT NULL DEFAULT 0, - created_at DATETIME - )`, - `CREATE TABLE session_members ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - member_type TEXT NOT NULL, - member_id TEXT NOT NULL, - role TEXT NOT NULL, - pinned INTEGER NOT NULL DEFAULT 0, - archived INTEGER NOT NULL DEFAULT 0, - muted INTEGER NOT NULL DEFAULT 0, - last_read_seq INTEGER NOT NULL DEFAULT 0, - joined_at DATETIME, - left_at DATETIME - )`, - `CREATE TABLE messages ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - seq_id INTEGER NOT NULL, - client_msg_id TEXT NOT NULL, - sender_type TEXT NOT NULL, - sender_id TEXT NOT NULL, - content_type TEXT NOT NULL, - content TEXT NOT NULL DEFAULT '', - reply_to_message_id TEXT, - recalled INTEGER NOT NULL DEFAULT 0, - edited BOOLEAN NOT NULL DEFAULT FALSE, - edited_at DATETIME, - created_at DATETIME - )`, - `CREATE TABLE message_pins ( - session_id TEXT NOT NULL, - message_id TEXT NOT NULL, - pinned_by_user_id TEXT NOT NULL, - pinned_at DATETIME, - PRIMARY KEY (session_id, message_id) - )`, - `CREATE TABLE message_attachments ( - session_id TEXT NOT NULL, - message_id TEXT NOT NULL, - attachment_id TEXT NOT NULL, - created_at DATETIME, - PRIMARY KEY (message_id, attachment_id) - )`, - `CREATE TABLE message_reactions ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - message_id TEXT NOT NULL, - user_id TEXT NOT NULL, - emoji TEXT NOT NULL, - created_at DATETIME, - UNIQUE (session_id, message_id, user_id, emoji) - )`, - `CREATE TABLE friendships ( - id TEXT PRIMARY KEY, - user_id TEXT NOT NULL, - friend_id TEXT NOT NULL, - status TEXT NOT NULL, - remark TEXT DEFAULT '', - request_message TEXT DEFAULT '', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE UNIQUE INDEX idx_friendships_user_friend ON friendships(user_id, friend_id)`, - // Additional models for full coverage - `CREATE TABLE notifications ( - id TEXT PRIMARY KEY, - user_id TEXT NOT NULL, - type TEXT NOT NULL, - payload TEXT NOT NULL DEFAULT '', - read INTEGER NOT NULL DEFAULT 0, - created_at DATETIME - )`, - `CREATE TABLE attachments ( - id TEXT PRIMARY KEY, - hash TEXT NOT NULL UNIQUE, - size INTEGER NOT NULL, - mime_type TEXT NOT NULL, - original_name TEXT DEFAULT '', - uploader_user_id TEXT NOT NULL, - metadata TEXT NOT NULL DEFAULT '{}', - created_at DATETIME - )`, - `CREATE TABLE agent_instances ( - id TEXT PRIMARY KEY, - agent_type TEXT NOT NULL, - custom_agent_id TEXT, - session_id TEXT NOT NULL, - inviter_user_id TEXT NOT NULL, - workspace_id TEXT, - display_name TEXT NOT NULL, - created_at DATETIME - )`, - `CREATE TABLE custom_agents ( - id TEXT PRIMARY KEY, - owner_user_id TEXT NOT NULL, - name TEXT NOT NULL, - avatar_url TEXT DEFAULT '', - tokendance_sub TEXT DEFAULT NULL, - tokendance_sub_linked_at DATETIME DEFAULT NULL, - agent_type TEXT NOT NULL, - system_prompt TEXT NOT NULL DEFAULT '', - capability_tags TEXT DEFAULT '[]', - tool_whitelist TEXT DEFAULT '[]', - model_params TEXT DEFAULT '{}', - output_schema TEXT DEFAULT NULL, - deleted_at DATETIME, - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE pending_agent_tasks ( - id TEXT PRIMARY KEY, - agent_instance_id TEXT NOT NULL, - triggered_by_user_id TEXT NOT NULL, - trigger_message_id TEXT NOT NULL, - target_id TEXT, - status TEXT NOT NULL, - edge_run_id TEXT DEFAULT '', - edge_device_id TEXT DEFAULT '', - error_message TEXT DEFAULT '', - model_params TEXT DEFAULT '{}', - created_at DATETIME, - dispatched_at DATETIME, - finished_at DATETIME, - expire_at DATETIME NOT NULL - )`, - `CREATE TABLE refresh_tokens ( - id TEXT PRIMARY KEY, - user_id TEXT NOT NULL, - device_type TEXT NOT NULL DEFAULT '', - device_id TEXT NOT NULL DEFAULT '', - token_hash TEXT NOT NULL UNIQUE, - expires_at DATETIME NOT NULL, - revoked INTEGER NOT NULL DEFAULT 0, - created_at DATETIME - )`, - `CREATE UNIQUE INDEX idx_rt_user_device ON refresh_tokens(user_id, device_type, device_id)`, - `CREATE TABLE agent_teams ( - id TEXT PRIMARY KEY, - owner_id TEXT NOT NULL, - name TEXT NOT NULL, - description TEXT DEFAULT '', - avatar_url TEXT DEFAULT '', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_members ( - id TEXT PRIMARY KEY, - team_id TEXT NOT NULL, - agent_profile_id TEXT, - role TEXT NOT NULL DEFAULT 'executor', - position INTEGER NOT NULL DEFAULT 0, - created_at DATETIME, - FOREIGN KEY (team_id) REFERENCES agent_teams(id) ON DELETE CASCADE - )`, - `CREATE TABLE agent_team_runs ( - id TEXT PRIMARY KEY, - team_id TEXT NOT NULL, - session_id TEXT, - trigger_user_id TEXT NOT NULL, - trigger_message TEXT DEFAULT '', - target_id TEXT, - mode TEXT NOT NULL DEFAULT 'supervisor', - status TEXT NOT NULL DEFAULT 'queued', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_assignments ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - from_member_id TEXT NOT NULL, - to_member_id TEXT NOT NULL, - type TEXT NOT NULL DEFAULT 'delegate', - task_prompt TEXT NOT NULL, - context TEXT DEFAULT '', - status TEXT NOT NULL DEFAULT 'pending', - run_id TEXT, - result TEXT DEFAULT '', - depth INTEGER NOT NULL DEFAULT 0, - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_tasks ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - assignment_id TEXT, - assignee_member_id TEXT NOT NULL, - parent_task_id TEXT, - status TEXT NOT NULL DEFAULT 'pending', - objective TEXT NOT NULL, - input_refs TEXT NOT NULL DEFAULT '{}', - run_id TEXT, - attempt INTEGER NOT NULL DEFAULT 1, - risk_level TEXT NOT NULL DEFAULT 'normal', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_events ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - seq INTEGER NOT NULL, - type TEXT NOT NULL, - payload TEXT NOT NULL DEFAULT '{}', - created_at DATETIME - )`, - // Mirrors migration 0056: seq must be unique per run so concurrent - // appends surface as unique violations instead of silent duplicates. - `CREATE UNIQUE INDEX uq_agent_team_events_run_seq ON agent_team_events(team_run_id, seq)`, - `CREATE TABLE agent_profiles ( - id TEXT PRIMARY KEY, - owner_id TEXT NOT NULL, - name TEXT NOT NULL, - description TEXT DEFAULT '', - runtime_id TEXT NOT NULL DEFAULT '', - model TEXT DEFAULT '', - provider TEXT DEFAULT '', - reasoning_effort TEXT DEFAULT 'medium', - model_mapping TEXT DEFAULT '{}', - skills TEXT DEFAULT '[]', - mcp_servers TEXT DEFAULT '[]', - tool_allowlist TEXT DEFAULT '[]', - approval_policy TEXT DEFAULT '{}', - permission_mode TEXT DEFAULT 'default', - target_preferences TEXT DEFAULT '{}', - context_budget_max_tokens INTEGER DEFAULT 200000, - is_public INTEGER DEFAULT 0, - install_count INTEGER DEFAULT 0, - rating_avg REAL DEFAULT 0, - rating_count INTEGER DEFAULT 0, - version INTEGER DEFAULT 1, - created_at DATETIME, - updated_at DATETIME, - deleted_at DATETIME - )`, - `CREATE TABLE agent_run_events ( - id TEXT PRIMARY KEY, - task_id TEXT NOT NULL, - edge_run_id TEXT DEFAULT '', - session_id TEXT NOT NULL, - agent_instance_id TEXT NOT NULL, - event_seq INTEGER NOT NULL, - event_type TEXT NOT NULL, - payload TEXT NOT NULL DEFAULT '', - created_at DATETIME - )`, - `CREATE TABLE agent_team_artifacts ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - team_task_id TEXT, - assignment_id TEXT, - member_id TEXT, - agent_task_id TEXT, - edge_run_id TEXT DEFAULT '', - source_event_id TEXT, - event_seq INTEGER NOT NULL DEFAULT 0, - path TEXT NOT NULL, - normalized_path TEXT NOT NULL, - action TEXT DEFAULT '', - tool_name TEXT DEFAULT '', - status TEXT DEFAULT '', - conflict_id TEXT, - created_at DATETIME, - updated_at DATETIME - )`, - } - for _, ddl := range tables { - require.NoError(t, db.Exec(ddl).Error, "DDL: %s", ddl[:60]) - } - return db -} - -func TestUpdateMessageContentMarksEditedAndStoresTimestamp(t *testing.T) { - db := setupSQLite(t) - require.NoError(t, db.Exec(`INSERT INTO messages ( - id, session_id, seq_id, client_msg_id, sender_type, sender_id, content_type, content, recalled, created_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - "msg-edit", "sess-edit", 1, "client-edit", "user", "user-1", "text", `{"text":"before"}`, false, time.Now(), - ).Error) - - require.NoError(t, UpdateMessageContent(db, "msg-edit", "text", `{"text":"after"}`)) - - var stored struct { - ContentType string - Content string - Edited bool - EditedAt *time.Time - } - require.NoError(t, db.Table("messages").Where("id = ?", "msg-edit").First(&stored).Error) - assert.Equal(t, "text", stored.ContentType) - assert.JSONEq(t, `{"text":"after"}`, stored.Content) - assert.True(t, stored.Edited) - if stored.EditedAt == nil { - t.Fatal("edited_at was not set") - } -} - -func TestUpdateMessageContentReturnsNotFoundWhenNoRowsChange(t *testing.T) { - db := setupSQLite(t) - require.ErrorIs(t, UpdateMessageContent(db, "missing", "text", `{"text":"after"}`), gorm.ErrRecordNotFound) -} - -// ============================================================================= -// User repository tests -// ============================================================================= - -func TestUserRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - user := &model.User{ - Username: "testuser", - PasswordHash: strPtr("hashed_password"), - Nickname: "Test User", - } - - // Create - err := CreateUser(db, user) - require.NoError(t, err) - assert.NotEmpty(t, user.ID) - - // Read by ID - fetched, err := GetUserByID(db, user.ID) - require.NoError(t, err) - assert.Equal(t, user.Username, fetched.Username) - assert.Equal(t, user.Nickname, fetched.Nickname) - - // Read by username - fetchedByUsername, err := GetUserByUsername(db, "testuser") - require.NoError(t, err) - assert.Equal(t, user.ID, fetchedByUsername.ID) - - // Read non-existent - _, err = GetUserByID(db, "nonexistent-id") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - - // Update - user.Nickname = "Updated Name" - err = UpdateUser(db, user) - require.NoError(t, err) - fetched, err = GetUserByID(db, user.ID) - require.NoError(t, err) - assert.Equal(t, "Updated Name", fetched.Nickname) -} - -func TestUserRepo_GetUsersByIDs(t *testing.T) { - db := setupSQLite(t) - - // Create multiple users - u1 := &model.User{Username: "user1", PasswordHash: strPtr("h1"), Nickname: "U1"} - u2 := &model.User{Username: "user2", PasswordHash: strPtr("h2"), Nickname: "U2"} - u3 := &model.User{Username: "user3", PasswordHash: strPtr("h3"), Nickname: "U3"} - require.NoError(t, CreateUser(db, u1)) - require.NoError(t, CreateUser(db, u2)) - require.NoError(t, CreateUser(db, u3)) - - // Fetch by IDs - m, err := GetUsersByIDs(db, []string{u1.ID, u2.ID}) - require.NoError(t, err) - assert.Len(t, m, 2) - assert.Equal(t, "U1", m[u1.ID].Nickname) - assert.Equal(t, "U2", m[u2.ID].Nickname) - - // Empty list - m, err = GetUsersByIDs(db, []string{}) - require.NoError(t, err) - assert.Empty(t, m) - - // Non-existent IDs - m, err = GetUsersByIDs(db, []string{"no-such-id"}) - require.NoError(t, err) - assert.Empty(t, m) -} - -func TestUserRepo_FindOrCreateByTokenDanceSubCreatesStableDistinctHubUsers(t *testing.T) { - db := setupSQLite(t) - - sub1 := "tokendance-subject-with-a-very-long-shared-prefix-" + strings.Repeat("1", 80) - sub2 := "tokendance-subject-with-a-very-long-shared-prefix-" + strings.Repeat("2", 80) - - u1, err := FindOrCreateByTokenDanceSub(db, sub1, "", "") - require.NoError(t, err) - u2, err := FindOrCreateByTokenDanceSub(db, sub2, "", "") - require.NoError(t, err) - - require.NotEqual(t, u1.ID, u2.ID) - require.NotEqual(t, u1.Username, u2.Username) - assert.True(t, strings.HasPrefix(u1.Username, "td_")) - assert.True(t, len(u1.Username) <= 32) - assert.True(t, len(u1.Nickname) <= 64) - require.NotNil(t, u1.TokenDanceSub) - assert.Equal(t, sub1, *u1.TokenDanceSub) - - again, err := FindOrCreateByTokenDanceSub(db, sub1, "", "") - require.NoError(t, err) - assert.Equal(t, u1.ID, again.ID) - assert.Equal(t, u1.Username, again.Username) -} - -func TestUserRepo_ReadsOIDCOnlyUserWithNullPasswordHash(t *testing.T) { - db := setupSQLite(t) - - sub := "tokendance-sub-null-password" - now := time.Now().UTC() - require.NoError(t, db.Exec(` - INSERT INTO users ( - id, username, password_hash, nickname, tokendance_sub, tokendance_sub_linked_at, created_at, updated_at - ) VALUES (?, ?, NULL, ?, ?, ?, ?, ?) - `, "user-oidc-null-password", "td_null_password", "OIDC Only", sub, now, now, now).Error) - - bySub, err := FindByTokenDanceSub(db, sub) - require.NoError(t, err) - require.NotNil(t, bySub) - assert.Equal(t, "user-oidc-null-password", bySub.ID) - assert.Nil(t, bySub.PasswordHash) - - byID, err := GetUserByID(db, "user-oidc-null-password") - require.NoError(t, err) - require.NotNil(t, byID) - assert.Equal(t, sub, *byID.TokenDanceSub) - assert.Nil(t, byID.PasswordHash) -} - -// ============================================================================= -// Device repository tests -// ============================================================================= - -func TestDeviceRepo_Upsert(t *testing.T) { - db := setupSQLite(t) - - device := &model.Device{ - ID: "dev-001", - UserID: "user-001", - DeviceType: "desktop", - AppVersion: "1.0.0", - Capabilities: `["chat","agent"]`, - } - - // First upsert: creates - err := UpsertDevice(db, device) - require.NoError(t, err) - - fetched, err := GetDeviceByID(db, "dev-001") - require.NoError(t, err) - assert.Equal(t, "desktop", fetched.DeviceType) - assert.Equal(t, "1.0.0", fetched.AppVersion) - - // Second upsert: same physical device updates by device ID. - device2 := &model.Device{ - ID: "dev-001", - UserID: "user-001", - DeviceType: "desktop", - AppVersion: "2.0.0", - Capabilities: `["chat","agent","file"]`, - } - err = UpsertDevice(db, device2) - require.NoError(t, err) - - // ON CONFLICT preserves the original row's ID but updates other columns. - // Verify the original row was updated. - fetched, err = GetDeviceByID(db, "dev-001") - require.NoError(t, err) - assert.Equal(t, "2.0.0", fetched.AppVersion) - - // A second desktop for the same user is a distinct device and must keep its - // own row so refresh_tokens.device_id can reference it. - device3 := &model.Device{ - ID: "dev-001-second-desktop", - UserID: "user-001", - DeviceType: "desktop", - AppVersion: "1.0.0", - Capabilities: `["chat"]`, - } - require.NoError(t, UpsertDevice(db, device3)) - fetched, err = GetDeviceByID(db, "dev-001-second-desktop") - require.NoError(t, err) - assert.Equal(t, "user-001", fetched.UserID) - assert.Equal(t, "desktop", fetched.DeviceType) - - stolen := &model.Device{ - ID: "dev-001", - UserID: "user-attacker", - DeviceType: "desktop", - AppVersion: "9.9.9", - Capabilities: `[]`, - } - require.Error(t, UpsertDevice(db, stolen)) - fetched, err = GetDeviceByID(db, "dev-001") - require.NoError(t, err) - assert.Equal(t, "user-001", fetched.UserID) - assert.Equal(t, "2.0.0", fetched.AppVersion) -} - -func TestDeviceRepo_GetByID(t *testing.T) { - db := setupSQLite(t) - - device := &model.Device{ - ID: "dev-002", - UserID: "user-002", - DeviceType: "mobile", - } - require.NoError(t, UpsertDevice(db, device)) - - fetched, err := GetDeviceByID(db, "dev-002") - require.NoError(t, err) - assert.Equal(t, "user-002", fetched.UserID) - assert.Equal(t, "mobile", fetched.DeviceType) - - _, err = GetDeviceByID(db, "nonexistent") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) -} - -// ============================================================================= -// Session repository tests -// ============================================================================= - -func TestSessionRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - session := &model.Session{ - Type: model.SessionTypePrivate, - Name: "Test Session", - OwnerUserID: strPtr("user-001"), - } - - // Create - err := CreateSession(db, session) - require.NoError(t, err) - assert.NotEmpty(t, session.ID) - - // Read - fetched, err := GetSessionByID(db, session.ID) - require.NoError(t, err) - assert.Equal(t, model.SessionTypePrivate, fetched.Type) - assert.Equal(t, "Test Session", fetched.Name) - - // Update - fetched.Name = "Updated Session" - err = UpdateSessionColumns(db, fetched, "name") - require.NoError(t, err) - - fetched2, err := GetSessionByID(db, session.ID) - require.NoError(t, err) - assert.Equal(t, "Updated Session", fetched2.Name) -} - -func TestSessionRepo_FindPrivateSessionBetween(t *testing.T) { - db := setupSQLite(t) - - session := &model.Session{ - Type: model.SessionTypePrivate, - } - require.NoError(t, CreateSession(db, session)) - - // Add both members - m1 := &model.SessionMember{ - SessionID: session.ID, - MemberType: model.MemberTypeUser, - MemberID: "user-a", - Role: model.MemberRoleMember, - } - m2 := &model.SessionMember{ - SessionID: session.ID, - MemberType: model.MemberTypeUser, - MemberID: "user-b", - Role: model.MemberRoleMember, - } - require.NoError(t, CreateSessionMember(db, m1)) - require.NoError(t, CreateSessionMember(db, m2)) - - // Find between - found, err := FindPrivateSessionBetween(db, "user-a", "user-b") - require.NoError(t, err) - require.NotNil(t, found) - assert.Equal(t, session.ID, found.ID) - - // Not found - found, err = FindPrivateSessionBetween(db, "user-a", "user-c") - require.NoError(t, err) - assert.Nil(t, found) -} - -func TestSessionRepo_TouchLastMessage(t *testing.T) { - db := setupSQLite(t) - - session := &model.Session{Type: model.SessionTypeGroup, Name: "Group"} - require.NoError(t, CreateSession(db, session)) - - // Initially nil - fetched, err := GetSessionByID(db, session.ID) - require.NoError(t, err) - assert.Nil(t, fetched.LastMessageAt) - - err = TouchSessionLastMessage(db, session.ID) - require.NoError(t, err) - - fetched, err = GetSessionByID(db, session.ID) - require.NoError(t, err) - assert.NotNil(t, fetched.LastMessageAt) -} - -func TestSessionRepo_ListUserSessions(t *testing.T) { - // member_count uses a correlated scalar subquery compatible with both PG and SQLite. - // Plan-level assertions are in session_member_count_test.go (PG-only EXPLAIN). - db := setupSQLite(t) - - s := &model.Session{Type: model.SessionTypeGroup, Name: "ListGroup"} - require.NoError(t, CreateSession(db, s)) - - m := &model.SessionMember{ - SessionID: s.ID, - MemberType: model.MemberTypeUser, - MemberID: "user-l", - Role: model.MemberRoleOwner, - } - require.NoError(t, CreateSessionMember(db, m)) - - result, err := ListUserSessions(db, "user-l") - require.NoError(t, err) - assert.Len(t, result, 1) - assert.Equal(t, s.ID, result[0].ID) - assert.Equal(t, model.MemberRoleOwner, result[0].Role) -} - -// ============================================================================= -// Message repository tests -// ============================================================================= - -func createTestSession(t *testing.T, db *gorm.DB) *model.Session { - t.Helper() - s := &model.Session{Type: model.SessionTypeGroup, Name: "MsgTest"} - require.NoError(t, CreateSession(db, s)) - return s -} - -func TestMessageRepo_InsertAndGet(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - msg := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "client-001", - SenderType: model.SenderTypeUser, - SenderID: "user-1", - ContentType: model.ContentTypeText, - Content: `{"text":"Hello"}`, - } - - err := InsertMessage(db, msg) - require.NoError(t, err) - assert.NotEmpty(t, msg.ID) - - fetched, err := GetMessageByID(db, msg.ID) - require.NoError(t, err) - assert.Equal(t, `{"text":"Hello"}`, fetched.Content) -} - -func TestMessageRepo_GetBySession(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - for i := 1; i <= 5; i++ { - msg := &model.Message{ - SessionID: s.ID, - SeqID: int64(i), - ClientMsgID: "client-" + string(rune('0'+i)), - SenderType: model.SenderTypeUser, - SenderID: "user-1", - ContentType: model.ContentTypeText, - Content: `{"text":"Message ` + string(rune('0'+i)) + `"}`, - } - require.NoError(t, InsertMessage(db, msg)) - } - - msgs, err := GetMessagesBySession(db, s.ID, 0, 10) - require.NoError(t, err) - assert.Len(t, msgs, 5) - - // Get with beforeSeq - msgs, err = GetMessagesBySession(db, s.ID, 4, 10) - require.NoError(t, err) - assert.Len(t, msgs, 3) // seq 1,2,3 (before seq 4) - - // Get with small limit - msgs, err = GetMessagesBySession(db, s.ID, 0, 2) - require.NoError(t, err) - assert.Len(t, msgs, 2) -} - -func TestMessageRepo_Increment(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - for i := 1; i <= 5; i++ { - msg := &model.Message{ - SessionID: s.ID, - SeqID: int64(i), - ClientMsgID: "inc-client-" + string(rune('0'+i)), - SenderType: model.SenderTypeUser, - SenderID: "user-1", - ContentType: model.ContentTypeText, - Content: `{"text":"Inc ` + string(rune('0'+i)) + `"}`, - } - require.NoError(t, InsertMessage(db, msg)) - } - - msgs, err := GetMessagesIncrement(db, s.ID, 2, 10) - require.NoError(t, err) - assert.Len(t, msgs, 3) // seq 3,4,5 (after seq 2) - assert.Equal(t, int64(3), msgs[0].SeqID) -} - -func TestMessageRepo_Recall(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - msg := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "recall-001", - SenderType: model.SenderTypeUser, - SenderID: "user-1", - ContentType: model.ContentTypeText, - Content: `{"text":"Recall me"}`, - } - require.NoError(t, InsertMessage(db, msg)) - - err := UpdateMessageRecalled(db, msg.ID) - require.NoError(t, err) - - fetched, err := GetMessageByID(db, msg.ID) - require.NoError(t, err) - assert.True(t, fetched.Recalled) -} - -func TestMessageRepo_DuplicateClientMsgID(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - msg := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "dup-client", - SenderType: model.SenderTypeUser, - SenderID: "user-1", - ContentType: model.ContentTypeText, - Content: `{"text":"First"}`, - } - require.NoError(t, InsertMessage(db, msg)) - - fetched, err := GetMessageByClientMsgID(db, s.ID, "dup-client") - require.NoError(t, err) - require.NotNil(t, fetched) - assert.Equal(t, `{"text":"First"}`, fetched.Content) - - // Non-existent returns nil, nil - fetched, err = GetMessageByClientMsgID(db, s.ID, "no-such") - require.NoError(t, err) - assert.Nil(t, fetched) -} - -func TestMessageRepo_Pins(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - pin := &model.MessagePin{ - SessionID: s.ID, - MessageID: "msg-001", - PinnedByUserID: "user-1", - } - - err := InsertPin(db, pin) - require.NoError(t, err) - - count, err := CountPinsBySession(db, s.ID) - require.NoError(t, err) - assert.Equal(t, int64(1), count) - - pins, err := ListPinsBySession(db, s.ID) - require.NoError(t, err) - assert.Len(t, pins, 1) - assert.Equal(t, "msg-001", pins[0].MessageID) - - // Delete pin - err = DeletePin(db, s.ID, "msg-001") - require.NoError(t, err) - - count, err = CountPinsBySession(db, s.ID) - require.NoError(t, err) - assert.Equal(t, int64(0), count) -} - -func TestMessageReactionRepo_AddListCountAndRemove(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - msg := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "reaction-client-1", - SenderType: model.SenderTypeUser, - SenderID: "sender-1", - ContentType: model.ContentTypeText, - Content: `{"text":"hello"}`, - } - require.NoError(t, InsertMessage(db, msg)) - - reaction := &model.MessageReaction{ - SessionID: s.ID, - MessageID: msg.ID, - UserID: "user-1", - Reaction: "thumbs_up", - } - require.NoError(t, AddReaction(db, reaction)) - require.NotEmpty(t, reaction.ID) - - reactions, err := ListReactionsByMessage(db, s.ID, msg.ID) - require.NoError(t, err) - require.Len(t, reactions, 1) - assert.Equal(t, "thumbs_up", reactions[0].Reaction) - assert.Equal(t, "user-1", reactions[0].UserID) - - byMessage, err := ListReactionsByMessages(db, s.ID, []string{msg.ID, "missing-message"}) - require.NoError(t, err) - require.Len(t, byMessage[msg.ID], 1) - assert.Empty(t, byMessage["missing-message"]) - - counts, err := ReactionCountsByMessage(db, s.ID, []string{msg.ID, "missing-message"}) - require.NoError(t, err) - assert.Equal(t, int64(1), counts[msg.ID]) - assert.Zero(t, counts["missing-message"]) - - summaries, err := ReactionSummariesByMessage(db, s.ID, msg.ID) - require.NoError(t, err) - require.Len(t, summaries, 1) - assert.Equal(t, "thumbs_up", summaries[0].Reaction) - assert.Equal(t, 1, summaries[0].Count) - assert.Equal(t, []string{"user-1"}, summaries[0].UserIDs) - - require.NoError(t, RemoveReaction(db, s.ID, msg.ID, "user-1", "thumbs_up")) - - reactions, err = ListReactionsByMessage(db, s.ID, msg.ID) - require.NoError(t, err) - assert.Empty(t, reactions) -} - -func TestMessageReactionRepo_AddReactionIsIdempotentForDuplicateUserReaction(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - msg := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "reaction-client-duplicate", - SenderType: model.SenderTypeUser, - SenderID: "sender-1", - ContentType: model.ContentTypeText, - Content: `{"text":"hello"}`, - } - require.NoError(t, InsertMessage(db, msg)) - - reaction := model.MessageReaction{ - SessionID: s.ID, - MessageID: msg.ID, - UserID: "user-1", - Reaction: "heart", - } - require.NoError(t, AddReaction(db, &reaction)) - require.NoError(t, AddReaction(db, &model.MessageReaction{ - SessionID: s.ID, - MessageID: msg.ID, - UserID: "user-1", - Reaction: "heart", - })) - - counts, err := ReactionCountsByMessage(db, s.ID, []string{msg.ID}) - require.NoError(t, err) - assert.Equal(t, int64(1), counts[msg.ID]) -} - -func TestMessageReactionRepo_RemoveReactionIsIdempotentForMissingRows(t *testing.T) { - db := setupSQLite(t) - - require.NoError(t, RemoveReaction(db, "missing-session", "missing-message", "missing-user", "heart")) -} - -func TestMessageReactionRepo_EmptyInputsReturnEmptyMaps(t *testing.T) { - db := setupSQLite(t) - - byMessage, err := ListReactionsByMessages(db, "session-1", nil) - require.NoError(t, err) - assert.Empty(t, byMessage) - - counts, err := ReactionCountsByMessage(db, "session-1", nil) - require.NoError(t, err) - assert.Empty(t, counts) - - summaries, err := ReactionSummariesByMessage(db, "session-1", "message-1") - require.NoError(t, err) - assert.Empty(t, summaries) -} - -func TestMessageAttachmentRepo_CreateAndAccess(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - member := &model.SessionMember{ - SessionID: s.ID, - MemberType: model.MemberTypeUser, - MemberID: "viewer-1", - Role: model.MemberRoleMember, - } - require.NoError(t, CreateSessionMember(db, member)) - - msg := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "attach-client-1", - SenderType: model.SenderTypeUser, - SenderID: "owner-1", - ContentType: model.ContentTypeFile, - Content: `{"attachment_id":"att-1"}`, - } - require.NoError(t, InsertMessage(db, msg)) - - refs := []model.MessageAttachment{{ - SessionID: s.ID, - MessageID: msg.ID, - AttachmentID: "att-1", - }} - require.NoError(t, CreateMessageAttachmentReferences(db, refs)) - require.NoError(t, CreateMessageAttachmentReferences(db, refs)) - - allowed, err := CanUserAccessReferencedAttachment(db, "viewer-1", "att-1") - require.NoError(t, err) - assert.True(t, allowed) - - allowed, err = CanUserAccessReferencedAttachment(db, "outsider-1", "att-1") - require.NoError(t, err) - assert.False(t, allowed) - - require.NoError(t, SoftDeleteMember(db, s.ID, model.MemberTypeUser, "viewer-1")) - allowed, err = CanUserAccessReferencedAttachment(db, "viewer-1", "att-1") - require.NoError(t, err) - assert.False(t, allowed) -} - -func TestMessageAttachmentRepo_ListAttachmentsByMessageIDs(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - msgWithAttachment := &model.Message{ - SessionID: s.ID, - SeqID: 1, - ClientMsgID: "attach-client-1", - SenderType: model.SenderTypeUser, - SenderID: "owner-1", - ContentType: model.ContentTypeFile, - Content: `{"attachment_id":"att-1"}`, - } - msgWithoutAttachment := &model.Message{ - SessionID: s.ID, - SeqID: 2, - ClientMsgID: "attach-client-2", - SenderType: model.SenderTypeUser, - SenderID: "owner-1", - ContentType: model.ContentTypeText, - Content: `{"text":"plain"}`, - } - require.NoError(t, InsertMessage(db, msgWithAttachment)) - require.NoError(t, InsertMessage(db, msgWithoutAttachment)) - - require.NoError(t, db.Exec( - `INSERT INTO attachments (id, hash, size, mime_type, original_name, uploader_user_id, metadata, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, - "att-1", "hash-1", 42, "text/plain", "notes.txt", "owner-1", `{"height":3,"width":2}`, time.Now(), - ).Error) - require.NoError(t, CreateMessageAttachmentReferences(db, []model.MessageAttachment{{ - SessionID: s.ID, - MessageID: msgWithAttachment.ID, - AttachmentID: "att-1", - }})) - - attachmentsByMessage, err := ListAttachmentsByMessageIDs(db, []string{msgWithAttachment.ID, msgWithoutAttachment.ID}) - require.NoError(t, err) - - require.Len(t, attachmentsByMessage[msgWithAttachment.ID], 1) - assert.Equal(t, "att-1", attachmentsByMessage[msgWithAttachment.ID][0].ID) - assert.Equal(t, "hash-1", attachmentsByMessage[msgWithAttachment.ID][0].Hash) - assert.Equal(t, int64(42), attachmentsByMessage[msgWithAttachment.ID][0].Size) - assert.Equal(t, "text/plain", attachmentsByMessage[msgWithAttachment.ID][0].MimeType) - assert.Equal(t, "notes.txt", attachmentsByMessage[msgWithAttachment.ID][0].OriginalName) - assert.JSONEq(t, `{"height":3,"width":2}`, attachmentsByMessage[msgWithAttachment.ID][0].Metadata) - assert.Empty(t, attachmentsByMessage[msgWithoutAttachment.ID]) - - empty, err := ListAttachmentsByMessageIDs(db, nil) - require.NoError(t, err) - assert.Empty(t, empty) -} - -func TestMessageRepo_GetByIDs(t *testing.T) { - db := setupSQLite(t) - s := createTestSession(t, db) - - msg1 := &model.Message{SessionID: s.ID, SeqID: 1, ClientMsgID: "c1", SenderType: model.SenderTypeUser, SenderID: "u1", ContentType: model.ContentTypeText, Content: `{}`} - msg2 := &model.Message{SessionID: s.ID, SeqID: 2, ClientMsgID: "c2", SenderType: model.SenderTypeUser, SenderID: "u1", ContentType: model.ContentTypeText, Content: `{}`} - require.NoError(t, InsertMessage(db, msg1)) - require.NoError(t, InsertMessage(db, msg2)) - - msgs, err := GetMessagesByIDs(db, []string{msg1.ID, msg2.ID}) - require.NoError(t, err) - assert.Len(t, msgs, 2) - - // Empty list - msgs, err = GetMessagesByIDs(db, []string{}) - require.NoError(t, err) - assert.Empty(t, msgs) -} - -// ============================================================================= -// Friendship repository tests -// ============================================================================= - -func TestFriendshipRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - f := &model.Friendship{ - UserID: "user-a", - FriendID: "user-b", - Status: model.StatusPending, - RequestMessage: "Please add me", - } - - // Create - err := CreateFriendship(db, f) - require.NoError(t, err) - assert.NotEmpty(t, f.ID) - - // Find between - found, err := FindFriendshipBetween(db, "user-a", "user-b") - require.NoError(t, err) - require.NotNil(t, found) - assert.Equal(t, model.StatusPending, found.Status) - - // Also find reversed - found, err = FindFriendshipBetween(db, "user-b", "user-a") - require.NoError(t, err) - require.NotNil(t, found) -} - -func TestFriendshipRepo_StatusTransitions(t *testing.T) { - db := setupSQLite(t) - - f := &model.Friendship{ - UserID: "user-1", - FriendID: "user-2", - Status: model.StatusPending, - } - require.NoError(t, CreateFriendship(db, f)) - - // Accept - err := UpdateFriendshipByID(db, f.ID, model.StatusAccepted) - require.NoError(t, err) - - fetched, err := GetFriendshipByID(db, f.ID) - require.NoError(t, err) - assert.Equal(t, model.StatusAccepted, fetched.Status) - - // Update remark - err = UpdateFriendshipRemark(db, "user-1", "user-2", "Bestie") - require.NoError(t, err) - - fetched, err = GetFriendshipByID(db, f.ID) - require.NoError(t, err) - assert.Equal(t, "Bestie", fetched.Remark) -} - -func TestFriendshipRepo_UpdateRemarkNoRows(t *testing.T) { - db := setupSQLite(t) - - // No friendship rows exist, so UpdateFriendshipRemark should return ErrRecordNotFound. - err := UpdateFriendshipRemark(db, "user-x", "user-y", "remark") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - - // Create a friendship with pending status - remark should NOT be updatable - f := &model.Friendship{UserID: "user-x", FriendID: "user-y", Status: model.StatusPending} - require.NoError(t, CreateFriendship(db, f)) - - err = UpdateFriendshipRemark(db, "user-x", "user-y", "remark") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound, "pending friendship should not allow remark update") - - // Update to accepted - remark SHOULD be updatable - require.NoError(t, UpdateFriendshipByID(db, f.ID, model.StatusAccepted)) - err = UpdateFriendshipRemark(db, "user-x", "user-y", "cleared") - require.NoError(t, err) - - fetched, err := GetFriendshipByID(db, f.ID) - require.NoError(t, err) - assert.Equal(t, "cleared", fetched.Remark) - - // Clearing remark to empty string - err = UpdateFriendshipRemark(db, "user-x", "user-y", "") - require.NoError(t, err) - fetched, err = GetFriendshipByID(db, f.ID) - require.NoError(t, err) - assert.Equal(t, "", fetched.Remark) -} - -func TestFriendshipRepo_Lists(t *testing.T) { - db := setupSQLite(t) - - // Pending incoming - f1 := &model.Friendship{UserID: "alice", FriendID: "bob", Status: model.StatusPending} - f2 := &model.Friendship{UserID: "carol", FriendID: "bob", Status: model.StatusPending} - // Accepted - f3 := &model.Friendship{UserID: "bob", FriendID: "dave", Status: model.StatusAccepted} - - require.NoError(t, CreateFriendship(db, f1)) - require.NoError(t, CreateFriendship(db, f2)) - require.NoError(t, CreateFriendship(db, f3)) - - // Pending (received for bob) - received, err := ListReceivedRequests(db, "bob") - require.NoError(t, err) - assert.Len(t, received, 2) - - // Pending (sent by alice) - sent, err := ListSentRequests(db, "alice") - require.NoError(t, err) - assert.Len(t, sent, 1) - - // Accepted friends (bob's connections) - accepted, err := ListAcceptedFriends(db, "bob") - require.NoError(t, err) - assert.Len(t, accepted, 1) - assert.Equal(t, "dave", accepted[0].FriendID) - - // Friend IDs - ids, err := GetFriendIDs(db, "bob") - require.NoError(t, err) - assert.Len(t, ids, 1) - assert.Equal(t, "dave", ids[0]) -} - -func TestFriendshipRepo_BlockAndDelete(t *testing.T) { - db := setupSQLite(t) - - f := &model.Friendship{UserID: "u1", FriendID: "u2", Status: model.StatusAccepted} - require.NoError(t, CreateFriendship(db, f)) - - // Block - err := UpdateFriendshipByID(db, f.ID, model.StatusBlocked) - require.NoError(t, err) - - blocked, err := IsBlockedBy(db, "u1", "u2") - require.NoError(t, err) - assert.True(t, blocked) - - // Delete pair - err = DeleteFriendshipPair(db, "u1", "u2") - require.NoError(t, err) - - found, err := FindFriendshipBetween(db, "u1", "u2") - require.NoError(t, err) - assert.Nil(t, found) -} - -func TestFriendshipRepo_DeleteFriendshipDeletesLoadedRow(t *testing.T) { - db := setupSQLite(t) - - target := &model.Friendship{UserID: "delete-u1", FriendID: "delete-u2", Status: model.StatusPending} - other := &model.Friendship{UserID: "keep-u1", FriendID: "keep-u2", Status: model.StatusBlocked} - require.NoError(t, CreateFriendship(db, target)) - require.NoError(t, CreateFriendship(db, other)) - - require.NoError(t, DeleteFriendship(db, target)) - - _, err := GetFriendshipByID(db, target.ID) - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - kept, err := GetFriendshipByID(db, other.ID) - require.NoError(t, err) - assert.Equal(t, other.ID, kept.ID) -} - -func TestFriendshipRepo_Upsert(t *testing.T) { - db := setupSQLite(t) - - f := &model.Friendship{ - UserID: "upsert-a", - FriendID: "upsert-b", - Status: model.StatusPending, - } - require.NoError(t, UpsertFriendship(db, f)) - - // Upsert with same user_id+friend_id updates - f2 := &model.Friendship{ - UserID: "upsert-a", - FriendID: "upsert-b", - Status: model.StatusAccepted, - } - require.NoError(t, UpsertFriendship(db, f2)) - - found, err := FindFriendshipBetween(db, "upsert-a", "upsert-b") - require.NoError(t, err) - require.NotNil(t, found) - assert.Equal(t, model.StatusAccepted, found.Status) -} - -// ============================================================================= -// Notification repository tests -// ============================================================================= - -func TestNotificationRepo_CreateAndList(t *testing.T) { - db := setupSQLite(t) - - n1 := &model.Notification{UserID: "user-n1", Type: model.TypeMention, Payload: `{"key":"1"}`} - n2 := &model.Notification{UserID: "user-n1", Type: model.TypeSystem, Payload: `{"key":"2"}`} - n3 := &model.Notification{UserID: "user-n2", Type: model.TypeFriendRequest, Payload: `{"key":"3"}`} - require.NoError(t, CreateNotification(db, n1)) - require.NoError(t, CreateNotification(db, n2)) - require.NoError(t, CreateNotification(db, n3)) - - // List all for user-n1 - result, err := ListNotifications(db, "user-n1", false, 10, 0) - require.NoError(t, err) - assert.Len(t, result, 2) - - // Mark first as read - require.NoError(t, MarkNotificationRead(db, "user-n1", n1.ID)) - - // List unread only - result, err = ListNotifications(db, "user-n1", true, 10, 0) - require.NoError(t, err) - assert.Len(t, result, 1) - assert.Equal(t, n2.ID, result[0].ID) -} - -func TestNotificationRepo_MarkReadRequiresOwningUser(t *testing.T) { - db := setupSQLite(t) - - n := &model.Notification{UserID: "user-owner", Type: model.TypeMention, Payload: `{}`} - require.NoError(t, CreateNotification(db, n)) - - err := MarkNotificationRead(db, "user-other", n.ID) - require.ErrorIs(t, err, gorm.ErrRecordNotFound) - - unread, err := ListNotifications(db, "user-owner", true, 10, 0) - require.NoError(t, err) - require.Len(t, unread, 1) - assert.Equal(t, n.ID, unread[0].ID) -} - -func TestNotificationRepo_MarkAllRead(t *testing.T) { - db := setupSQLite(t) - - n1 := &model.Notification{UserID: "user-all", Type: model.TypeMention, Payload: `{}`} - n2 := &model.Notification{UserID: "user-all", Type: model.TypeSystem, Payload: `{}`} - require.NoError(t, CreateNotification(db, n1)) - require.NoError(t, CreateNotification(db, n2)) - - unread, err := ListNotifications(db, "user-all", true, 10, 0) - require.NoError(t, err) - assert.Len(t, unread, 2) - - require.NoError(t, MarkAllNotificationsRead(db, "user-all")) - - unread, err = ListNotifications(db, "user-all", true, 10, 0) - require.NoError(t, err) - assert.Len(t, unread, 0) -} - -// ============================================================================= -// Attachment repository tests -// ============================================================================= - -func TestAttachmentRepo_CreateAndGet(t *testing.T) { - db := setupSQLite(t) - - a := &model.Attachment{ - Hash: "abc123hash", - Size: 2048, - MimeType: "image/png", - OriginalName: "screenshot.png", - UploaderUserID: "user-att", - Metadata: `{"origin":"test"}`, - } - err := CreateAttachment(db, a) - require.NoError(t, err) - assert.NotEmpty(t, a.ID) - - // Get by ID - fetched, err := GetAttachmentByID(db, a.ID) - require.NoError(t, err) - assert.Equal(t, "abc123hash", fetched.Hash) - assert.Equal(t, int64(2048), fetched.Size) - assert.JSONEq(t, `{"origin":"test"}`, fetched.Metadata) - - // Get by hash - fetched, err = GetAttachmentByHash(db, "abc123hash") - require.NoError(t, err) - assert.Equal(t, a.ID, fetched.ID) - - // Non-existent - _, err = GetAttachmentByID(db, "nonexistent") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - - _, err = GetAttachmentByHash(db, "nonexistent") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) -} - -// ============================================================================= -// AgentInstance repository tests -// ============================================================================= - -func TestAgentInstanceRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - ai := &model.AgentInstance{ - AgentType: "code-explorer", - SessionID: "session-ai-1", - InviterUserID: "user-inviter", - DisplayName: "Code Explorer", - } - err := CreateAgentInstance(db, ai) - require.NoError(t, err) - assert.NotEmpty(t, ai.ID) - - // Get by ID - fetched, err := GetAgentInstanceByID(db, ai.ID) - require.NoError(t, err) - assert.Equal(t, "code-explorer", fetched.AgentType) - - // Create a second agent instance in the same session - ai2 := &model.AgentInstance{ - AgentType: "code-reviewer", - SessionID: "session-ai-1", - InviterUserID: "user-inviter", - DisplayName: "Code Reviewer", - } - require.NoError(t, CreateAgentInstance(db, ai2)) - - // List by session - list, err := ListAgentInstancesBySession(db, "session-ai-1") - require.NoError(t, err) - assert.Len(t, list, 2) - - // List by inviter - list, err = ListAgentInstancesByInviter(db, "session-ai-1", "user-inviter") - require.NoError(t, err) - assert.Len(t, list, 2) - - // List by different inviter - list, err = ListAgentInstancesByInviter(db, "session-ai-1", "other-user") - require.NoError(t, err) - assert.Len(t, list, 0) - - // Delete - require.NoError(t, DeleteAgentInstance(db, ai2.ID)) - list, err = ListAgentInstancesBySession(db, "session-ai-1") - require.NoError(t, err) - assert.Len(t, list, 1) -} - -// ============================================================================= -// CustomAgent repository tests -// ============================================================================= - -func TestCustomAgentRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - ca := &model.CustomAgent{ - OwnerUserID: "user-ca", - Name: "My Agent", - AgentType: "code-explorer", - SystemPrompt: "You are a helpful assistant.", - CapabilityTags: `["code"]`, - ToolWhitelist: `["read","write"]`, - ModelParams: `{}`, - } - err := CreateCustomAgent(db, ca) - require.NoError(t, err) - assert.NotEmpty(t, ca.ID) - - // Get by ID - fetched, err := GetCustomAgentByID(db, ca.ID) - require.NoError(t, err) - assert.Equal(t, "My Agent", fetched.Name) - - // List by owner - list, err := ListCustomAgentsByOwner(db, "user-ca") - require.NoError(t, err) - assert.Len(t, list, 1) - - // Create another - ca2 := &model.CustomAgent{ - OwnerUserID: "user-ca", - Name: "Agent 2", - AgentType: "code-reviewer", - SystemPrompt: "Review code.", - } - require.NoError(t, CreateCustomAgent(db, ca2)) - list, err = ListCustomAgentsByOwner(db, "user-ca") - require.NoError(t, err) - assert.Len(t, list, 2) - - // Update - ca.Name = "Renamed Agent" - err = UpdateCustomAgent(db, ca) - require.NoError(t, err) - fetched, err = GetCustomAgentByID(db, ca.ID) - require.NoError(t, err) - assert.Equal(t, "Renamed Agent", fetched.Name) - - // Soft delete - require.NoError(t, SoftDeleteCustomAgent(db, ca2.ID)) - _, err = GetCustomAgentByID(db, ca2.ID) - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - - // But the first is still there - fetched, err = GetCustomAgentByID(db, ca.ID) - require.NoError(t, err) - assert.NotNil(t, fetched) -} - -// ============================================================================= -// PendingAgentTask repository tests -// ============================================================================= - -func TestPendingTaskRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - expireAt := time.Now().Add(time.Hour) - task := &model.PendingAgentTask{ - AgentInstanceID: "agent-inst-1", - TriggeredByUserID: "user-trigger", - TriggerMessageID: "msg-trigger-1", - Status: model.TaskStatusQueued, - ExpireAt: expireAt, - } - err := CreatePendingTask(db, task) - require.NoError(t, err) - assert.NotEmpty(t, task.ID) - - // Get by ID - fetched, err := GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusQueued, fetched.Status) - - // Update status to dispatched - err = UpdatePendingTaskStatus(db, task.ID, model.TaskStatusDispatched, "") - require.NoError(t, err) - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusDispatched, fetched.Status) - assert.NotNil(t, fetched.DispatchedAt) - - err = UpdatePendingTaskDispatched(db, task.ID, "device-edge-1") - require.NoError(t, err) - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusDispatched, fetched.Status) - assert.Equal(t, "device-edge-1", fetched.EdgeDeviceID) - - // Update status to running and persist the Edge run mapping. - err = UpdatePendingTaskStatusWithEdgeRunID(db, task.ID, model.TaskStatusRunning, "", "run-edge-1") - require.NoError(t, err) - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusRunning, fetched.Status) - assert.Equal(t, "run-edge-1", fetched.EdgeRunID) - - // Update status to done - err = UpdatePendingTaskStatus(db, task.ID, model.TaskStatusDone, "") - require.NoError(t, err) - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusDone, fetched.Status) - assert.NotNil(t, fetched.FinishedAt) - - err = UpdatePendingTaskDispatched(db, task.ID, "device-after-done") - require.ErrorIs(t, err, gorm.ErrRecordNotFound) - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusDone, fetched.Status, "terminal task should not be moved back to dispatched") - assert.Equal(t, "device-edge-1", fetched.EdgeDeviceID) - - // Update with error - task2 := &model.PendingAgentTask{ - AgentInstanceID: "agent-inst-2", - TriggeredByUserID: "user-trigger", - TriggerMessageID: "msg-trigger-2", - Status: model.TaskStatusQueued, - ExpireAt: expireAt, - } - require.NoError(t, CreatePendingTask(db, task2)) - err = UpdatePendingTaskStatus(db, task2.ID, model.TaskStatusFailed, "something went wrong") - require.NoError(t, err) - fetched, err = GetPendingTaskByID(db, task2.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusFailed, fetched.Status) - assert.Equal(t, "something went wrong", fetched.ErrorMessage) -} - -func TestPendingTaskRepo_CancelTasksByAgent(t *testing.T) { - db := setupSQLite(t) - - expireAt := time.Now().Add(time.Hour) - task1 := &model.PendingAgentTask{ - AgentInstanceID: "agent-cancel", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-1", - Status: model.TaskStatusQueued, - ExpireAt: expireAt, - } - task2 := &model.PendingAgentTask{ - AgentInstanceID: "agent-cancel", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-2", - Status: model.TaskStatusDispatched, - ExpireAt: expireAt, - } - task3 := &model.PendingAgentTask{ - AgentInstanceID: "agent-other", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-3", - Status: model.TaskStatusQueued, - ExpireAt: expireAt, - } - require.NoError(t, CreatePendingTask(db, task1)) - require.NoError(t, CreatePendingTask(db, task2)) - require.NoError(t, CreatePendingTask(db, task3)) - - err := CancelTasksByAgentInstance(db, "agent-cancel") - require.NoError(t, err) - - // Tasks 1 and 2 should be cancelled - fetched, _ := GetPendingTaskByID(db, task1.ID) - require.NotNil(t, fetched) - assert.Equal(t, model.TaskStatusCancelled, fetched.Status) - - fetched, _ = GetPendingTaskByID(db, task2.ID) - require.NotNil(t, fetched) - assert.Equal(t, model.TaskStatusCancelled, fetched.Status) - - // Task 3 (different agent) should still be queued - fetched, _ = GetPendingTaskByID(db, task3.ID) - require.NotNil(t, fetched) - assert.Equal(t, model.TaskStatusQueued, fetched.Status) -} - -func TestPendingTaskRepo_ScanExpiredTasks(t *testing.T) { - db := setupSQLite(t) - - // Create tasks with different statuses and expire times - // Expired queued - task1 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-1", - Status: model.TaskStatusQueued, - ExpireAt: time.Now().Add(-time.Hour), // expired - } - // Expired dispatched - task2 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-2", - Status: model.TaskStatusDispatched, - ExpireAt: time.Now().Add(-time.Hour), // expired - } - // Expired running (#132: running tasks should also be scanned) - task3 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-3", - Status: model.TaskStatusRunning, - ExpireAt: time.Now().Add(-time.Hour), // expired - } - // Not expired yet - task4 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-4", - Status: model.TaskStatusQueued, - ExpireAt: time.Now().Add(time.Hour), // not expired - } - // Expired but already done (terminal states excluded) - task5 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-5", - Status: model.TaskStatusDone, - ExpireAt: time.Now().Add(-time.Hour), // expired but done - } - // Expired but already failed (terminal states excluded) - task6 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-6", - Status: model.TaskStatusFailed, - ExpireAt: time.Now().Add(-time.Hour), // expired but failed - } - // Expired but already cancelled (terminal states excluded) - task7 := &model.PendingAgentTask{ - AgentInstanceID: "agent-scan", - TriggeredByUserID: "user-scan", - TriggerMessageID: "msg-7", - Status: model.TaskStatusCancelled, - ExpireAt: time.Now().Add(-time.Hour), // expired but cancelled - } - - require.NoError(t, CreatePendingTask(db, task1)) - require.NoError(t, CreatePendingTask(db, task2)) - require.NoError(t, CreatePendingTask(db, task3)) - require.NoError(t, CreatePendingTask(db, task4)) - require.NoError(t, CreatePendingTask(db, task5)) - require.NoError(t, CreatePendingTask(db, task6)) - require.NoError(t, CreatePendingTask(db, task7)) - - tasks, err := ScanExpiredTasks(db) - require.NoError(t, err) - - // Should only return tasks 1, 2, 3 (expired and in non-terminal states) - assert.Len(t, tasks, 3) - - taskIDs := make(map[string]bool) - for _, t := range tasks { - taskIDs[t.ID] = true - } - assert.True(t, taskIDs[task1.ID], "expired queued task should be returned") - assert.True(t, taskIDs[task2.ID], "expired dispatched task should be returned") - assert.True(t, taskIDs[task3.ID], "expired running task should be returned (#132)") - assert.False(t, taskIDs[task4.ID], "non-expired task should not be returned") - assert.False(t, taskIDs[task5.ID], "expired done task should not be returned") - assert.False(t, taskIDs[task6.ID], "expired failed task should not be returned") - assert.False(t, taskIDs[task7.ID], "expired cancelled task should not be returned") -} - -// ============================================================================= -// SessionMember repository tests -// ============================================================================= - -func TestSessionMemberRepo_CRUD(t *testing.T) { - db := setupSQLite(t) - - s := &model.Session{Type: model.SessionTypeGroup, Name: "MemberTest"} - require.NoError(t, CreateSession(db, s)) - - // Create single member - m := &model.SessionMember{ - SessionID: s.ID, - MemberType: model.MemberTypeUser, - MemberID: "member-1", - Role: model.MemberRoleMember, - } - require.NoError(t, CreateSessionMember(db, m)) - assert.NotEmpty(t, m.ID) - - // Get member - fetched, err := GetMember(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - assert.Equal(t, model.MemberRoleMember, fetched.Role) - - // Get active member - active, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - assert.NotNil(t, active) - - // Is member active - isActive, err := IsMemberActive(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - assert.True(t, isActive) - - // Non-existent member - _, err = IsMemberSoftDeleted(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - // member-1 is active, so IsMemberSoftDeleted returns false - deleted, err := IsMemberSoftDeleted(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - assert.False(t, deleted) - - // Batch create - m2 := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeAgent, MemberID: "agent-1", Role: model.MemberRoleMember} - m3 := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeAgent, MemberID: "agent-2", Role: model.MemberRoleMember} - require.NoError(t, BatchCreateMembers(db, []*model.SessionMember{m2, m3})) - - all, err := ListActiveMembers(db, s.ID) - require.NoError(t, err) - assert.Len(t, all, 3) -} - -func TestSessionMemberRepo_SettingsAndTransfer(t *testing.T) { - db := setupSQLite(t) - - s := &model.Session{Type: model.SessionTypeGroup, Name: "SettingsTest"} - require.NoError(t, CreateSession(db, s)) - - owner := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "owner-1", Role: model.MemberRoleOwner} - member := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "member-1", Role: model.MemberRoleMember} - require.NoError(t, CreateSessionMember(db, owner)) - require.NoError(t, CreateSessionMember(db, member)) - - // Update member settings - pinned := true - muted := true - require.NoError(t, UpdateMemberSettings(db, s.ID, model.MemberTypeUser, "member-1", &pinned, nil, &muted)) - - fetched, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - assert.True(t, fetched.Pinned) - assert.True(t, fetched.Muted) - - // Transfer ownership - require.NoError(t, TransferOwnership(db, s.ID, "owner-1", "member-1")) - - // Old owner becomes member - oldOwner, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "owner-1") - require.NoError(t, err) - assert.Equal(t, model.MemberRoleMember, oldOwner.Role) - - // New owner - newOwner, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "member-1") - require.NoError(t, err) - assert.Equal(t, model.MemberRoleOwner, newOwner.Role) - - // Session owner_user_id updated - session, err := GetSessionByID(db, s.ID) - require.NoError(t, err) - assert.Equal(t, "member-1", *session.OwnerUserID) -} - -func TestSessionMemberRepo_LeaveAndReactivate(t *testing.T) { - db := setupSQLite(t) - - s := &model.Session{Type: model.SessionTypePrivate, OwnerUserID: strPtr("user-x")} - require.NoError(t, CreateSession(db, s)) - - m := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "user-x", Role: model.MemberRoleOwner} - require.NoError(t, CreateSessionMember(db, m)) - - // Soft delete (leave) - require.NoError(t, SoftDeleteMember(db, s.ID, model.MemberTypeUser, "user-x")) - - // No longer active - _, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "user-x") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - - // Is soft deleted - deleted, err := IsMemberSoftDeleted(db, s.ID, model.MemberTypeUser, "user-x") - require.NoError(t, err) - assert.True(t, deleted) - - // Reactivate - require.NoError(t, ReactivateMember(db, s.ID, model.MemberTypeUser, "user-x", model.MemberRoleMember)) - - // Active again - active, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "user-x") - require.NoError(t, err) - assert.Equal(t, model.MemberRoleMember, active.Role) - - // GetOtherMemberInPrivate - member2 := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "user-y", Role: model.MemberRoleMember} - require.NoError(t, CreateSessionMember(db, member2)) - - other, err := GetOtherMemberInPrivate(db, s.ID, "user-x") - require.NoError(t, err) - require.NotNil(t, other) - assert.Equal(t, "user-y", other.MemberID) - - // No other member - _, err = GetOtherMemberInPrivate(db, s.ID, "user-z") - require.NoError(t, err) - // user-z does not exist as a member, but the query still runs; the result is nil - // Actually: the query excludes user-z, so it returns one of the active members - // Let's verify: both user-x and user-y are active, excluding user-z returns both. - // But First() only returns one. - // Actually wait - user-z is not "user-x", so it returns... hmm. - // The query is: WHERE member_id != `user-z` AND left_at IS NULL - // Both user-x and user-y match that. First() picks one. - // So we should get a result, it's just that user-z is not the excluded one. - // This is fine, GetOtherMemberInPrivate returns the "other" member. -} - -func TestSessionMemberRepo_LastReadSeq(t *testing.T) { - db := setupSQLite(t) - - s := &model.Session{Type: model.SessionTypeGroup, Name: "ReadSeqTest"} - require.NoError(t, CreateSession(db, s)) - - m := &model.SessionMember{ - SessionID: s.ID, - MemberType: model.MemberTypeUser, - MemberID: "reader-1", - Role: model.MemberRoleMember, - } - require.NoError(t, CreateSessionMember(db, m)) - - err := UpdateLastReadSeq(db, s.ID, "reader-1", 5) - require.NoError(t, err) - - fetched, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "reader-1") - require.NoError(t, err) - assert.Equal(t, int64(5), fetched.LastReadSeq) - - // Update with a smaller seq — should not overwrite (WHERE last_read_seq < ?) - err = UpdateLastReadSeq(db, s.ID, "reader-1", 3) - require.NoError(t, err) - fetched, err = GetActiveMember(db, s.ID, model.MemberTypeUser, "reader-1") - require.NoError(t, err) - assert.Equal(t, int64(5), fetched.LastReadSeq) - - // Update with a larger seq - err = UpdateLastReadSeq(db, s.ID, "reader-1", 10) - require.NoError(t, err) - fetched, err = GetActiveMember(db, s.ID, model.MemberTypeUser, "reader-1") - require.NoError(t, err) - assert.Equal(t, int64(10), fetched.LastReadSeq) -} - -// ============================================================================= -// RefreshToken repository tests -// ============================================================================= - -func TestRefreshTokenRepo_UpsertAndGet(t *testing.T) { - db := setupSQLite(t) - - expiresAt := time.Now().Add(24 * time.Hour) - rt := &model.RefreshToken{ - UserID: "user-rt", - DeviceType: "desktop", - DeviceID: "dev-rt-1", - TokenHash: "hash123abc", - ExpiresAt: expiresAt, - } - err := UpsertRefreshToken(db, rt) - require.NoError(t, err) - assert.NotEmpty(t, rt.ID) - - // Find by hash - found, err := FindRefreshTokenByHash(db, "hash123abc") - require.NoError(t, err) - assert.Equal(t, "user-rt", found.UserID) - - // Non-existent hash - _, err = FindRefreshTokenByHash(db, "nonexistent") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) - - // Upsert again (same user_id + device_type + device_id) - rt2 := &model.RefreshToken{ - UserID: "user-rt", - DeviceType: "desktop", - DeviceID: "dev-rt-1", - TokenHash: "hash456def", - ExpiresAt: expiresAt.Add(time.Hour), - } - err = UpsertRefreshToken(db, rt2) - require.NoError(t, err) - // Should have same ID (updated existing) - assert.Equal(t, rt.ID, rt2.ID) - - // Find by new hash - found, err = FindRefreshTokenByHash(db, "hash456def") - require.NoError(t, err) - assert.Equal(t, rt.ID, found.ID) - - // Old hash no longer exists (overwritten) - _, err = FindRefreshTokenByHash(db, "hash123abc") - assert.ErrorIs(t, err, gorm.ErrRecordNotFound) -} - -func TestRefreshTokenRepo_Revoke(t *testing.T) { - db := setupSQLite(t) - - expiresAt := time.Now().Add(24 * time.Hour) - rt1 := &model.RefreshToken{UserID: "user-rev", DeviceType: "desktop", DeviceID: "dev-1", TokenHash: "h1", ExpiresAt: expiresAt} - rt2 := &model.RefreshToken{UserID: "user-rev", DeviceType: "desktop", DeviceID: "dev-2", TokenHash: "h2", ExpiresAt: expiresAt} - require.NoError(t, UpsertRefreshToken(db, rt1)) - require.NoError(t, UpsertRefreshToken(db, rt2)) - - // Revoke by device - require.NoError(t, RevokeRefreshTokensByUserDevice(db, "user-rev", "dev-1")) - - found, err := FindRefreshTokenByHash(db, "h1") - require.NoError(t, err) - assert.True(t, found.Revoked) - - // rt2 still not revoked - found, err = FindRefreshTokenByHash(db, "h2") - require.NoError(t, err) - assert.False(t, found.Revoked) - - // Revoke all for user - rt3 := &model.RefreshToken{UserID: "user-all-rev", DeviceType: "mobile", DeviceID: "dev-3", TokenHash: "h3", ExpiresAt: expiresAt} - require.NoError(t, UpsertRefreshToken(db, rt3)) - - require.NoError(t, RevokeAllUserTokens(db, "user-all-rev")) - found, err = FindRefreshTokenByHash(db, "h3") - require.NoError(t, err) - assert.True(t, found.Revoked) -} - -// ============================================================================= -// B2 data-integrity tests — atomic status, password+token, pin limit -// ============================================================================= - -func TestPendingTaskRepo_AtomicStatusUpdate(t *testing.T) { - db := setupSQLite(t) - - expireAt := time.Now().Add(time.Hour) - task := &model.PendingAgentTask{ - AgentInstanceID: "agent-atomic", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-1", - Status: model.TaskStatusDispatched, - ExpireAt: expireAt, - } - require.NoError(t, CreatePendingTask(db, task)) - - // Atomic transition: dispatched → running (should succeed) - rows, err := UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusDispatched, model.TaskStatusRunning, "") - require.NoError(t, err) - assert.Equal(t, int64(1), rows) - - fetched, err := GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusRunning, fetched.Status) - - // Atomic transition with wrong old status (should fail — 0 rows) - rows, err = UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusDispatched, model.TaskStatusDone, "") - require.NoError(t, err) - assert.Equal(t, int64(0), rows, "should not transition from running when oldStatus is dispatched") - - // Verify status unchanged - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusRunning, fetched.Status) - - // Atomic transition: running → done (should succeed) - rows, err = UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusRunning, model.TaskStatusDone, "") - require.NoError(t, err) - assert.Equal(t, int64(1), rows) - - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusDone, fetched.Status) - assert.NotNil(t, fetched.FinishedAt) -} - -func TestPendingTaskRepo_AtomicWithEdgeRunID(t *testing.T) { - db := setupSQLite(t) - - expireAt := time.Now().Add(time.Hour) - task := &model.PendingAgentTask{ - AgentInstanceID: "agent-edge", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-1", - Status: model.TaskStatusDispatched, - ExpireAt: expireAt, - } - require.NoError(t, CreatePendingTask(db, task)) - - // Atomic: dispatched → running with edgeRunID - rows, err := UpdatePendingTaskStatusAtomicWithEdgeRunID(db, task.ID, model.TaskStatusDispatched, model.TaskStatusRunning, "", "run-001") - require.NoError(t, err) - assert.Equal(t, int64(1), rows) - - fetched, err := GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusRunning, fetched.Status) - assert.Equal(t, "run-001", fetched.EdgeRunID) - - // Edge run id backfill on an already populated edgeRunID is a no-op. - rows, err = UpdatePendingTaskEdgeRunID(db, task.ID, "run-002") - // edge_run_id is already "run-001" so WHERE edge_run_id = '' matches 0 rows - require.NoError(t, err) - assert.Equal(t, int64(0), rows) - fetched, err = GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, "run-001", fetched.EdgeRunID, "edgeRunID should not be overwritten") - - backfillTask := &model.PendingAgentTask{ - AgentInstanceID: "agent-edge", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-2", - Status: model.TaskStatusRunning, - ExpireAt: expireAt, - } - require.NoError(t, CreatePendingTask(db, backfillTask)) - - rows, err = UpdatePendingTaskEdgeRunID(db, backfillTask.ID, "run-002") - require.NoError(t, err) - assert.Equal(t, int64(1), rows) - fetched, err = GetPendingTaskByID(db, backfillTask.ID) - require.NoError(t, err) - assert.Equal(t, "run-002", fetched.EdgeRunID) -} - -func TestPendingTaskRepo_AtomicFailClosed(t *testing.T) { - db := setupSQLite(t) - - expireAt := time.Now().Add(time.Hour) - task := &model.PendingAgentTask{ - AgentInstanceID: "agent-failclosed", - TriggeredByUserID: "user-t", - TriggerMessageID: "msg-1", - Status: model.TaskStatusRunning, - ExpireAt: expireAt, - } - require.NoError(t, CreatePendingTask(db, task)) - - // Simulate a race: first writer marks it done - rows, err := UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusRunning, model.TaskStatusDone, "") - require.NoError(t, err) - assert.Equal(t, int64(1), rows) - - // Second writer tries to mark it failed — 0 rows (fail closed) - rows, err = UpdatePendingTaskStatusAtomic(db, task.ID, model.TaskStatusRunning, model.TaskStatusFailed, "boom") - require.NoError(t, err) - assert.Equal(t, int64(0), rows, "second writer should get 0 rows affected") - - // Status should remain done - fetched, err := GetPendingTaskByID(db, task.ID) - require.NoError(t, err) - assert.Equal(t, model.TaskStatusDone, fetched.Status) -} - -func TestMessageRepo_PinMessageAtomic(t *testing.T) { - db := setupSQLite(t) - - s := &model.Session{Type: model.SessionTypeGroup, Name: "pin-test"} - require.NoError(t, CreateSession(db, s)) - - pin := func(msgID, userID string) error { - return PinMessageAtomic(db, &model.MessagePin{ - SessionID: s.ID, - MessageID: msgID, - PinnedByUserID: userID, - }, 3) // low limit for testing - } - - // Pin 3 messages — all should succeed - require.NoError(t, pin("msg-1", "user-1")) - require.NoError(t, pin("msg-2", "user-1")) - require.NoError(t, pin("msg-3", "user-1")) - - // 4th pin should fail with limit exceeded - err := pin("msg-4", "user-1") - assert.ErrorIs(t, err, ErrPinLimitExceeded) - - // Verify count - count, err := CountPinsBySession(db, s.ID) - require.NoError(t, err) - assert.Equal(t, int64(3), count) -} - -// ============================================================================= -// Helpers -// ============================================================================= - -func strPtr(s string) *string { - return &s -} diff --git a/hub-server/internal/repository/session_member_test.go b/hub-server/internal/repository/session_member_test.go new file mode 100644 index 000000000..d453e07fe --- /dev/null +++ b/hub-server/internal/repository/session_member_test.go @@ -0,0 +1,191 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// ============================================================================= +// SessionMember repository tests +// ============================================================================= + +func TestSessionMemberRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + s := &model.Session{Type: model.SessionTypeGroup, Name: "MemberTest"} + require.NoError(t, CreateSession(db, s)) + + // Create single member + m := &model.SessionMember{ + SessionID: s.ID, + MemberType: model.MemberTypeUser, + MemberID: "member-1", + Role: model.MemberRoleMember, + } + require.NoError(t, CreateSessionMember(db, m)) + assert.NotEmpty(t, m.ID) + + // Get member + fetched, err := GetMember(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + assert.Equal(t, model.MemberRoleMember, fetched.Role) + + // Get active member + active, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + assert.NotNil(t, active) + + // Is member active + isActive, err := IsMemberActive(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + assert.True(t, isActive) + + // Non-existent member + _, err = IsMemberSoftDeleted(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + // member-1 is active, so IsMemberSoftDeleted returns false + deleted, err := IsMemberSoftDeleted(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + assert.False(t, deleted) + + // Batch create + m2 := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeAgent, MemberID: "agent-1", Role: model.MemberRoleMember} + m3 := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeAgent, MemberID: "agent-2", Role: model.MemberRoleMember} + require.NoError(t, BatchCreateMembers(db, []*model.SessionMember{m2, m3})) + + all, err := ListActiveMembers(db, s.ID) + require.NoError(t, err) + assert.Len(t, all, 3) +} + +func TestSessionMemberRepo_SettingsAndTransfer(t *testing.T) { + db := setupSQLite(t) + + s := &model.Session{Type: model.SessionTypeGroup, Name: "SettingsTest"} + require.NoError(t, CreateSession(db, s)) + + owner := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "owner-1", Role: model.MemberRoleOwner} + member := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "member-1", Role: model.MemberRoleMember} + require.NoError(t, CreateSessionMember(db, owner)) + require.NoError(t, CreateSessionMember(db, member)) + + // Update member settings + pinned := true + muted := true + require.NoError(t, UpdateMemberSettings(db, s.ID, model.MemberTypeUser, "member-1", &pinned, nil, &muted)) + + fetched, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + assert.True(t, fetched.Pinned) + assert.True(t, fetched.Muted) + + // Transfer ownership + require.NoError(t, TransferOwnership(db, s.ID, "owner-1", "member-1")) + + // Old owner becomes member + oldOwner, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "owner-1") + require.NoError(t, err) + assert.Equal(t, model.MemberRoleMember, oldOwner.Role) + + // New owner + newOwner, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "member-1") + require.NoError(t, err) + assert.Equal(t, model.MemberRoleOwner, newOwner.Role) + + // Session owner_user_id updated + session, err := GetSessionByID(db, s.ID) + require.NoError(t, err) + assert.Equal(t, "member-1", *session.OwnerUserID) +} + +func TestSessionMemberRepo_LeaveAndReactivate(t *testing.T) { + db := setupSQLite(t) + + s := &model.Session{Type: model.SessionTypePrivate, OwnerUserID: strPtr("user-x")} + require.NoError(t, CreateSession(db, s)) + + m := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "user-x", Role: model.MemberRoleOwner} + require.NoError(t, CreateSessionMember(db, m)) + + // Soft delete (leave) + require.NoError(t, SoftDeleteMember(db, s.ID, model.MemberTypeUser, "user-x")) + + // No longer active + _, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "user-x") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + + // Is soft deleted + deleted, err := IsMemberSoftDeleted(db, s.ID, model.MemberTypeUser, "user-x") + require.NoError(t, err) + assert.True(t, deleted) + + // Reactivate + require.NoError(t, ReactivateMember(db, s.ID, model.MemberTypeUser, "user-x", model.MemberRoleMember)) + + // Active again + active, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "user-x") + require.NoError(t, err) + assert.Equal(t, model.MemberRoleMember, active.Role) + + // GetOtherMemberInPrivate + member2 := &model.SessionMember{SessionID: s.ID, MemberType: model.MemberTypeUser, MemberID: "user-y", Role: model.MemberRoleMember} + require.NoError(t, CreateSessionMember(db, member2)) + + other, err := GetOtherMemberInPrivate(db, s.ID, "user-x") + require.NoError(t, err) + require.NotNil(t, other) + assert.Equal(t, "user-y", other.MemberID) + + // No other member + _, err = GetOtherMemberInPrivate(db, s.ID, "user-z") + require.NoError(t, err) + // user-z does not exist as a member, but the query still runs; the result is nil + // Actually: the query excludes user-z, so it returns one of the active members + // Let's verify: both user-x and user-y are active, excluding user-z returns both. + // But First() only returns one. + // Actually wait - user-z is not "user-x", so it returns... hmm. + // The query is: WHERE member_id != `user-z` AND left_at IS NULL + // Both user-x and user-y match that. First() picks one. + // So we should get a result, it's just that user-z is not the excluded one. + // This is fine, GetOtherMemberInPrivate returns the "other" member. +} + +func TestSessionMemberRepo_LastReadSeq(t *testing.T) { + db := setupSQLite(t) + + s := &model.Session{Type: model.SessionTypeGroup, Name: "ReadSeqTest"} + require.NoError(t, CreateSession(db, s)) + + m := &model.SessionMember{ + SessionID: s.ID, + MemberType: model.MemberTypeUser, + MemberID: "reader-1", + Role: model.MemberRoleMember, + } + require.NoError(t, CreateSessionMember(db, m)) + + err := UpdateLastReadSeq(db, s.ID, "reader-1", 5) + require.NoError(t, err) + + fetched, err := GetActiveMember(db, s.ID, model.MemberTypeUser, "reader-1") + require.NoError(t, err) + assert.Equal(t, int64(5), fetched.LastReadSeq) + + // Update with a smaller seq — should not overwrite (WHERE last_read_seq < ?) + err = UpdateLastReadSeq(db, s.ID, "reader-1", 3) + require.NoError(t, err) + fetched, err = GetActiveMember(db, s.ID, model.MemberTypeUser, "reader-1") + require.NoError(t, err) + assert.Equal(t, int64(5), fetched.LastReadSeq) + + // Update with a larger seq + err = UpdateLastReadSeq(db, s.ID, "reader-1", 10) + require.NoError(t, err) + fetched, err = GetActiveMember(db, s.ID, model.MemberTypeUser, "reader-1") + require.NoError(t, err) + assert.Equal(t, int64(10), fetched.LastReadSeq) +} diff --git a/hub-server/internal/repository/session_test.go b/hub-server/internal/repository/session_test.go new file mode 100644 index 000000000..1b95d6384 --- /dev/null +++ b/hub-server/internal/repository/session_test.go @@ -0,0 +1,121 @@ +package repository + +import ( + "testing" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// ============================================================================= +// Session repository tests +// ============================================================================= + +func TestSessionRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + session := &model.Session{ + Type: model.SessionTypePrivate, + Name: "Test Session", + OwnerUserID: strPtr("user-001"), + } + + // Create + err := CreateSession(db, session) + require.NoError(t, err) + assert.NotEmpty(t, session.ID) + + // Read + fetched, err := GetSessionByID(db, session.ID) + require.NoError(t, err) + assert.Equal(t, model.SessionTypePrivate, fetched.Type) + assert.Equal(t, "Test Session", fetched.Name) + + // Update + fetched.Name = "Updated Session" + err = UpdateSessionColumns(db, fetched, "name") + require.NoError(t, err) + + fetched2, err := GetSessionByID(db, session.ID) + require.NoError(t, err) + assert.Equal(t, "Updated Session", fetched2.Name) +} + +func TestSessionRepo_FindPrivateSessionBetween(t *testing.T) { + db := setupSQLite(t) + + session := &model.Session{ + Type: model.SessionTypePrivate, + } + require.NoError(t, CreateSession(db, session)) + + // Add both members + m1 := &model.SessionMember{ + SessionID: session.ID, + MemberType: model.MemberTypeUser, + MemberID: "user-a", + Role: model.MemberRoleMember, + } + m2 := &model.SessionMember{ + SessionID: session.ID, + MemberType: model.MemberTypeUser, + MemberID: "user-b", + Role: model.MemberRoleMember, + } + require.NoError(t, CreateSessionMember(db, m1)) + require.NoError(t, CreateSessionMember(db, m2)) + + // Find between + found, err := FindPrivateSessionBetween(db, "user-a", "user-b") + require.NoError(t, err) + require.NotNil(t, found) + assert.Equal(t, session.ID, found.ID) + + // Not found + found, err = FindPrivateSessionBetween(db, "user-a", "user-c") + require.NoError(t, err) + assert.Nil(t, found) +} + +func TestSessionRepo_TouchLastMessage(t *testing.T) { + db := setupSQLite(t) + + session := &model.Session{Type: model.SessionTypeGroup, Name: "Group"} + require.NoError(t, CreateSession(db, session)) + + // Initially nil + fetched, err := GetSessionByID(db, session.ID) + require.NoError(t, err) + assert.Nil(t, fetched.LastMessageAt) + + err = TouchSessionLastMessage(db, session.ID) + require.NoError(t, err) + + fetched, err = GetSessionByID(db, session.ID) + require.NoError(t, err) + assert.NotNil(t, fetched.LastMessageAt) +} + +func TestSessionRepo_ListUserSessions(t *testing.T) { + // member_count uses a correlated scalar subquery compatible with both PG and SQLite. + // Plan-level assertions are in session_member_count_test.go (PG-only EXPLAIN). + db := setupSQLite(t) + + s := &model.Session{Type: model.SessionTypeGroup, Name: "ListGroup"} + require.NoError(t, CreateSession(db, s)) + + m := &model.SessionMember{ + SessionID: s.ID, + MemberType: model.MemberTypeUser, + MemberID: "user-l", + Role: model.MemberRoleOwner, + } + require.NoError(t, CreateSessionMember(db, m)) + + result, err := ListUserSessions(db, "user-l") + require.NoError(t, err) + assert.Len(t, result, 1) + assert.Equal(t, s.ID, result[0].ID) + assert.Equal(t, model.MemberRoleOwner, result[0].Role) +} diff --git a/hub-server/internal/repository/user_test.go b/hub-server/internal/repository/user_test.go new file mode 100644 index 000000000..118382a38 --- /dev/null +++ b/hub-server/internal/repository/user_test.go @@ -0,0 +1,132 @@ +package repository + +import ( + "strings" + "testing" + "time" + + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +// ============================================================================= +// User repository tests +// ============================================================================= + +func TestUserRepo_CRUD(t *testing.T) { + db := setupSQLite(t) + + user := &model.User{ + Username: "testuser", + PasswordHash: strPtr("hashed_password"), + Nickname: "Test User", + } + + // Create + err := CreateUser(db, user) + require.NoError(t, err) + assert.NotEmpty(t, user.ID) + + // Read by ID + fetched, err := GetUserByID(db, user.ID) + require.NoError(t, err) + assert.Equal(t, user.Username, fetched.Username) + assert.Equal(t, user.Nickname, fetched.Nickname) + + // Read by username + fetchedByUsername, err := GetUserByUsername(db, "testuser") + require.NoError(t, err) + assert.Equal(t, user.ID, fetchedByUsername.ID) + + // Read non-existent + _, err = GetUserByID(db, "nonexistent-id") + assert.ErrorIs(t, err, gorm.ErrRecordNotFound) + + // Update + user.Nickname = "Updated Name" + err = UpdateUser(db, user) + require.NoError(t, err) + fetched, err = GetUserByID(db, user.ID) + require.NoError(t, err) + assert.Equal(t, "Updated Name", fetched.Nickname) +} + +func TestUserRepo_GetUsersByIDs(t *testing.T) { + db := setupSQLite(t) + + // Create multiple users + u1 := &model.User{Username: "user1", PasswordHash: strPtr("h1"), Nickname: "U1"} + u2 := &model.User{Username: "user2", PasswordHash: strPtr("h2"), Nickname: "U2"} + u3 := &model.User{Username: "user3", PasswordHash: strPtr("h3"), Nickname: "U3"} + require.NoError(t, CreateUser(db, u1)) + require.NoError(t, CreateUser(db, u2)) + require.NoError(t, CreateUser(db, u3)) + + // Fetch by IDs + m, err := GetUsersByIDs(db, []string{u1.ID, u2.ID}) + require.NoError(t, err) + assert.Len(t, m, 2) + assert.Equal(t, "U1", m[u1.ID].Nickname) + assert.Equal(t, "U2", m[u2.ID].Nickname) + + // Empty list + m, err = GetUsersByIDs(db, []string{}) + require.NoError(t, err) + assert.Empty(t, m) + + // Non-existent IDs + m, err = GetUsersByIDs(db, []string{"no-such-id"}) + require.NoError(t, err) + assert.Empty(t, m) +} + +func TestUserRepo_FindOrCreateByTokenDanceSubCreatesStableDistinctHubUsers(t *testing.T) { + db := setupSQLite(t) + + sub1 := "tokendance-subject-with-a-very-long-shared-prefix-" + strings.Repeat("1", 80) + sub2 := "tokendance-subject-with-a-very-long-shared-prefix-" + strings.Repeat("2", 80) + + u1, err := FindOrCreateByTokenDanceSub(db, sub1, "", "") + require.NoError(t, err) + u2, err := FindOrCreateByTokenDanceSub(db, sub2, "", "") + require.NoError(t, err) + + require.NotEqual(t, u1.ID, u2.ID) + require.NotEqual(t, u1.Username, u2.Username) + assert.True(t, strings.HasPrefix(u1.Username, "td_")) + assert.True(t, len(u1.Username) <= 32) + assert.True(t, len(u1.Nickname) <= 64) + require.NotNil(t, u1.TokenDanceSub) + assert.Equal(t, sub1, *u1.TokenDanceSub) + + again, err := FindOrCreateByTokenDanceSub(db, sub1, "", "") + require.NoError(t, err) + assert.Equal(t, u1.ID, again.ID) + assert.Equal(t, u1.Username, again.Username) +} + +func TestUserRepo_ReadsOIDCOnlyUserWithNullPasswordHash(t *testing.T) { + db := setupSQLite(t) + + sub := "tokendance-sub-null-password" + now := time.Now().UTC() + require.NoError(t, db.Exec(` + INSERT INTO users ( + id, username, password_hash, nickname, tokendance_sub, tokendance_sub_linked_at, created_at, updated_at + ) VALUES (?, ?, NULL, ?, ?, ?, ?, ?) + `, "user-oidc-null-password", "td_null_password", "OIDC Only", sub, now, now, now).Error) + + bySub, err := FindByTokenDanceSub(db, sub) + require.NoError(t, err) + require.NotNil(t, bySub) + assert.Equal(t, "user-oidc-null-password", bySub.ID) + assert.Nil(t, bySub.PasswordHash) + + byID, err := GetUserByID(db, "user-oidc-null-password") + require.NoError(t, err) + require.NotNil(t, byID) + assert.Equal(t, sub, *byID.TokenDanceSub) + assert.Nil(t, byID.PasswordHash) +} diff --git a/hub-server/internal/repository/user_upsert_predicate_test.go b/hub-server/internal/repository/user_upsert_predicate_test.go index 920decf0c..1477512ff 100644 --- a/hub-server/internal/repository/user_upsert_predicate_test.go +++ b/hub-server/internal/repository/user_upsert_predicate_test.go @@ -27,7 +27,7 @@ func TestTokenDanceSubPredicateSSOT(t *testing.T) { {"model.User gorm tag", "../model/user.go"}, // SQLite fixtures. If any of these drops the predicate it silently // re-masks the inference requirement above. - {"fixture repository_test.go", "repository_test.go"}, + {"fixture helpers_test.go", "helpers_test.go"}, {"fixture upsert_toctou_test.go", "upsert_toctou_test.go"}, {"fixture service/oidc/oidc_test.go", "../service/oidc/oidc_test.go"}, {"fixture tests/integration/tokendance_oidc_e2e_test.go", "../../tests/integration/tokendance_oidc_e2e_test.go"}, diff --git a/hub-server/internal/service/agentteam/agent_team_approval_test.go b/hub-server/internal/service/agentteam/agent_team_approval_test.go new file mode 100644 index 000000000..00f84e9c8 --- /dev/null +++ b/hub-server/internal/service/agentteam/agent_team_approval_test.go @@ -0,0 +1,531 @@ +package agentteam + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAgentTeamService_ResolveConflictAppendsEventAndUpdatesReplay(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + reviewerProfileID := "profile-reviewer" + reviewer := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &reviewerProfileID, + Role: model.TeamMemberRoleReviewer, + } + require.NoError(t, repository.AddTeamMember(db, reviewer)) + firstTaskID := "agent-task-one" + secondTaskID := "agent-task-two" + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusDone, + Objective: "Change shared file", + RunID: &firstTaskID, + })) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: reviewer.ID, + Status: model.TeamTaskStatusDone, + Objective: "Review shared file", + RunID: &secondTaskID, + })) + for _, event := range []model.AgentRunEvent{ + { + TaskID: firstTaskID, + EdgeRunID: "edge-one", + SessionID: run.SessionID, + AgentInstanceID: "agent-one", + EventType: "run.agent.file_change", + Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, + }, + { + TaskID: secondTaskID, + EdgeRunID: "edge-two", + SessionID: run.SessionID, + AgentInstanceID: "agent-two", + EventType: "run.agent.file_change", + Payload: `{"path":"./shared.txt","action":"modified","status":"completed"}`, + }, + } { + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &event)) + } + + conflictID := conflictIDForPath("shared.txt") + resolved, err := svc.ResolveConflict(context.Background(), "user-1", team.ID, run.ID, model.TeamConflictResolution{ + ConflictID: conflictID, + Resolution: model.TeamConflictResolutionAcceptAgentTask, + SelectedAgentTaskID: firstTaskID, + Reason: "Use executor result", + }) + require.NoError(t, err) + require.NotNil(t, resolved) + assert.Equal(t, model.TeamConflictStatusResolved, resolved.Status) + assert.Equal(t, model.TeamConflictResolutionAcceptAgentTask, resolved.Resolution) + assert.Equal(t, firstTaskID, resolved.SelectedTask) + assert.Equal(t, "user-1", resolved.ResolvedBy) + require.NotNil(t, resolved.ResolvedAt) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventConflictResolved, events[0].Type) + assert.Contains(t, events[0].Payload, firstTaskID) + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.Conflicts, 1) + assert.Equal(t, model.TeamConflictStatusResolved, state.Conflicts[0].Status) + assert.Equal(t, model.TeamConflictResolutionAcceptAgentTask, state.Conflicts[0].Resolution) + assert.Equal(t, firstTaskID, state.Conflicts[0].SelectedTask) +} + +func TestAgentTeamService_ResolveConflictRejectsTaskOutsideConflict(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + reviewerProfileID := "profile-reviewer" + reviewer := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &reviewerProfileID, + Role: model.TeamMemberRoleReviewer, + } + require.NoError(t, repository.AddTeamMember(db, reviewer)) + taskID := "agent-task-one" + otherTaskID := "agent-task-two" + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusDone, + Objective: "Change shared file", + RunID: &taskID, + })) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: reviewer.ID, + Status: model.TeamTaskStatusDone, + Objective: "Review shared file", + RunID: &otherTaskID, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: taskID, + EdgeRunID: "edge-one", + SessionID: run.SessionID, + AgentInstanceID: "agent-one", + EventType: "run.agent.file_change", + Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: otherTaskID, + EdgeRunID: "edge-two", + SessionID: run.SessionID, + AgentInstanceID: "agent-two", + EventType: "run.agent.file_change", + Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, + })) + + resolved, err := svc.ResolveConflict(context.Background(), "user-1", team.ID, run.ID, model.TeamConflictResolution{ + ConflictID: conflictIDForPath("shared.txt"), + Resolution: model.TeamConflictResolutionAcceptAgentTask, + SelectedAgentTaskID: "missing-task", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, resolved) +} + +func TestAgentTeamService_DecideApprovalAppendsEventAndUpdatesReplay(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + pending := &model.PendingAgentTask{ + ID: "agent-task-approval", + AgentInstanceID: "agent-executor", + TriggeredByUserID: "user-1", + TriggerMessageID: "msg-approval", + Status: model.TaskStatusRunning, + EdgeRunID: "edge-run-approval", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(pending).Error) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Run gated command", + RunID: &pending.ID, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: pending.EdgeRunID, + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-approval","toolUseId":"tool-approval","toolName":"Bash","status":"pending"}`, + })) + + decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-approval", model.TeamApprovalDecision{ + Decision: "allow", + Reason: "Known safe command", + }) + require.NoError(t, err) + require.NotNil(t, decided) + assert.Equal(t, "req-approval", decided.ApprovalID) + assert.Equal(t, "allow", decided.Status) + assert.Equal(t, "user-1", decided.DecidedBy) + require.NotNil(t, decided.EdgeControl) + assert.Equal(t, pending.EdgeRunID, decided.EdgeControl.RunID) + assert.Equal(t, "req-approval", decided.EdgeControl.RequestID) + assert.Equal(t, "allow", decided.EdgeControl.Decision) + assert.Equal(t, "Known safe command", decided.EdgeControl.Reason) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventApprovalDecided, events[0].Type) + assert.Contains(t, events[0].Payload, `"edge_control"`) + assert.Contains(t, events[0].Payload, `"runId":"edge-run-approval"`) + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.Approvals, 1) + assert.Equal(t, "req-approval", state.Approvals[0].ApprovalID) + assert.Equal(t, "allow", state.Approvals[0].Status) + assert.Equal(t, "Known safe command", state.Approvals[0].Reason) + assert.Equal(t, "user-1", state.Approvals[0].DecidedBy) + require.NotNil(t, state.Approvals[0].DecidedAt) + require.NotNil(t, state.Approvals[0].EdgeControl) + assert.Equal(t, "edge-run-approval", state.Approvals[0].EdgeControl.RunID) +} + +func TestAgentTeamService_DecideApprovalDeliversControlToExactEdgeDevice(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + controlSvc := &mockAgentTeamControlSvc{} + svc := NewAgentTeamService(db, nil, nil) + svc.SetControlService(controlSvc) + team, _, executor, run := seedAgentTeamRun(t, db) + pending := &model.PendingAgentTask{ + ID: "agent-task-control", + AgentInstanceID: "agent-executor", + TriggeredByUserID: "user-1", + TriggerMessageID: "msg-control", + TargetID: "target-local", + Status: model.TaskStatusRunning, + EdgeRunID: "edge-run-control", + EdgeDeviceID: "edge-device-control", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(pending).Error) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Run gated command", + RunID: &pending.ID, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: pending.EdgeRunID, + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-control","toolUseId":"tool-control","toolName":"Bash","status":"pending"}`, + })) + + _, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-control", model.TeamApprovalDecision{ + Decision: "allow", + Reason: "Known safe command", + }) + require.NoError(t, err) + require.Len(t, controlSvc.calls, 1) + call := controlSvc.calls[0] + assert.Equal(t, "user-1", call.userID) + assert.Equal(t, "edge-device-control", call.deviceID) + assert.Equal(t, model.AgentControlKindPermissionDecide, call.payload.Kind) + assert.Equal(t, pending.ID, call.payload.AgentTaskID) + assert.Equal(t, "target-local", call.payload.TargetID) + assert.Equal(t, "edge-device-control", call.payload.EdgeDeviceID) + assert.Equal(t, "req-control", call.payload.ApprovalID) + require.NotNil(t, call.payload.EdgeControl) + assert.Equal(t, "edge-run-control", call.payload.EdgeControl.RunID) + assert.Equal(t, "req-control", call.payload.EdgeControl.RequestID) + assert.Equal(t, "allow", call.payload.EdgeControl.Decision) +} + +func TestAgentTeamService_DecideApprovalRedeliversSameDecisionWithoutDuplicateEvent(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + controlSvc := &mockAgentTeamControlSvc{} + svc := NewAgentTeamService(db, nil, nil) + svc.SetControlService(controlSvc) + team, _, executor, run := seedAgentTeamRun(t, db) + pending := &model.PendingAgentTask{ + ID: "agent-task-redeliver-control", + AgentInstanceID: "agent-executor", + TriggeredByUserID: "user-1", + TriggerMessageID: "msg-redeliver-control", + TargetID: "target-local", + Status: model.TaskStatusRunning, + EdgeRunID: "edge-run-redeliver-control", + EdgeDeviceID: "edge-device-redeliver-control", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(pending).Error) + teamTask := &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Run gated command", + RunID: &pending.ID, + } + require.NoError(t, repository.CreateTeamTask(db, teamTask)) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: pending.EdgeRunID, + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-redeliver","toolUseId":"tool-redeliver","toolName":"Bash","status":"pending"}`, + })) + decidedAt := time.Now().UTC() + record := model.TeamApprovalDecision{ + ApprovalID: "req-redeliver", + AgentTaskID: pending.ID, + TeamTaskID: teamTask.ID, + MemberID: executor.ID, + EdgeRunID: pending.EdgeRunID, + RequestID: "req-redeliver", + ToolName: "Bash", + ToolUseID: "tool-redeliver", + Decision: "allow", + Reason: "Known safe command", + DecidedBy: "user-1", + DecidedAt: decidedAt, + EdgeControl: &model.TeamApprovalEdgeControl{ + RunID: pending.EdgeRunID, + RequestID: "req-redeliver", + Decision: "allow", + Reason: "Known safe command", + }, + } + payload, err := json.Marshal(record) + require.NoError(t, err) + require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ + TeamRunID: run.ID, + Type: model.TeamEventApprovalDecided, + Payload: string(payload), + })) + + decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-redeliver", model.TeamApprovalDecision{ + Decision: "allow", + Reason: "Known safe command", + }) + + require.NoError(t, err) + require.NotNil(t, decided) + assert.Equal(t, "allow", decided.Status) + require.Len(t, controlSvc.calls, 1) + call := controlSvc.calls[0] + assert.Equal(t, "user-1", call.userID) + assert.Equal(t, pending.EdgeDeviceID, call.deviceID) + assert.Equal(t, pending.ID, call.payload.AgentTaskID) + assert.Equal(t, "req-redeliver", call.payload.ApprovalID) + require.NotNil(t, call.payload.EdgeControl) + assert.Equal(t, pending.EdgeRunID, call.payload.EdgeControl.RunID) + assert.Equal(t, "allow", call.payload.EdgeControl.Decision) + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) +} + +func TestAgentTeamService_DecideApprovalRejectsMissingEdgeDeviceForControl(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + controlSvc := &mockAgentTeamControlSvc{} + svc := NewAgentTeamService(db, nil, nil) + svc.SetControlService(controlSvc) + team, _, executor, run := seedAgentTeamRun(t, db) + pending := &model.PendingAgentTask{ + ID: "agent-task-no-device", + AgentInstanceID: "agent-executor", + TriggeredByUserID: "user-1", + TriggerMessageID: "msg-no-device", + Status: model.TaskStatusRunning, + EdgeRunID: "edge-run-no-device", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(pending).Error) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Run gated command", + RunID: &pending.ID, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: pending.EdgeRunID, + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-no-device","toolUseId":"tool-no-device","toolName":"Bash","status":"pending"}`, + })) + + decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-no-device", model.TeamApprovalDecision{ + Decision: "allow", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, decided) + assert.Empty(t, controlSvc.calls) +} + +func TestAgentTeamService_DecideApprovalRejectsAlreadyDecided(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + taskID := "agent-task-decided" + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Run gated command", + RunID: &taskID, + })) + for _, event := range []model.AgentRunEvent{ + { + TaskID: taskID, + EdgeRunID: "edge-run-decided", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-decided","toolUseId":"tool-decided","toolName":"Bash","status":"pending"}`, + }, + { + TaskID: taskID, + EdgeRunID: "edge-run-decided", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_decided", + Payload: `{"requestId":"req-decided","decision":"deny","reason":"too broad"}`, + }, + } { + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &event)) + } + + decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-decided", model.TeamApprovalDecision{ + Decision: "allow", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, decided) +} + +func TestAgentTeamService_MemberReadableTeamCannotMutateRunDecisions(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + addReadableTeamMemberForUser(t, db, team.ID, "member-user") + + state, err := svc.GetTeamRunState(context.Background(), "member-user", team.ID, run.ID) + require.NoError(t, err) + require.NotNil(t, state) + + reviewerProfileID := "profile-reviewer" + reviewer := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &reviewerProfileID, + Role: model.TeamMemberRoleReviewer, + } + require.NoError(t, repository.AddTeamMember(db, reviewer)) + firstTaskID := "agent-task-conflict-one" + secondTaskID := "agent-task-conflict-two" + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusDone, + Objective: "Change shared file", + RunID: &firstTaskID, + })) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: reviewer.ID, + Status: model.TeamTaskStatusDone, + Objective: "Review shared file", + RunID: &secondTaskID, + })) + for _, event := range []model.AgentRunEvent{ + { + TaskID: firstTaskID, + EdgeRunID: "edge-conflict-one", + SessionID: run.SessionID, + AgentInstanceID: "agent-one", + EventType: "run.agent.file_change", + Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, + }, + { + TaskID: secondTaskID, + EdgeRunID: "edge-conflict-two", + SessionID: run.SessionID, + AgentInstanceID: "agent-two", + EventType: "run.agent.file_change", + Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, + }, + } { + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &event)) + } + + resolved, err := svc.ResolveConflict(context.Background(), "member-user", team.ID, run.ID, model.TeamConflictResolution{ + ConflictID: conflictIDForPath("shared.txt"), + Resolution: model.TeamConflictResolutionAcceptAgentTask, + SelectedAgentTaskID: firstTaskID, + }) + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) + assert.Nil(t, resolved) + + pending := &model.PendingAgentTask{ + ID: "agent-task-member-approval", + AgentInstanceID: "agent-executor", + TriggeredByUserID: "user-1", + TriggerMessageID: "msg-member-approval", + Status: model.TaskStatusRunning, + EdgeRunID: "edge-run-member-approval", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(pending).Error) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Run gated command", + RunID: &pending.ID, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: pending.EdgeRunID, + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-member-approval","toolUseId":"tool-member-approval","toolName":"Bash","status":"pending"}`, + })) + + decided, err := svc.DecideApproval(context.Background(), "member-user", team.ID, run.ID, "req-member-approval", model.TeamApprovalDecision{ + Decision: "allow", + }) + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) + assert.Nil(t, decided) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + assert.Empty(t, events) +} diff --git a/hub-server/internal/service/agentteam/agent_team_compete_test.go b/hub-server/internal/service/agentteam/agent_team_compete_test.go new file mode 100644 index 000000000..c1f8568d4 --- /dev/null +++ b/hub-server/internal/service/agentteam/agent_team_compete_test.go @@ -0,0 +1,220 @@ +package agentteam + +import ( + "context" + "fmt" + "strings" + "testing" + + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// mockCompeteAggregator implements CompeteAggregator for tests. +type mockCompeteAggregator struct { + summary string +} + +func (m *mockCompeteAggregator) CompareResults(_ context.Context, _ string, entries []model.CompeteSummaryEntry) (string, error) { + if m.summary != "" { + return m.summary, nil + } + return "Comparison: " + strings.Join(entryMemberIDs(entries), " vs "), nil +} + +func TestCompeteMode(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + aggregator := &mockCompeteAggregator{} + svc := NewAgentTeamService(db, nil, nil) + svc.SetCompeteAggregator(aggregator) + + team := &model.AgentTeam{OwnerID: "user-1", Name: "Compete Team"} + require.NoError(t, repository.CreateTeam(db, team)) + + supervisor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-supervisor"), + Role: model.TeamMemberRoleSupervisor, + } + executor1 := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-executor-1"), + Role: model.TeamMemberRoleExecutor, + } + executor2 := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-executor-2"), + Role: model.TeamMemberRoleExecutor, + } + require.NoError(t, repository.AddTeamMember(db, supervisor)) + require.NoError(t, repository.AddTeamMember(db, executor1)) + require.NoError(t, repository.AddTeamMember(db, executor2)) + + run := &model.AgentTeamRun{ + TeamID: team.ID, + SessionID: "session-compete", + TriggerUserID: "user-1", + TriggerMessage: "compare implementations of factorial", + Mode: model.TeamRunModeCompete, + Status: model.TeamRunStatusRunning, + } + require.NoError(t, repository.CreateTeamRun(db, run)) + + // Submit a compete route decision targeting two executors. + workerIDs := executor1.ID + "," + executor2.ID + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "compete", + NextWorker: workerIDs, + Instructions: "Implement factorial in your best style", + Reasoning: "Let's see different approaches", + }) + require.NoError(t, err) + require.NotNil(t, assignment) + assert.Equal(t, model.AssignmentTypeCompete, assignment.Type) + + // Both executors should have assignments. + assignments, err := svc.ListAssignments(context.Background(), "user-1", run.ID) + require.NoError(t, err) + assert.Len(t, assignments, 2) + assert.Equal(t, model.AssignmentTypeCompete, assignments[0].Type) + assert.Equal(t, model.AssignmentTypeCompete, assignments[1].Type) + + // Complete both assignments with results. + as1 := &assignments[0] + as1.Status = model.AssignmentStatusRunning + require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", as1.ID).Update("status", model.AssignmentStatusRunning).Error) + require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", as1.ID, "func fact(n int) int { if n <= 1 { return 1 }; return n * fact(n-1) }")) + + as2 := &assignments[1] + as2.Status = model.AssignmentStatusRunning + require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", as2.ID).Update("status", model.AssignmentStatusRunning).Error) + require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", as2.ID, "def factorial(n): return 1 if n <= 1 else n * factorial(n-1)")) + + // Generate compete summary. + resp, err := svc.GenerateCompeteSummary(context.Background(), "user-1", run.ID, model.CompeteSummaryRequest{}) + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, run.ID, resp.TeamRunID) + assert.Contains(t, resp.Summary, "Comparison:") + assert.Len(t, resp.Entries, 2) + + // Verify events include compete dispatched and aggregated. + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + hasCompeteDispatched := false + hasCompeteAggregated := false + for _, ev := range events { + if ev.Type == model.TeamEventCompeteDispatched { + hasCompeteDispatched = true + } + if ev.Type == model.TeamEventCompeteAggregated { + hasCompeteAggregated = true + } + } + assert.True(t, hasCompeteDispatched, "expected team.compete.dispatched event") + assert.True(t, hasCompeteAggregated, "expected team.compete.aggregated event") +} + +func TestCompeteModeAutoSelectsExecutors(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + aggregator := &mockCompeteAggregator{} + svc := NewAgentTeamService(db, nil, nil) + svc.SetCompeteAggregator(aggregator) + + team := &model.AgentTeam{OwnerID: "user-1", Name: "Auto Compete Team"} + require.NoError(t, repository.CreateTeam(db, team)) + + supervisor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-supervisor"), + Role: model.TeamMemberRoleSupervisor, + } + executor1 := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-executor-1"), + Role: model.TeamMemberRoleExecutor, + } + executor2 := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-executor-2"), + Role: model.TeamMemberRoleExecutor, + } + require.NoError(t, repository.AddTeamMember(db, supervisor)) + require.NoError(t, repository.AddTeamMember(db, executor1)) + require.NoError(t, repository.AddTeamMember(db, executor2)) + + run := &model.AgentTeamRun{ + TeamID: team.ID, + SessionID: "session-auto-compete", + TriggerUserID: "user-1", + TriggerMessage: "auto compete", + Mode: model.TeamRunModeCompete, + Status: model.TeamRunStatusRunning, + } + require.NoError(t, repository.CreateTeamRun(db, run)) + + // Submit compete decision with empty NextWorker — should pick both executors. + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "compete", + Instructions: "Write hello world", + Reasoning: "Auto-select", + }) + require.NoError(t, err) + require.NotNil(t, assignment) + + assignments, err := svc.ListAssignments(context.Background(), "user-1", run.ID) + require.NoError(t, err) + assert.Len(t, assignments, 2) + // Both should be executors (not supervisor). + for _, a := range assignments { + assert.NotEqual(t, supervisor.ID, a.ToMemberID) + } +} + +func TestCompeteModeRejectsExceedingMaxAgents(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetCompeteMaxAgents(2) + + team := &model.AgentTeam{OwnerID: "user-1", Name: "Max Compete Team"} + require.NoError(t, repository.CreateTeam(db, team)) + + supervisor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr("profile-supervisor"), + Role: model.TeamMemberRoleSupervisor, + } + require.NoError(t, repository.AddTeamMember(db, supervisor)) + + workerIDs := make([]string, 3) + for i := range 3 { + e := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: strPtr(fmt.Sprintf("profile-executor-%d", i)), + Role: model.TeamMemberRoleExecutor, + } + require.NoError(t, repository.AddTeamMember(db, e)) + workerIDs[i] = e.ID + } + + run := &model.AgentTeamRun{ + TeamID: team.ID, + SessionID: "session-max-compete", + TriggerUserID: "user-1", + TriggerMessage: "too many", + Mode: model.TeamRunModeCompete, + Status: model.TeamRunStatusRunning, + } + require.NoError(t, repository.CreateTeamRun(db, run)) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "compete", + NextWorker: strings.Join(workerIDs, ","), + Instructions: "Too many workers", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) +} diff --git a/hub-server/internal/service/agentteam/agent_team_crud_test.go b/hub-server/internal/service/agentteam/agent_team_crud_test.go new file mode 100644 index 000000000..d11df3a04 --- /dev/null +++ b/hub-server/internal/service/agentteam/agent_team_crud_test.go @@ -0,0 +1,275 @@ +package agentteam + +import ( + "context" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func TestAgentTeamService_CreateTeam(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + mock.ExpectExec(`INSERT INTO "agent_teams"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + team, err := svc.CreateTeam(context.Background(), "user-1", "My Team", "A test team") + require.NoError(t, err) + assert.Equal(t, "My Team", team.Name) + assert.Equal(t, "user-1", team.OwnerID) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_CreateTeamEmptyName(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + _, err := svc.CreateTeam(context.Background(), "user-1", "", "desc") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_GetTeam(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "desc", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(rows) + + team, err := svc.GetTeam(context.Background(), "user-1", "team-1") + require.NoError(t, err) + assert.Equal(t, "My Team", team.Name) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_GetTeamWrongOwner(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + + team := &model.AgentTeam{OwnerID: "user-2", Name: "My Team", Description: "desc"} + require.NoError(t, repository.CreateTeam(db, team)) + + _, err := svc.GetTeam(context.Background(), "user-1", team.ID) + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) +} + +func TestAgentTeamService_GetTeamAllowsAgentProfileOwnerMemberRead(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + + team := &model.AgentTeam{OwnerID: "owner-user", Name: "Shared Team"} + require.NoError(t, repository.CreateTeam(db, team)) + memberAgent := &model.CustomAgent{ + OwnerUserID: "member-user", + Name: "Member Agent", + AgentType: "codex", + SystemPrompt: "Help the team", + } + require.NoError(t, repository.CreateCustomAgent(db, memberAgent)) + require.NoError(t, repository.AddTeamMember(db, &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &memberAgent.ID, + Role: model.TeamMemberRoleExecutor, + })) + + got, err := svc.GetTeam(context.Background(), "member-user", team.ID) + require.NoError(t, err) + assert.Equal(t, team.ID, got.ID) + + _, err = svc.GetTeam(context.Background(), "intruder-user", team.ID) + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) + + err = svc.UpdateTeam(context.Background(), "member-user", team.ID, "Should Not Write", "") + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) +} + +func TestAgentTeamService_GetTeamNotFound(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnError(gorm.ErrRecordNotFound) + + _, err := svc.GetTeam(context.Background(), "user-1", "team-1") + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_ListTeams(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + + require.NoError(t, repository.CreateTeam(db, &model.AgentTeam{OwnerID: "user-1", Name: "Team 1"})) + require.NoError(t, repository.CreateTeam(db, &model.AgentTeam{OwnerID: "user-1", Name: "Team 2"})) + + teams, err := svc.ListTeams(context.Background(), "user-1") + require.NoError(t, err) + assert.Len(t, teams, 2) +} + +func TestAgentTeamService_ListTeamsIncludesReadableMemberTeamsWithoutLeaking(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + + ownedTeam := &model.AgentTeam{OwnerID: "member-user", Name: "Owned Team"} + require.NoError(t, repository.CreateTeam(db, ownedTeam)) + + sharedTeam := &model.AgentTeam{OwnerID: "owner-user", Name: "Shared Team"} + require.NoError(t, repository.CreateTeam(db, sharedTeam)) + memberAgent := &model.CustomAgent{ + OwnerUserID: "member-user", + Name: "Member Agent", + AgentType: "codex", + SystemPrompt: "Read shared team", + } + require.NoError(t, repository.CreateCustomAgent(db, memberAgent)) + require.NoError(t, repository.AddTeamMember(db, &model.AgentTeamMember{ + TeamID: sharedTeam.ID, + AgentProfileID: &memberAgent.ID, + Role: model.TeamMemberRoleExecutor, + })) + + foreignTeam := &model.AgentTeam{OwnerID: "foreign-owner", Name: "Foreign Team"} + require.NoError(t, repository.CreateTeam(db, foreignTeam)) + foreignAgent := &model.CustomAgent{ + OwnerUserID: "foreign-member", + Name: "Foreign Agent", + AgentType: "codex", + SystemPrompt: "Not visible", + } + require.NoError(t, repository.CreateCustomAgent(db, foreignAgent)) + require.NoError(t, repository.AddTeamMember(db, &model.AgentTeamMember{ + TeamID: foreignTeam.ID, + AgentProfileID: &foreignAgent.ID, + Role: model.TeamMemberRoleExecutor, + })) + + teams, err := svc.ListTeams(context.Background(), "member-user") + require.NoError(t, err) + gotIDs := make([]string, 0, len(teams)) + for _, team := range teams { + gotIDs = append(gotIDs, team.ID) + } + assert.ElementsMatch(t, []string{ownedTeam.ID, sharedTeam.ID}, gotIDs) +} + +func TestAgentTeamService_UpdateTeam(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team first + rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "Old Name", "old desc", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(rows) + + // Update + mock.ExpectExec(`UPDATE "agent_teams"`). + WillReturnResult(sqlmock.NewResult(0, 1)) + + err := svc.UpdateTeam(context.Background(), "user-1", "team-1", "New Name", "new desc") + require.NoError(t, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_UpdateTeamWrongOwner(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-2", "Old Name", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(rows) + + err := svc.UpdateTeam(context.Background(), "user-1", "team-1", "New Name", "") + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_DeleteTeam(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team first + rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(rows) + + // No run history → deletable. + mock.ExpectQuery(`SELECT * FROM "agent_team_runs" WHERE team_id`). + WillReturnRows(sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "mode", "status", "created_at", "updated_at"})) + + // Transaction: batch-delete members in one statement (#2102 F10) + delete team. + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM "agent_team_members" WHERE team_id`). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec(`DELETE FROM "agent_teams"`). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + err := svc.DeleteTeam(context.Background(), "user-1", "team-1") + require.NoError(t, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_DeleteTeamWithRuns(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team first + rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(rows) + + // One historical run → 409, no delete attempted. + mock.ExpectQuery(`SELECT * FROM "agent_team_runs" WHERE team_id`). + WillReturnRows(sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "mode", "status", "created_at", "updated_at"}). + AddRow("run-1", "team-1", "sess-1", "user-1", "go", "supervisor", "completed", time.Now(), time.Now())) + + err := svc.DeleteTeam(context.Background(), "user-1", "team-1") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.TeamHasRuns) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_GetTeamWithMembers(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // List members + memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). + AddRow("member-1", "team-1", "agent-1", "executor", 0, time.Now()). + AddRow("member-2", "team-1", "agent-2", "supervisor", 1, time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). + WillReturnRows(memberRows) + + detail, err := svc.GetTeamWithMembers(context.Background(), "user-1", "team-1") + require.NoError(t, err) + assert.Equal(t, "My Team", detail.Name) + assert.Len(t, detail.Members, 2) + assert.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/hub-server/internal/service/agentteam/agent_team_member_test.go b/hub-server/internal/service/agentteam/agent_team_member_test.go new file mode 100644 index 000000000..f2811cf02 --- /dev/null +++ b/hub-server/internal/service/agentteam/agent_team_member_test.go @@ -0,0 +1,118 @@ +package agentteam + +import ( + "context" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAgentTeamService_AddTeamMember(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // Get agent profile + agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). + AddRow("agent-1", "user-1", "Agent 1", "", "codex", "prompt", "[]", "[]", "{}", nil, time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "custom_agents"`). + WillReturnRows(agentRows) + + // Add member — the service first lists existing members to reject + // duplicates with 409 instead of leaking a 23505 as a 500. + memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}) + mock.ExpectQuery(`SELECT * FROM "agent_team_members" WHERE team_id`). + WillReturnRows(memberRows) + + mock.ExpectExec(`INSERT INTO "agent_team_members"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + err := svc.AddTeamMember(context.Background(), "user-1", "team-1", "agent-1", model.TeamMemberRoleExecutor) + require.NoError(t, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_AddTeamMemberDuplicate(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // Get agent profile + agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). + AddRow("agent-1", "user-1", "Agent 1", "", "codex", "prompt", "[]", "[]", "{}", nil, time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "custom_agents"`). + WillReturnRows(agentRows) + + // Existing members include the same profile → 409, no insert. + memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). + AddRow("member-1", "team-1", "agent-1", model.TeamMemberRoleExecutor, 0, time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_team_members" WHERE team_id`). + WillReturnRows(memberRows) + + err := svc.AddTeamMember(context.Background(), "user-1", "team-1", "agent-1", model.TeamMemberRoleSupervisor) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.TeamMemberAlready) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_AddTeamMemberInvalidRole(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // Get agent profile + agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). + AddRow("agent-1", "user-1", "Agent 1", "", "codex", "prompt", "[]", "[]", "{}", nil, time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "custom_agents"`). + WillReturnRows(agentRows) + + err := svc.AddTeamMember(context.Background(), "user-1", "team-1", "agent-1", "invalid_role") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_RemoveTeamMember(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // Get member + memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). + AddRow("member-1", "team-1", "agent-1", "executor", 0, time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). + WillReturnRows(memberRows) + + // Remove member + mock.ExpectExec(`DELETE FROM "agent_team_members"`). + WillReturnResult(sqlmock.NewResult(0, 1)) + + err := svc.RemoveTeamMember(context.Background(), "user-1", "team-1", "member-1") + require.NoError(t, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/hub-server/internal/service/agentteam/agent_team_review_test.go b/hub-server/internal/service/agentteam/agent_team_review_test.go new file mode 100644 index 000000000..409f3c684 --- /dev/null +++ b/hub-server/internal/service/agentteam/agent_team_review_test.go @@ -0,0 +1,281 @@ +package agentteam + +import ( + "context" + "testing" + + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// ── Human Review Gate (ADR-008) ──────────────────────────────── + +func TestHumanReviewGate(t *testing.T) { + t.Run("disabled by default rejects review API", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + + // HandleRouteDecision should proceed normally (no pending_review). + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + require.NotNil(t, assignment) + + // Verify the run is still in running state, not pending_review. + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) + + // Calling ReviewDagPlan when disabled should fail. + _, err = svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: model.ReviewActionApprove, + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + }) + + t.Run("enabled sets pending_review after route decision", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + require.NotNil(t, assignment) + + // Verify the run is now pending_review. + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusPendingReview, gotRun.Status) + + // Verify review_pending event was recorded. + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + foundPending := false + for _, e := range events { + if e.Type == model.TeamEventReviewPending { + foundPending = true + break + } + } + assert.True(t, foundPending, "expected team.review.pending event") + }) + + t.Run("enabled approve transitions back to running", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + + // Review: approve + state, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: model.ReviewActionApprove, + Comment: "looks good", + }) + require.NoError(t, err) + assert.Equal(t, model.ReviewActionApprove, state.Action) + assert.Equal(t, "looks good", state.Comment) + + // Verify the run is back to running. + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) + + // Verify review_decided event was recorded. + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + foundDecided := false + for _, e := range events { + if e.Type == model.TeamEventReviewDecided { + foundDecided = true + break + } + } + assert.True(t, foundDecided, "expected team.review.decided event") + + // The replay projection must leave the review gate too: before the + // ReviewDecided replay fix the client-facing GetTeamRunState stayed + // pending_review forever even though the DB row was already running. + runState, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusRunning, runState.Status) + require.Len(t, runState.Reviews, 1) + assert.Equal(t, model.ReviewActionApprove, runState.Reviews[0].Action) + }) + + t.Run("decided review restores running in projection for discuss and modify", func(t *testing.T) { + for _, action := range []string{model.ReviewActionDiscuss, model.ReviewActionModify} { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + + _, err = svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: action, + Comment: "not yet", + }) + require.NoError(t, err) + + // Write side sets the DB row back to running for every decided + // action; the projection must agree. + runState, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusRunning, runState.Status, "action=%s", action) + } + }) + + t.Run("enabled discuss cancels pending assignments", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + + // Verify assignment is pending before review. + assignments, err := repository.ListAssignmentsByTeamRun(db, run.ID) + require.NoError(t, err) + require.Len(t, assignments, 1) + assert.Equal(t, model.AssignmentStatusPending, assignments[0].Status) + + // Review: discuss + state, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: model.ReviewActionDiscuss, + Comment: "needs more thought", + }) + require.NoError(t, err) + assert.Equal(t, model.ReviewActionDiscuss, state.Action) + + // Verify the run is back to running (so supervisor can re-plan). + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) + + // Verify assignments were cancelled. + assignments, err = repository.ListAssignmentsByTeamRun(db, run.ID) + require.NoError(t, err) + require.Len(t, assignments, 1) + assert.Equal(t, model.AssignmentStatusCancelled, assignments[0].Status) + }) + + t.Run("enabled modify cancels pending assignments with changes recorded", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + + state, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: model.ReviewActionModify, + Changes: []model.HumanReviewChange{ + {Field: "instructions", Value: "Do the other thing"}, + {Field: "next_worker", Value: "worker-2"}, + }, + Comment: "change the target", + }) + require.NoError(t, err) + assert.Equal(t, model.ReviewActionModify, state.Action) + assert.Len(t, state.Changes, 2) + assert.Equal(t, "instructions", state.Changes[0].Field) + + // Verify the run is back to running. + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) + + // Verify assignments were cancelled. + assignments, err := repository.ListAssignmentsByTeamRun(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.AssignmentStatusCancelled, assignments[0].Status) + }) + + t.Run("review rejects invalid action", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + + _, err = svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: "bogus", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + }) + + t.Run("review rejects non-pending_review status", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + _, _, _, run := seedAgentTeamRun(t, db) // run status is "running" + + _, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ + Action: model.ReviewActionApprove, + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + }) + + t.Run("review rejects non-owner user", func(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + svc.SetHumanReviewEnabled(true) + team, _, executor, run := seedAgentTeamRun(t, db) + + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do the thing", + }) + require.NoError(t, err) + + _, err = svc.ReviewDagPlan(context.Background(), "intruder-user", run.ID, model.HumanReviewDecision{ + Action: model.ReviewActionApprove, + }) + require.Error(t, err) + assert.Equal(t, errcode.AgentTaskNotFound, err) + }) +} diff --git a/hub-server/internal/service/agentteam/agent_team_run_test.go b/hub-server/internal/service/agentteam/agent_team_run_test.go index 306aafd21..7962639b8 100644 --- a/hub-server/internal/service/agentteam/agent_team_run_test.go +++ b/hub-server/internal/service/agentteam/agent_team_run_test.go @@ -8,9 +8,13 @@ import ( "testing" "time" + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" + "github.com/agenthub/hub-server/internal/bus" + "github.com/agenthub/hub-server/internal/errcode" "github.com/agenthub/hub-server/internal/model" "github.com/agenthub/hub-server/internal/repository" ) @@ -313,3 +317,567 @@ func TestGetTeamRunState_ProjectionUnchangedByParallelReads(t *testing.T) { // run-started event must be reflected exactly as in the serial version. require.NotEmpty(t, state.Status) } + +func TestAgentTeamService_GetTeamRun(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // Get run + runRows := sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "status", "created_at", "updated_at"}). + AddRow("run-1", "team-1", "session-1", "user-1", "hello", "completed", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_team_runs"`). + WillReturnRows(runRows) + + run, err := svc.GetTeamRun(context.Background(), "user-1", "team-1", "run-1") + require.NoError(t, err) + assert.Equal(t, "completed", run.Status) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_GetTeamRunWrongTeam(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // Get run (different team) + runRows := sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "status", "created_at", "updated_at"}). + AddRow("run-1", "team-2", "session-2", "user-1", "hello", "completed", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_team_runs"`). + WillReturnRows(runRows) + + _, err := svc.GetTeamRun(context.Background(), "user-1", "team-1", "run-1") + require.Error(t, err) + assert.Equal(t, errcode.AgentTaskNotFound, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_ListTeamRuns(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // List runs + runRows := sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "status", "created_at", "updated_at"}). + AddRow("run-1", "team-1", "session-1", "user-1", "msg1", "completed", time.Now(), time.Now()). + AddRow("run-2", "team-1", "session-2", "user-1", "msg2", "running", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_team_runs"`). + WillReturnRows(runRows) + + runs, err := svc.ListTeamRuns(context.Background(), "user-1", "team-1") + require.NoError(t, err) + assert.Len(t, runs, 2) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +// --- StartTeamRun tests --- + +func TestAgentTeamService_StartTeamRun_TeamNotFound(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnError(gorm.ErrRecordNotFound) + + _, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "") + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_StartTeamRun_EmptyMembers(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + svc := NewAgentTeamService(db, nil, nil) + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // List members (empty) + mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). + WillReturnRows(sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"})) + + _, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_StartTeamRun_Success(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + agentSvc := &mockAgentTeamAgentSvc{} + svc := NewAgentTeamService(db, agentSvc, nil) + eventBus, err := bus.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { eventBus.Close(context.Background()) }) + events := make(chan bus.Event, 1) + eventBus.Subscribe("team.run.started", func(ctx context.Context, event bus.Event) { + events <- event + }) + svc.SetBus(eventBus) + + now := time.Now() + agentProfileID := "agent-1" + + // Get team + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "desc", "", now, now) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + // List members (one supervisor) + memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). + AddRow("member-1", "team-1", agentProfileID, "supervisor", 0, now) + mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). + WillReturnRows(memberRows) + + // Transaction: Begin + mock.ExpectBegin() + + // CreateSession + mock.ExpectExec(`INSERT INTO "sessions"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + // CreateSessionMember (owner) + mock.ExpectExec(`INSERT INTO "session_members"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + // Batch query custom agents + agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). + AddRow("agent-1", "user-1", "My Agent", "", "codex", "prompt", "[]", "[]", "{}", nil, now, now) + mock.ExpectQuery(`SELECT * FROM "custom_agents" WHERE id IN`). + WillReturnRows(agentRows) + + // CreateAgentInstance + mock.ExpectExec(`INSERT INTO "agent_instances"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + // CreateSessionMember (agent) + mock.ExpectExec(`INSERT INTO "session_members"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + // AllocateSeqID (UPDATE ... RETURNING next_seq) + mock.ExpectQuery(`UPDATE sessions SET next_seq`). + WillReturnRows(sqlmock.NewRows([]string{"next_seq"}).AddRow(1)) + + // InsertMessage + mock.ExpectExec(`INSERT INTO "messages"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + // CreateTeamRun + mock.ExpectExec(`INSERT INTO "agent_team_runs"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + + // Transaction: Commit + mock.ExpectCommit() + + // AppendTeamEvent(team.run.started) after successful trigger. + mock.ExpectBegin() + mock.ExpectQuery(`SELECT id FROM agent_team_runs WHERE id`). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow("run-placeholder")) + mock.ExpectQuery(`COALESCE(MAX(seq)`). + WillReturnRows(sqlmock.NewRows([]string{"coalesce"}).AddRow(0)) + mock.ExpectExec(`INSERT INTO "agent_team_events"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + run, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "") + require.NoError(t, err) + assert.NotNil(t, run) + assert.Equal(t, "team-1", run.TeamID) + assert.Equal(t, model.TeamRunStatusRunning, run.Status) + assert.NotEmpty(t, agentSvc.triggerMessageID) + assert.Contains(t, agentSvc.modelParams, "structured_output_schema") + assert.Contains(t, agentSvc.modelParams, "AgentHub TeamRun supervisor mode") + event := readAgentTeamEvent(t, events) + assert.Equal(t, "team.run.started", event.Type) + payload, ok := event.Payload.(map[string]interface{}) + require.True(t, ok) + assert.Equal(t, "team-1", payload["team_id"]) + assert.Equal(t, run.ID, payload["run_id"]) + assert.Equal(t, run.SessionID, payload["session_id"]) + assert.Equal(t, "user-1", payload["user_id"]) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_StartTeamRunPassesTargetIDToSupervisor(t *testing.T) { + db, mock := newMockAgentTeamDB(t) + agentSvc := &mockAgentTeamAgentSvc{} + svc := NewAgentTeamService(db, agentSvc, nil) + + now := time.Now() + agentProfileID := "agent-1" + + teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). + AddRow("team-1", "user-1", "My Team", "desc", "", now, now) + mock.ExpectQuery(`SELECT * FROM "agent_teams"`). + WillReturnRows(teamRows) + + memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). + AddRow("member-1", "team-1", agentProfileID, "supervisor", 0, now) + mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). + WillReturnRows(memberRows) + + mock.ExpectBegin() + mock.ExpectExec(`INSERT INTO "sessions"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec(`INSERT INTO "session_members"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). + AddRow("agent-1", "user-1", "My Agent", "", "codex", "prompt", "[]", "[]", "{}", nil, now, now) + mock.ExpectQuery(`SELECT * FROM "custom_agents" WHERE id IN`). + WillReturnRows(agentRows) + mock.ExpectExec(`INSERT INTO "agent_instances"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec(`INSERT INTO "session_members"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectQuery(`UPDATE sessions SET next_seq`). + WillReturnRows(sqlmock.NewRows([]string{"next_seq"}).AddRow(1)) + mock.ExpectExec(`INSERT INTO "messages"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec(`INSERT INTO "agent_team_runs"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT id FROM agent_team_runs WHERE id`). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow("run-placeholder")) + mock.ExpectQuery(`COALESCE(MAX(seq)`). + WillReturnRows(sqlmock.NewRows([]string{"coalesce"}).AddRow(0)) + mock.ExpectExec(`INSERT INTO "agent_team_events"`). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + run, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "target-local-edge-1") + require.NoError(t, err) + require.NotNil(t, run.TargetID) + assert.Equal(t, "target-local-edge-1", *run.TargetID) + assert.Equal(t, "target-local-edge-1", agentSvc.targetID) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAgentTeamService_GetTeamRunStateReplaysEvents(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + + team := &model.AgentTeam{OwnerID: "user-1", Name: "State Team"} + require.NoError(t, repository.CreateTeam(db, team)) + + supervisorProfileID := "profile-supervisor" + executorProfileID := "profile-executor" + supervisor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &supervisorProfileID, + Role: model.TeamMemberRoleSupervisor, + } + executor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &executorProfileID, + Role: model.TeamMemberRoleExecutor, + } + require.NoError(t, repository.AddTeamMember(db, supervisor)) + require.NoError(t, repository.AddTeamMember(db, executor)) + + run := &model.AgentTeamRun{ + TeamID: team.ID, + SessionID: "session-1", + TriggerUserID: "user-1", + TriggerMessage: "ship it", + Status: model.TeamRunStatusRunning, + } + require.NoError(t, repository.CreateTeamRun(db, run)) + + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "Implement replay", + Status: model.AssignmentStatusDone, + Result: "done", + Depth: 1, + RunID: stringPtr("edge-run-1"), + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + + require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ + TeamRunID: run.ID, + Type: model.TeamEventRunStarted, + Payload: `{"status":"running"}`, + })) + require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ + TeamRunID: run.ID, + Type: model.TeamEventRouteDecided, + Payload: `{"action":"delegate","next_worker":"` + executor.ID + `","instructions":"Implement replay","reasoning":"needs executor"}`, + })) + require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ + TeamRunID: run.ID, + Type: model.TeamEventRunCompleted, + Payload: `{"summary":"done"}`, + })) + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + assert.Equal(t, run.ID, state.RunID) + assert.Equal(t, team.ID, state.TeamID) + assert.Equal(t, model.TeamRunStatusCompleted, state.Status) + assert.Equal(t, "done", state.TerminalReason) + require.Len(t, state.Members, 2) + assert.Equal(t, 1, state.Members[1].CompletedTasks) + require.Len(t, state.Assignments, 1) + assert.Equal(t, assignment.ID, state.Assignments[0].AssignmentID) + assert.Equal(t, "edge-run-1", state.Assignments[0].RunID) + require.Len(t, state.RouteLog, 1) + assert.Equal(t, "delegate", state.RouteLog[0].Action) + assert.Equal(t, executor.ID, state.RouteLog[0].NextWorker) +} + +func TestAgentTeamService_ListTeamTasksIsOwnerScoped(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusPending, + Objective: "Build task board", + InputRefs: "{}", + Attempt: 1, + RiskLevel: model.TeamTaskRiskNormal, + })) + + tasks, err := svc.ListTeamTasks(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, tasks, 1) + assert.Equal(t, "Build task board", tasks[0].Objective) + + _, err = svc.ListTeamTasks(context.Background(), "other-user", team.ID, run.ID) + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) +} + +func TestAgentTeamService_GetTeamRunStateProjectsDependenciesAndBudget(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + reviewerProfileID := "profile-reviewer" + reviewer := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &reviewerProfileID, + Role: model.TeamMemberRoleReviewer, + } + require.NoError(t, repository.AddTeamMember(db, reviewer)) + pending := &model.PendingAgentTask{ + AgentInstanceID: "agent-executor", + TriggeredByUserID: "user-1", + TriggerMessageID: "message-1", + Status: model.TaskStatusRunning, + EdgeRunID: "edge-run-budget", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(pending).Error) + + root := &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusRunning, + Objective: "Root task", + RunID: &pending.ID, + } + require.NoError(t, repository.CreateTeamTask(db, root)) + child := &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + ParentTaskID: &root.ID, + Status: model.TeamTaskStatusPending, + Objective: "Child task", + } + require.NoError(t, repository.CreateTeamTask(db, child)) + conflictingPending := &model.PendingAgentTask{ + AgentInstanceID: "agent-reviewer", + TriggeredByUserID: "user-1", + TriggerMessageID: "message-2", + Status: model.TaskStatusDone, + EdgeRunID: "edge-run-reviewer", + ExpireAt: time.Now().Add(time.Hour), + } + require.NoError(t, db.Create(conflictingPending).Error) + conflictingTask := &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: reviewer.ID, + Status: model.TeamTaskStatusDone, + Objective: "Review same file", + RunID: &conflictingPending.ID, + } + require.NoError(t, repository.CreateTeamTask(db, conflictingTask)) + + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: "edge-run-budget", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.result", + Payload: `{"success":true,"usage":{"input_tokens":1200,"output_tokens":800},"tokenLimit":200000,"tokensRemaining":198000,"usagePercent":1.0}`, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: "edge-run-budget", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.context_warning", + Payload: `{"usagePercent":86.5,"threshold":85}`, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: "edge-run-budget", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_requested", + Payload: `{"requestId":"req-1","toolUseId":"tool-1","toolName":"Bash","status":"pending"}`, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: "edge-run-budget", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.permission_decided", + Payload: `{"requestId":"req-1","decision":"allow","reason":"safe command"}`, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: "edge-run-budget", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.file_change", + Payload: `{"path":"hub-server/internal/service/agent_team.go","action":"modified","toolName":"apply_patch","status":"completed"}`, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: conflictingPending.ID, + EdgeRunID: "edge-run-reviewer", + SessionID: run.SessionID, + AgentInstanceID: "agent-reviewer", + EventType: "run.agent.file_change", + Payload: `{"path":"./hub-server/internal/service/agent_team.go","action":"modified","toolName":"review_patch","status":"completed"}`, + })) + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.Dependencies, 1) + assert.Equal(t, child.ID, state.Dependencies[0].TaskID) + assert.Equal(t, root.ID, state.Dependencies[0].DependsOnTaskID) + assert.Equal(t, "parent_task", state.Dependencies[0].Kind) + require.NotNil(t, state.Budget) + assert.Equal(t, int64(2000), state.Budget.TotalTokensUsed) + assert.Equal(t, int64(1200), state.Budget.InputTokens) + assert.Equal(t, int64(800), state.Budget.OutputTokens) + assert.Equal(t, int64(200000), state.Budget.TokenLimit) + assert.Equal(t, int64(198000), state.Budget.RemainingTokens) + assert.Equal(t, 86.5, state.Budget.UsagePercent) + assert.Equal(t, 2, state.Budget.RunCount) + assert.Equal(t, 1, state.Budget.ContextWarnings) + require.Len(t, state.Approvals, 1) + assert.Equal(t, pending.ID, state.Approvals[0].AgentTaskID) + assert.Equal(t, root.ID, state.Approvals[0].TeamTaskID) + assert.Equal(t, executor.ID, state.Approvals[0].MemberID) + assert.Equal(t, "req-1", state.Approvals[0].RequestID) + assert.Equal(t, "Bash", state.Approvals[0].ToolName) + assert.Equal(t, "tool-1", state.Approvals[0].ToolUseID) + assert.Equal(t, "allow", state.Approvals[0].Status) + assert.Equal(t, "safe command", state.Approvals[0].Reason) + require.NotNil(t, state.Approvals[0].DecidedAt) + require.Len(t, state.Artifacts, 2) + assert.Equal(t, pending.ID, state.Artifacts[0].AgentTaskID) + assert.Equal(t, root.ID, state.Artifacts[0].TeamTaskID) + assert.Equal(t, executor.ID, state.Artifacts[0].MemberID) + assert.Equal(t, "hub-server/internal/service/agent_team.go", state.Artifacts[0].Path) + assert.Equal(t, "modified", state.Artifacts[0].Action) + assert.Equal(t, "apply_patch", state.Artifacts[0].ToolName) + assert.Equal(t, "completed", state.Artifacts[0].Status) + assert.Equal(t, state.Artifacts[0].ConflictID, state.Artifacts[1].ConflictID) + require.Len(t, state.Conflicts, 1) + assert.Equal(t, "hub-server/internal/service/agent_team.go", state.Conflicts[0].Path) + assert.Equal(t, "pending", state.Conflicts[0].Status) + assert.ElementsMatch(t, []string{pending.ID, conflictingPending.ID}, state.Conflicts[0].AgentTaskIDs) + assert.ElementsMatch(t, []string{root.ID, conflictingTask.ID}, state.Conflicts[0].TeamTaskIDs) + assert.ElementsMatch(t, []string{executor.ID, reviewer.ID}, state.Conflicts[0].MemberIDs) + assert.ElementsMatch(t, []string{"modified"}, state.Conflicts[0].Actions) + + indexed, err := repository.ListTeamArtifactsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, indexed, 2) + assert.Equal(t, run.ID, indexed[0].TeamRunID) + require.NotNil(t, indexed[0].TeamTaskID) + assert.Equal(t, root.ID, *indexed[0].TeamTaskID) + require.NotNil(t, indexed[0].MemberID) + assert.Equal(t, executor.ID, *indexed[0].MemberID) + require.NotNil(t, indexed[0].AgentTaskID) + assert.Equal(t, pending.ID, *indexed[0].AgentTaskID) + require.NotNil(t, indexed[0].SourceEventID) + assert.NotEmpty(t, *indexed[0].SourceEventID) + assert.Equal(t, "apply_patch", indexed[0].ToolName) + assert.Equal(t, "hub-server/internal/service/agent_team.go", indexed[0].NormalizedPath) + assert.Equal(t, state.Artifacts[0].ConflictID, indexed[0].ConflictID) +} + +func TestAgentTeamService_ListTeamEventsIsOwnerScoped(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, _, run := seedAgentTeamRun(t, db) + require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ + TeamRunID: run.ID, + Type: model.TeamEventRouteRejected, + Payload: `{"reason":"invalid action"}`, + })) + + page, err := svc.ListTeamEvents(context.Background(), "user-1", team.ID, run.ID, 0, 50) + require.NoError(t, err) + require.Len(t, page.Items, 1) + assert.Equal(t, model.TeamEventRouteRejected, page.Items[0].Type) + assert.False(t, page.HasMore) + assert.Equal(t, 1, page.NextSeq) + + _, err = svc.ListTeamEvents(context.Background(), "other-user", team.ID, run.ID, 0, 50) + require.Error(t, err) + assert.Equal(t, errcode.AgentNotFound, err) +} + +func TestAgentTeamService_GetTeamRunStateKeepsTerminalAssignmentOverPendingProjection(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, supervisor, executor, run := seedAgentTeamRun(t, db) + pendingID := "pending-task-terminal-1" + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, TaskPrompt: "already failed", Status: model.AssignmentStatusFailed, + Result: "assignment timeout reached", RunID: &pendingID, + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + require.NoError(t, db.Exec( + "INSERT INTO pending_agent_tasks (id, agent_instance_id, trigger_message_id, triggered_by_user_id, status, expire_at) VALUES (?, ?, ?, ?, ?, ?)", + pendingID, "agent-executor", "msg-1", "user-1", model.TaskStatusRunning, time.Now().Add(time.Hour), + ).Error) + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.Assignments, 1) + assert.Equal(t, model.AssignmentStatusFailed, state.Assignments[0].Status) +} diff --git a/hub-server/internal/service/agentteam/agent_team_test.go b/hub-server/internal/service/agentteam/agent_team_test.go deleted file mode 100644 index d0ec278e7..000000000 --- a/hub-server/internal/service/agentteam/agent_team_test.go +++ /dev/null @@ -1,3176 +0,0 @@ -//nolint:gosec // 测试 fixture:凭据模式字符串用于构造测试用例,非真实凭据 -package agentteam - -import ( - "context" - "encoding/json" - "fmt" - "path/filepath" - "strings" - "testing" - "time" - - "github.com/DATA-DOG/go-sqlmock" - "github.com/agenthub/hub-server/internal/bus" - "github.com/agenthub/hub-server/internal/errcode" - "github.com/agenthub/hub-server/internal/model" - "github.com/agenthub/hub-server/internal/repository" - "github.com/glebarez/sqlite" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "gorm.io/gorm" -) - -func TestAgentTeamService_CreateTeam(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - mock.ExpectExec(`INSERT INTO "agent_teams"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - team, err := svc.CreateTeam(context.Background(), "user-1", "My Team", "A test team") - require.NoError(t, err) - assert.Equal(t, "My Team", team.Name) - assert.Equal(t, "user-1", team.OwnerID) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_CreateTeamEmptyName(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - _, err := svc.CreateTeam(context.Background(), "user-1", "", "desc") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_GetTeam(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "desc", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(rows) - - team, err := svc.GetTeam(context.Background(), "user-1", "team-1") - require.NoError(t, err) - assert.Equal(t, "My Team", team.Name) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_GetTeamWrongOwner(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - - team := &model.AgentTeam{OwnerID: "user-2", Name: "My Team", Description: "desc"} - require.NoError(t, repository.CreateTeam(db, team)) - - _, err := svc.GetTeam(context.Background(), "user-1", team.ID) - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) -} - -func TestAgentTeamService_GetTeamAllowsAgentProfileOwnerMemberRead(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - - team := &model.AgentTeam{OwnerID: "owner-user", Name: "Shared Team"} - require.NoError(t, repository.CreateTeam(db, team)) - memberAgent := &model.CustomAgent{ - OwnerUserID: "member-user", - Name: "Member Agent", - AgentType: "codex", - SystemPrompt: "Help the team", - } - require.NoError(t, repository.CreateCustomAgent(db, memberAgent)) - require.NoError(t, repository.AddTeamMember(db, &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &memberAgent.ID, - Role: model.TeamMemberRoleExecutor, - })) - - got, err := svc.GetTeam(context.Background(), "member-user", team.ID) - require.NoError(t, err) - assert.Equal(t, team.ID, got.ID) - - _, err = svc.GetTeam(context.Background(), "intruder-user", team.ID) - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) - - err = svc.UpdateTeam(context.Background(), "member-user", team.ID, "Should Not Write", "") - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) -} - -func TestAgentTeamService_GetTeamNotFound(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnError(gorm.ErrRecordNotFound) - - _, err := svc.GetTeam(context.Background(), "user-1", "team-1") - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_ListTeams(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - - require.NoError(t, repository.CreateTeam(db, &model.AgentTeam{OwnerID: "user-1", Name: "Team 1"})) - require.NoError(t, repository.CreateTeam(db, &model.AgentTeam{OwnerID: "user-1", Name: "Team 2"})) - - teams, err := svc.ListTeams(context.Background(), "user-1") - require.NoError(t, err) - assert.Len(t, teams, 2) -} - -func TestAgentTeamService_ListTeamsIncludesReadableMemberTeamsWithoutLeaking(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - - ownedTeam := &model.AgentTeam{OwnerID: "member-user", Name: "Owned Team"} - require.NoError(t, repository.CreateTeam(db, ownedTeam)) - - sharedTeam := &model.AgentTeam{OwnerID: "owner-user", Name: "Shared Team"} - require.NoError(t, repository.CreateTeam(db, sharedTeam)) - memberAgent := &model.CustomAgent{ - OwnerUserID: "member-user", - Name: "Member Agent", - AgentType: "codex", - SystemPrompt: "Read shared team", - } - require.NoError(t, repository.CreateCustomAgent(db, memberAgent)) - require.NoError(t, repository.AddTeamMember(db, &model.AgentTeamMember{ - TeamID: sharedTeam.ID, - AgentProfileID: &memberAgent.ID, - Role: model.TeamMemberRoleExecutor, - })) - - foreignTeam := &model.AgentTeam{OwnerID: "foreign-owner", Name: "Foreign Team"} - require.NoError(t, repository.CreateTeam(db, foreignTeam)) - foreignAgent := &model.CustomAgent{ - OwnerUserID: "foreign-member", - Name: "Foreign Agent", - AgentType: "codex", - SystemPrompt: "Not visible", - } - require.NoError(t, repository.CreateCustomAgent(db, foreignAgent)) - require.NoError(t, repository.AddTeamMember(db, &model.AgentTeamMember{ - TeamID: foreignTeam.ID, - AgentProfileID: &foreignAgent.ID, - Role: model.TeamMemberRoleExecutor, - })) - - teams, err := svc.ListTeams(context.Background(), "member-user") - require.NoError(t, err) - gotIDs := make([]string, 0, len(teams)) - for _, team := range teams { - gotIDs = append(gotIDs, team.ID) - } - assert.ElementsMatch(t, []string{ownedTeam.ID, sharedTeam.ID}, gotIDs) -} - -func TestAgentTeamService_UpdateTeam(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team first - rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "Old Name", "old desc", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(rows) - - // Update - mock.ExpectExec(`UPDATE "agent_teams"`). - WillReturnResult(sqlmock.NewResult(0, 1)) - - err := svc.UpdateTeam(context.Background(), "user-1", "team-1", "New Name", "new desc") - require.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_UpdateTeamWrongOwner(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-2", "Old Name", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(rows) - - err := svc.UpdateTeam(context.Background(), "user-1", "team-1", "New Name", "") - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_DeleteTeam(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team first - rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(rows) - - // No run history → deletable. - mock.ExpectQuery(`SELECT * FROM "agent_team_runs" WHERE team_id`). - WillReturnRows(sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "mode", "status", "created_at", "updated_at"})) - - // Transaction: batch-delete members in one statement (#2102 F10) + delete team. - mock.ExpectBegin() - mock.ExpectExec(`DELETE FROM "agent_team_members" WHERE team_id`). - WillReturnResult(sqlmock.NewResult(0, 0)) - mock.ExpectExec(`DELETE FROM "agent_teams"`). - WillReturnResult(sqlmock.NewResult(0, 1)) - mock.ExpectCommit() - - err := svc.DeleteTeam(context.Background(), "user-1", "team-1") - require.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_DeleteTeamWithRuns(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team first - rows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(rows) - - // One historical run → 409, no delete attempted. - mock.ExpectQuery(`SELECT * FROM "agent_team_runs" WHERE team_id`). - WillReturnRows(sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "mode", "status", "created_at", "updated_at"}). - AddRow("run-1", "team-1", "sess-1", "user-1", "go", "supervisor", "completed", time.Now(), time.Now())) - - err := svc.DeleteTeam(context.Background(), "user-1", "team-1") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.TeamHasRuns) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_AddTeamMember(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // Get agent profile - agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). - AddRow("agent-1", "user-1", "Agent 1", "", "codex", "prompt", "[]", "[]", "{}", nil, time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "custom_agents"`). - WillReturnRows(agentRows) - - // Add member — the service first lists existing members to reject - // duplicates with 409 instead of leaking a 23505 as a 500. - memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}) - mock.ExpectQuery(`SELECT * FROM "agent_team_members" WHERE team_id`). - WillReturnRows(memberRows) - - mock.ExpectExec(`INSERT INTO "agent_team_members"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - err := svc.AddTeamMember(context.Background(), "user-1", "team-1", "agent-1", model.TeamMemberRoleExecutor) - require.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_AddTeamMemberDuplicate(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // Get agent profile - agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). - AddRow("agent-1", "user-1", "Agent 1", "", "codex", "prompt", "[]", "[]", "{}", nil, time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "custom_agents"`). - WillReturnRows(agentRows) - - // Existing members include the same profile → 409, no insert. - memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). - AddRow("member-1", "team-1", "agent-1", model.TeamMemberRoleExecutor, 0, time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_team_members" WHERE team_id`). - WillReturnRows(memberRows) - - err := svc.AddTeamMember(context.Background(), "user-1", "team-1", "agent-1", model.TeamMemberRoleSupervisor) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.TeamMemberAlready) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_AddTeamMemberInvalidRole(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // Get agent profile - agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). - AddRow("agent-1", "user-1", "Agent 1", "", "codex", "prompt", "[]", "[]", "{}", nil, time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "custom_agents"`). - WillReturnRows(agentRows) - - err := svc.AddTeamMember(context.Background(), "user-1", "team-1", "agent-1", "invalid_role") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_RemoveTeamMember(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // Get member - memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). - AddRow("member-1", "team-1", "agent-1", "executor", 0, time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). - WillReturnRows(memberRows) - - // Remove member - mock.ExpectExec(`DELETE FROM "agent_team_members"`). - WillReturnResult(sqlmock.NewResult(0, 1)) - - err := svc.RemoveTeamMember(context.Background(), "user-1", "team-1", "member-1") - require.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_GetTeamWithMembers(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // List members - memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). - AddRow("member-1", "team-1", "agent-1", "executor", 0, time.Now()). - AddRow("member-2", "team-1", "agent-2", "supervisor", 1, time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). - WillReturnRows(memberRows) - - detail, err := svc.GetTeamWithMembers(context.Background(), "user-1", "team-1") - require.NoError(t, err) - assert.Equal(t, "My Team", detail.Name) - assert.Len(t, detail.Members, 2) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_GetTeamRun(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // Get run - runRows := sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "status", "created_at", "updated_at"}). - AddRow("run-1", "team-1", "session-1", "user-1", "hello", "completed", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_team_runs"`). - WillReturnRows(runRows) - - run, err := svc.GetTeamRun(context.Background(), "user-1", "team-1", "run-1") - require.NoError(t, err) - assert.Equal(t, "completed", run.Status) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_GetTeamRunWrongTeam(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // Get run (different team) - runRows := sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "status", "created_at", "updated_at"}). - AddRow("run-1", "team-2", "session-2", "user-1", "hello", "completed", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_team_runs"`). - WillReturnRows(runRows) - - _, err := svc.GetTeamRun(context.Background(), "user-1", "team-1", "run-1") - require.Error(t, err) - assert.Equal(t, errcode.AgentTaskNotFound, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_ListTeamRuns(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // List runs - runRows := sqlmock.NewRows([]string{"id", "team_id", "session_id", "trigger_user_id", "trigger_message", "status", "created_at", "updated_at"}). - AddRow("run-1", "team-1", "session-1", "user-1", "msg1", "completed", time.Now(), time.Now()). - AddRow("run-2", "team-1", "session-2", "user-1", "msg2", "running", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_team_runs"`). - WillReturnRows(runRows) - - runs, err := svc.ListTeamRuns(context.Background(), "user-1", "team-1") - require.NoError(t, err) - assert.Len(t, runs, 2) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -// newMockAgentTeamDB creates a sqlmock-backed gorm.DB for agent team tests. -func newMockAgentTeamDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock) { - t.Helper() - db, mock, _ := newMockDBAgent(t) - return db, mock -} - -// mockAgentTeamAgentSvc implements agentTeamAgentSvc for tests. -type mockAgentTeamAgentSvc struct { - triggerMessageID string - targetID string - modelParams string - returnTaskID string -} - -func (m *mockAgentTeamAgentSvc) AddAgentToSession(ctx context.Context, userID, sessionID, agentType, customAgentID, displayName string) (*model.AgentInstance, error) { - return &model.AgentInstance{}, nil -} - -func (m *mockAgentTeamAgentSvc) TriggerAgentTask(ctx context.Context, userID, triggerMessageID, targetAgentInstanceID, targetAgentType, targetCustomAgentID, modelParams, targetID string) (*model.PendingAgentTask, error) { - m.triggerMessageID = triggerMessageID - m.targetID = targetID - m.modelParams = modelParams - taskID := m.returnTaskID - if taskID == "" { - taskID = "task-1" - } - return &model.PendingAgentTask{ID: taskID}, nil -} - -type mockAgentTeamControlSvc struct { - calls []agentTeamControlCall -} - -type agentTeamControlCall struct { - userID string - deviceID string - payload model.AgentControlPayload -} - -func (m *mockAgentTeamControlSvc) DeliverToDesktopDevice(ctx context.Context, userID, deviceID string, payload model.AgentControlPayload) error { - m.calls = append(m.calls, agentTeamControlCall{ - userID: userID, - deviceID: deviceID, - payload: payload, - }) - return nil -} - -// --- StartTeamRun tests --- - -func TestAgentTeamService_StartTeamRun_TeamNotFound(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnError(gorm.ErrRecordNotFound) - - _, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "") - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_StartTeamRun_EmptyMembers(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - svc := NewAgentTeamService(db, nil, nil) - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "", "", time.Now(), time.Now()) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // List members (empty) - mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). - WillReturnRows(sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"})) - - _, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_StartTeamRun_Success(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - agentSvc := &mockAgentTeamAgentSvc{} - svc := NewAgentTeamService(db, agentSvc, nil) - eventBus, err := bus.New() - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { eventBus.Close(context.Background()) }) - events := make(chan bus.Event, 1) - eventBus.Subscribe("team.run.started", func(ctx context.Context, event bus.Event) { - events <- event - }) - svc.SetBus(eventBus) - - now := time.Now() - agentProfileID := "agent-1" - - // Get team - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "desc", "", now, now) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - // List members (one supervisor) - memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). - AddRow("member-1", "team-1", agentProfileID, "supervisor", 0, now) - mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). - WillReturnRows(memberRows) - - // Transaction: Begin - mock.ExpectBegin() - - // CreateSession - mock.ExpectExec(`INSERT INTO "sessions"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - // CreateSessionMember (owner) - mock.ExpectExec(`INSERT INTO "session_members"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - // Batch query custom agents - agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). - AddRow("agent-1", "user-1", "My Agent", "", "codex", "prompt", "[]", "[]", "{}", nil, now, now) - mock.ExpectQuery(`SELECT * FROM "custom_agents" WHERE id IN`). - WillReturnRows(agentRows) - - // CreateAgentInstance - mock.ExpectExec(`INSERT INTO "agent_instances"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - // CreateSessionMember (agent) - mock.ExpectExec(`INSERT INTO "session_members"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - // AllocateSeqID (UPDATE ... RETURNING next_seq) - mock.ExpectQuery(`UPDATE sessions SET next_seq`). - WillReturnRows(sqlmock.NewRows([]string{"next_seq"}).AddRow(1)) - - // InsertMessage - mock.ExpectExec(`INSERT INTO "messages"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - // CreateTeamRun - mock.ExpectExec(`INSERT INTO "agent_team_runs"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - - // Transaction: Commit - mock.ExpectCommit() - - // AppendTeamEvent(team.run.started) after successful trigger. - mock.ExpectBegin() - mock.ExpectQuery(`SELECT id FROM agent_team_runs WHERE id`). - WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow("run-placeholder")) - mock.ExpectQuery(`COALESCE(MAX(seq)`). - WillReturnRows(sqlmock.NewRows([]string{"coalesce"}).AddRow(0)) - mock.ExpectExec(`INSERT INTO "agent_team_events"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectCommit() - - run, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "") - require.NoError(t, err) - assert.NotNil(t, run) - assert.Equal(t, "team-1", run.TeamID) - assert.Equal(t, model.TeamRunStatusRunning, run.Status) - assert.NotEmpty(t, agentSvc.triggerMessageID) - assert.Contains(t, agentSvc.modelParams, "structured_output_schema") - assert.Contains(t, agentSvc.modelParams, "AgentHub TeamRun supervisor mode") - event := readAgentTeamEvent(t, events) - assert.Equal(t, "team.run.started", event.Type) - payload, ok := event.Payload.(map[string]interface{}) - require.True(t, ok) - assert.Equal(t, "team-1", payload["team_id"]) - assert.Equal(t, run.ID, payload["run_id"]) - assert.Equal(t, run.SessionID, payload["session_id"]) - assert.Equal(t, "user-1", payload["user_id"]) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_StartTeamRunPassesTargetIDToSupervisor(t *testing.T) { - db, mock := newMockAgentTeamDB(t) - agentSvc := &mockAgentTeamAgentSvc{} - svc := NewAgentTeamService(db, agentSvc, nil) - - now := time.Now() - agentProfileID := "agent-1" - - teamRows := sqlmock.NewRows([]string{"id", "owner_id", "name", "description", "avatar_url", "created_at", "updated_at"}). - AddRow("team-1", "user-1", "My Team", "desc", "", now, now) - mock.ExpectQuery(`SELECT * FROM "agent_teams"`). - WillReturnRows(teamRows) - - memberRows := sqlmock.NewRows([]string{"id", "team_id", "agent_profile_id", "role", "position", "created_at"}). - AddRow("member-1", "team-1", agentProfileID, "supervisor", 0, now) - mock.ExpectQuery(`SELECT * FROM "agent_team_members"`). - WillReturnRows(memberRows) - - mock.ExpectBegin() - mock.ExpectExec(`INSERT INTO "sessions"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectExec(`INSERT INTO "session_members"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - agentRows := sqlmock.NewRows([]string{"id", "owner_user_id", "name", "avatar_url", "agent_type", "system_prompt", "capability_tags", "tool_whitelist", "model_params", "deleted_at", "created_at", "updated_at"}). - AddRow("agent-1", "user-1", "My Agent", "", "codex", "prompt", "[]", "[]", "{}", nil, now, now) - mock.ExpectQuery(`SELECT * FROM "custom_agents" WHERE id IN`). - WillReturnRows(agentRows) - mock.ExpectExec(`INSERT INTO "agent_instances"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectExec(`INSERT INTO "session_members"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectQuery(`UPDATE sessions SET next_seq`). - WillReturnRows(sqlmock.NewRows([]string{"next_seq"}).AddRow(1)) - mock.ExpectExec(`INSERT INTO "messages"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectExec(`INSERT INTO "agent_team_runs"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectCommit() - - mock.ExpectBegin() - mock.ExpectQuery(`SELECT id FROM agent_team_runs WHERE id`). - WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow("run-placeholder")) - mock.ExpectQuery(`COALESCE(MAX(seq)`). - WillReturnRows(sqlmock.NewRows([]string{"coalesce"}).AddRow(0)) - mock.ExpectExec(`INSERT INTO "agent_team_events"`). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectCommit() - - run, err := svc.StartTeamRun(context.Background(), "user-1", "team-1", "hello", "target-local-edge-1") - require.NoError(t, err) - require.NotNil(t, run.TargetID) - assert.Equal(t, "target-local-edge-1", *run.TargetID) - assert.Equal(t, "target-local-edge-1", agentSvc.targetID) - assert.NoError(t, mock.ExpectationsWereMet()) -} - -func TestAgentTeamService_CompleteAssignmentPublishesEvent(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - eventBus, err := bus.New() - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { eventBus.Close(context.Background()) }) - events := make(chan bus.Event, 1) - eventBus.Subscribe(bus.EventTypeTeamAssignmentDone, func(ctx context.Context, event bus.Event) { - events <- event - }) - svc.SetBus(eventBus) - - _, supervisor, executor, run := seedAgentTeamRun(t, db) - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "Ship it", - Status: model.AssignmentStatusRunning, - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - - require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", assignment.ID, "done text")) - - event := readAgentTeamEvent(t, events) - assert.Equal(t, bus.EventTypeTeamAssignmentDone, event.Type) - payload, ok := event.Payload.(map[string]interface{}) - require.True(t, ok) - assert.Equal(t, run.ID, payload["team_run_id"]) - assert.Equal(t, assignment.ID, payload["assignment_id"]) - assert.Equal(t, run.SessionID, payload["session_id"]) - assert.Equal(t, "done text", payload["result"]) -} - -func TestAgentTeamService_FailAssignmentPublishesEvent(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - eventBus, err := bus.New() - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { eventBus.Close(context.Background()) }) - events := make(chan bus.Event, 1) - eventBus.Subscribe("team.assignment.failed", func(ctx context.Context, event bus.Event) { - events <- event - }) - svc.SetBus(eventBus) - - _, supervisor, executor, run := seedAgentTeamRun(t, db) - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "Ship it", - Status: model.AssignmentStatusRunning, - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - - require.NoError(t, svc.FailAssignment(context.Background(), "user-1", assignment.ID, "blocked")) - - event := readAgentTeamEvent(t, events) - assert.Equal(t, "team.assignment.failed", event.Type) - payload, ok := event.Payload.(map[string]interface{}) - require.True(t, ok) - assert.Equal(t, run.ID, payload["team_run_id"]) - assert.Equal(t, assignment.ID, payload["assignment_id"]) - assert.Equal(t, run.SessionID, payload["session_id"]) - assert.Equal(t, "blocked", payload["reason"]) -} - -func TestAgentTeamService_GetTeamRunStateReplaysEvents(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - - team := &model.AgentTeam{OwnerID: "user-1", Name: "State Team"} - require.NoError(t, repository.CreateTeam(db, team)) - - supervisorProfileID := "profile-supervisor" - executorProfileID := "profile-executor" - supervisor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &supervisorProfileID, - Role: model.TeamMemberRoleSupervisor, - } - executor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &executorProfileID, - Role: model.TeamMemberRoleExecutor, - } - require.NoError(t, repository.AddTeamMember(db, supervisor)) - require.NoError(t, repository.AddTeamMember(db, executor)) - - run := &model.AgentTeamRun{ - TeamID: team.ID, - SessionID: "session-1", - TriggerUserID: "user-1", - TriggerMessage: "ship it", - Status: model.TeamRunStatusRunning, - } - require.NoError(t, repository.CreateTeamRun(db, run)) - - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "Implement replay", - Status: model.AssignmentStatusDone, - Result: "done", - Depth: 1, - RunID: stringPtr("edge-run-1"), - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - - require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ - TeamRunID: run.ID, - Type: model.TeamEventRunStarted, - Payload: `{"status":"running"}`, - })) - require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ - TeamRunID: run.ID, - Type: model.TeamEventRouteDecided, - Payload: `{"action":"delegate","next_worker":"` + executor.ID + `","instructions":"Implement replay","reasoning":"needs executor"}`, - })) - require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ - TeamRunID: run.ID, - Type: model.TeamEventRunCompleted, - Payload: `{"summary":"done"}`, - })) - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - assert.Equal(t, run.ID, state.RunID) - assert.Equal(t, team.ID, state.TeamID) - assert.Equal(t, model.TeamRunStatusCompleted, state.Status) - assert.Equal(t, "done", state.TerminalReason) - require.Len(t, state.Members, 2) - assert.Equal(t, 1, state.Members[1].CompletedTasks) - require.Len(t, state.Assignments, 1) - assert.Equal(t, assignment.ID, state.Assignments[0].AssignmentID) - assert.Equal(t, "edge-run-1", state.Assignments[0].RunID) - require.Len(t, state.RouteLog, 1) - assert.Equal(t, "delegate", state.RouteLog[0].Action) - assert.Equal(t, executor.ID, state.RouteLog[0].NextWorker) -} - -func TestAgentTeamService_HandleRouteDecisionCreatesAssignmentAndAuditEvents(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, supervisor, executor, run := seedAgentTeamRun(t, db) - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Implement the replay UI", - Reasoning: "executor owns UI work", - Reason: "worker owns fixture implementation", - Context: "state endpoint is ready", - AgentID: executor.ID, - ParentTaskID: "parent-task-1", - CorrelationID: "corr-route-1", - }) - require.NoError(t, err) - require.NotNil(t, assignment) - assert.Equal(t, supervisor.ID, assignment.FromMemberID) - assert.Equal(t, executor.ID, assignment.ToMemberID) - assert.Equal(t, model.AssignmentTypeDelegate, assignment.Type) - assert.Equal(t, model.AssignmentStatusPending, assignment.Status) - assert.Equal(t, 1, assignment.Depth) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 3) - assert.Equal(t, model.TeamEventRouteDecided, events[0].Type) - assert.Equal(t, model.TeamEventAssignmentCreated, events[1].Type) - assert.Equal(t, model.TeamEventTaskCreated, events[2].Type) - var accepted model.CoordinatorRouteDecision - require.NoError(t, json.Unmarshal([]byte(events[0].Payload), &accepted)) - require.NotEmpty(t, accepted.SubtaskID) - assert.True(t, accepted.Accepted) - assert.Equal(t, executor.ID, accepted.AgentID) - assert.Equal(t, "parent-task-1", accepted.ParentTaskID) - assert.Equal(t, "worker owns fixture implementation", accepted.Reason) - assert.Equal(t, "corr-route-1", accepted.CorrelationID) - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.RouteLog, 1) - assert.Equal(t, "delegate", state.RouteLog[0].Action) - assert.Equal(t, "corr-route-1", state.RouteLog[0].CorrelationID) - require.Len(t, state.RouteAuditLog, 1) - assert.Equal(t, "accepted", state.RouteAuditLog[0].Status) - assert.Equal(t, "corr-route-1", state.RouteAuditLog[0].CorrelationID) - assert.Equal(t, accepted.SubtaskID, state.RouteAuditLog[0].SubtaskID) - assert.Equal(t, "parent-task-1", state.RouteAuditLog[0].ParentTaskID) - assert.Equal(t, executor.ID, state.RouteAuditLog[0].AgentID) - assert.Equal(t, "worker owns fixture implementation", state.RouteAuditLog[0].Reason) - require.Len(t, state.Assignments, 1) - assert.Equal(t, assignment.ID, state.Assignments[0].AssignmentID) - assert.Equal(t, 1, state.Members[1].ActiveTasks) - require.Len(t, state.Tasks, 1) - assert.Equal(t, assignment.ID, state.Tasks[0].AssignmentID) - assert.Equal(t, executor.ID, state.Tasks[0].AssigneeMemberID) - assert.Equal(t, "parent-task-1", state.Tasks[0].ParentTaskID) - assert.Equal(t, "Implement the replay UI", state.Tasks[0].Objective) - assert.Equal(t, model.TeamTaskStatusPending, state.Tasks[0].Status) -} - -func TestAgentTeamService_HandleRouteDecisionRejectsMissingWorkerWithAuditEvent(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, _, run := seedAgentTeamRun(t, db) - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: "missing-member", - Instructions: "Do work", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) - assert.Contains(t, events[0].Payload, "next_worker") - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.RouteAuditLog, 1) - assert.Equal(t, "rejected", state.RouteAuditLog[0].Status) - assert.Equal(t, "missing-member", state.RouteAuditLog[0].AgentID) - assert.Equal(t, "next_worker is not a team member", state.RouteAuditLog[0].Reason) -} - -func TestAgentTeamService_HandleRouteDecisionRejectsWhenTaskLimitReached(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, supervisor, executor, run := seedAgentTeamRun(t, db) - for i := 0; i < model.MaxTasksPerTeamRun; i++ { - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "existing task", - Status: model.AssignmentStatusDone, - Depth: 1, - })) - } - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "one more task", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) - assert.Contains(t, events[0].Payload, "task limit") -} - -func TestAgentTeamService_HandleRouteDecisionRejectsWhenActiveSubagentLimitReached(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, supervisor, executor, run := seedAgentTeamRun(t, db) - for i := 0; i < model.MaxActiveSubAgentsPerRun; i++ { - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "active task", - Status: model.AssignmentStatusRunning, - Depth: 1, - })) - } - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "blocked by active limit", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) - assert.Contains(t, events[0].Payload, "active subagent limit") -} - -func TestAgentTeamService_HandleRouteDecisionRejectsWhenRouteRepeatLimitReached(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - repeated := model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "repeat me", - Reasoning: "same route", - } - for i := 0; i < model.MaxRouteRepeats; i++ { - require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ - TeamRunID: run.ID, - Type: model.TeamEventRouteDecided, - Payload: mustJSON(t, repeated), - })) - } - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, repeated) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, model.MaxRouteRepeats+1) - assert.Equal(t, model.TeamEventRouteRejected, events[len(events)-1].Type) - assert.Contains(t, events[len(events)-1].Payload, "route repeat limit") -} - -func TestAgentTeamService_HandleRouteDecisionRejectsTimedOutAssignment(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, supervisor, executor, run := seedAgentTeamRun(t, db) - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "stale task", - Status: model.AssignmentStatusRunning, - Depth: 1, - CreatedAt: time.Now().Add(-model.DefaultAssignmentTimeout - time.Minute), - })) - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "blocked by timeout", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) - assert.Contains(t, events[0].Payload, "assignment timeout") -} - -func TestAgentTeamService_HandleRouteDecisionRejectsBudgetExceeded(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - agentTaskID := "budget-task-1" - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusDone, - Objective: "spent budget", - RunID: &agentTaskID, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: agentTaskID, - EdgeRunID: "edge-budget", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.result", - Payload: `{"success":true,"usage":{"input_tokens":600,"output_tokens":400},"tokenLimit":1000}`, - })) - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "blocked by budget", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) - assert.Contains(t, events[0].Payload, "budget exceeded") -} - -func TestAgentTeamService_CreateAssignmentRejectsDelegationDepthLimit(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - _, supervisor, _, run := seedAgentTeamRun(t, db) - supervisor2 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-2") - supervisor3 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-3") - supervisor4 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-4") - supervisor5 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-5") - - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: supervisor2.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "depth 1", - Status: model.AssignmentStatusDone, - Depth: 1, - })) - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor2.ID, - ToMemberID: supervisor3.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "depth 2", - Status: model.AssignmentStatusDone, - Depth: 2, - })) - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor3.ID, - ToMemberID: supervisor4.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "depth 3", - Status: model.AssignmentStatusDone, - Depth: model.MaxDelegationDepth, - })) - - assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor4.ID, supervisor5.ID, model.AssignmentTypeDelegate, "too deep", "") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) -} - -func TestAgentTeamService_CreateAssignmentRejectsTeamRunTaskLimit(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - - for i := 0; i < model.MaxTasksPerTeamRun; i++ { - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "completed task", - Status: model.AssignmentStatusDone, - Depth: 1, - })) - } - - assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor.ID, executor.ID, model.AssignmentTypeDelegate, "too many", "") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) -} - -func TestAgentTeamService_CreateAssignmentRejectsDelegationCycle(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - _, supervisor, _, run := seedAgentTeamRun(t, db) - supervisor2 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-cycle-2") - supervisor3 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-cycle-3") - - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: supervisor2.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "cycle depth 1", - Status: model.AssignmentStatusDone, - Depth: 1, - })) - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor2.ID, - ToMemberID: supervisor3.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "cycle depth 2", - Status: model.AssignmentStatusDone, - Depth: 2, - })) - - assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor3.ID, supervisor.ID, model.AssignmentTypeDelegate, "cycle", "") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) -} - -func TestAgentTeamService_CreateAssignmentUsesConfiguredDelegationDepthLimit(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamServiceWithGuardrails(db, nil, nil, AgentTeamGuardrails{ - MaxDelegationDepth: 1, - }) - _, supervisor, _, run := seedAgentTeamRun(t, db) - supervisor2 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-config-depth-2") - supervisor3 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-config-depth-3") - - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: supervisor2.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "depth 1", - Status: model.AssignmentStatusDone, - Depth: 1, - })) - - assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor2.ID, supervisor3.ID, model.AssignmentTypeDelegate, "too deep for config", "") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) -} - -func TestAgentTeamService_CreateAssignmentUsesConfiguredTeamRunTaskLimit(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamServiceWithGuardrails(db, nil, nil, AgentTeamGuardrails{ - MaxTasksPerTeamRun: 1, - }) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - - require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "first task", - Status: model.AssignmentStatusDone, - Depth: 1, - })) - - assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor.ID, executor.ID, model.AssignmentTypeDelegate, "over configured task limit", "") - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, assignment) -} - -func TestAgentTeamService_ListTeamTasksIsOwnerScoped(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusPending, - Objective: "Build task board", - InputRefs: "{}", - Attempt: 1, - RiskLevel: model.TeamTaskRiskNormal, - })) - - tasks, err := svc.ListTeamTasks(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, tasks, 1) - assert.Equal(t, "Build task board", tasks[0].Objective) - - _, err = svc.ListTeamTasks(context.Background(), "other-user", team.ID, run.ID) - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) -} - -func TestAgentTeamService_DispatchAssignmentBindsTeamTaskToPendingAgentTask(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - agentSvc := &mockAgentTeamAgentSvc{returnTaskID: "task-dispatch-1"} - svc := NewAgentTeamService(db, agentSvc, nil) - team, supervisor, executor, run := seedAgentTeamRun(t, db) - seedTeamRunSession(t, db, run.SessionID, "user-1", executor) - - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "Implement replay", - Context: "include events", - Status: model.AssignmentStatusPending, - Depth: 1, - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - assignmentID := assignment.ID - teamTask := &model.AgentTeamTask{ - TeamRunID: run.ID, - AssignmentID: &assignmentID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusPending, - Objective: assignment.TaskPrompt, - } - require.NoError(t, repository.CreateTeamTask(db, teamTask)) - - require.NoError(t, svc.DispatchAssignment(context.Background(), "user-1", assignment.ID)) - - var reloadedAssignment model.AgentTeamAssignment - require.NoError(t, db.Where("id = ?", assignment.ID).First(&reloadedAssignment).Error) - require.NotNil(t, reloadedAssignment.RunID) - assert.Equal(t, model.AssignmentStatusRunning, reloadedAssignment.Status) - - // Seed the pending task that the mock agent service would have created. - triggerMsgID := "msg-trigger-dispatch" - require.NoError(t, db.Exec( - "INSERT INTO pending_agent_tasks (id, agent_instance_id, trigger_message_id, triggered_by_user_id, status, expire_at) VALUES (?, ?, ?, ?, ?, ?)", - *reloadedAssignment.RunID, "agent-executor", triggerMsgID, "user-1", model.TaskStatusRunning, time.Now().Add(1*time.Hour), - ).Error) - require.NoError(t, db.Exec( - "INSERT INTO messages (id, session_id, seq_id, client_msg_id, sender_type, sender_id, content_type, content, created_at) VALUES (?, ?, 1, ?, 'user', ?, 'text', ?, ?)", - triggerMsgID, run.SessionID, triggerMsgID, "user-1", "Task: Implement replay\nContext: include events", time.Now(), - ).Error) - - var reloadedTask model.AgentTeamTask - require.NoError(t, db.Where("id = ?", teamTask.ID).First(&reloadedTask).Error) - require.NotNil(t, reloadedTask.RunID) - assert.Equal(t, *reloadedAssignment.RunID, *reloadedTask.RunID) - assert.Equal(t, model.TeamTaskStatusDispatched, reloadedTask.Status) - - var pending model.PendingAgentTask - require.NoError(t, db.Where("id = ?", *reloadedTask.RunID).First(&pending).Error) - assert.Equal(t, "agent-executor", pending.AgentInstanceID) - assert.NotEmpty(t, pending.TriggerMessageID) - - var triggerMessage model.Message - require.NoError(t, db.Where("id = ?", pending.TriggerMessageID).First(&triggerMessage).Error) - assert.Contains(t, triggerMessage.Content, "Implement replay") - assert.Contains(t, triggerMessage.Content, "include events") - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventAssignmentDispatched, events[0].Type) - - require.NoError(t, repository.UpdatePendingTaskStatusWithEdgeRunID(db, pending.ID, model.TaskStatusRunning, "", "edge-run-1")) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: "edge-run-1", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: model.RunEventTypeOutputBatch, - Payload: `{"content":"runtime output"}`, - })) - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.Assignments, 1) - assert.Equal(t, model.AssignmentStatusRunning, state.Assignments[0].Status) - assert.Equal(t, pending.ID, state.Assignments[0].AgentTaskID) - assert.Equal(t, "edge-run-1", state.Assignments[0].EdgeRunID) - require.Len(t, state.Tasks, 1) - assert.Equal(t, model.TeamTaskStatusRunning, state.Tasks[0].Status) - assert.Equal(t, pending.ID, state.Tasks[0].AgentTaskID) - assert.Equal(t, "edge-run-1", state.Tasks[0].EdgeRunID) - require.Len(t, state.RunEvents, 1) - assert.Equal(t, pending.ID, state.RunEvents[0].AgentTaskID) - assert.Equal(t, "edge-run-1", state.RunEvents[0].EdgeRunID) - assert.Equal(t, model.RunEventTypeOutputBatch, state.RunEvents[0].EventType) - assert.JSONEq(t, `{"content":"runtime output"}`, state.RunEvents[0].Payload) -} - -func TestAgentTeamService_DispatchAssignmentPassesTeamRunTargetID(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - agentSvc := &mockAgentTeamAgentSvc{returnTaskID: "task-dispatch-1"} - svc := NewAgentTeamService(db, agentSvc, nil) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - targetID := "target-local-edge-1" - run.TargetID = &targetID - require.NoError(t, db.Model(&model.AgentTeamRun{}).Where("id = ?", run.ID).Update("target_id", targetID).Error) - seedTeamRunSession(t, db, run.SessionID, "user-1", executor) - - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "Implement replay", - Context: "include events", - Status: model.AssignmentStatusPending, - Depth: 1, - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - - require.NoError(t, svc.DispatchAssignment(context.Background(), "user-1", assignment.ID)) - - assert.Equal(t, "target-local-edge-1", agentSvc.targetID) - var reloadedAssignment model.AgentTeamAssignment - require.NoError(t, db.Where("id = ?", assignment.ID).First(&reloadedAssignment).Error) - require.NotNil(t, reloadedAssignment.RunID) - assert.Equal(t, "task-dispatch-1", *reloadedAssignment.RunID) -} - -func TestAgentTeamService_GetTeamRunStateProjectsDependenciesAndBudget(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - reviewerProfileID := "profile-reviewer" - reviewer := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &reviewerProfileID, - Role: model.TeamMemberRoleReviewer, - } - require.NoError(t, repository.AddTeamMember(db, reviewer)) - pending := &model.PendingAgentTask{ - AgentInstanceID: "agent-executor", - TriggeredByUserID: "user-1", - TriggerMessageID: "message-1", - Status: model.TaskStatusRunning, - EdgeRunID: "edge-run-budget", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(pending).Error) - - root := &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Root task", - RunID: &pending.ID, - } - require.NoError(t, repository.CreateTeamTask(db, root)) - child := &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - ParentTaskID: &root.ID, - Status: model.TeamTaskStatusPending, - Objective: "Child task", - } - require.NoError(t, repository.CreateTeamTask(db, child)) - conflictingPending := &model.PendingAgentTask{ - AgentInstanceID: "agent-reviewer", - TriggeredByUserID: "user-1", - TriggerMessageID: "message-2", - Status: model.TaskStatusDone, - EdgeRunID: "edge-run-reviewer", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(conflictingPending).Error) - conflictingTask := &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: reviewer.ID, - Status: model.TeamTaskStatusDone, - Objective: "Review same file", - RunID: &conflictingPending.ID, - } - require.NoError(t, repository.CreateTeamTask(db, conflictingTask)) - - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: "edge-run-budget", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.result", - Payload: `{"success":true,"usage":{"input_tokens":1200,"output_tokens":800},"tokenLimit":200000,"tokensRemaining":198000,"usagePercent":1.0}`, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: "edge-run-budget", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.context_warning", - Payload: `{"usagePercent":86.5,"threshold":85}`, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: "edge-run-budget", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-1","toolUseId":"tool-1","toolName":"Bash","status":"pending"}`, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: "edge-run-budget", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_decided", - Payload: `{"requestId":"req-1","decision":"allow","reason":"safe command"}`, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: "edge-run-budget", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.file_change", - Payload: `{"path":"hub-server/internal/service/agent_team.go","action":"modified","toolName":"apply_patch","status":"completed"}`, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: conflictingPending.ID, - EdgeRunID: "edge-run-reviewer", - SessionID: run.SessionID, - AgentInstanceID: "agent-reviewer", - EventType: "run.agent.file_change", - Payload: `{"path":"./hub-server/internal/service/agent_team.go","action":"modified","toolName":"review_patch","status":"completed"}`, - })) - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.Dependencies, 1) - assert.Equal(t, child.ID, state.Dependencies[0].TaskID) - assert.Equal(t, root.ID, state.Dependencies[0].DependsOnTaskID) - assert.Equal(t, "parent_task", state.Dependencies[0].Kind) - require.NotNil(t, state.Budget) - assert.Equal(t, int64(2000), state.Budget.TotalTokensUsed) - assert.Equal(t, int64(1200), state.Budget.InputTokens) - assert.Equal(t, int64(800), state.Budget.OutputTokens) - assert.Equal(t, int64(200000), state.Budget.TokenLimit) - assert.Equal(t, int64(198000), state.Budget.RemainingTokens) - assert.Equal(t, 86.5, state.Budget.UsagePercent) - assert.Equal(t, 2, state.Budget.RunCount) - assert.Equal(t, 1, state.Budget.ContextWarnings) - require.Len(t, state.Approvals, 1) - assert.Equal(t, pending.ID, state.Approvals[0].AgentTaskID) - assert.Equal(t, root.ID, state.Approvals[0].TeamTaskID) - assert.Equal(t, executor.ID, state.Approvals[0].MemberID) - assert.Equal(t, "req-1", state.Approvals[0].RequestID) - assert.Equal(t, "Bash", state.Approvals[0].ToolName) - assert.Equal(t, "tool-1", state.Approvals[0].ToolUseID) - assert.Equal(t, "allow", state.Approvals[0].Status) - assert.Equal(t, "safe command", state.Approvals[0].Reason) - require.NotNil(t, state.Approvals[0].DecidedAt) - require.Len(t, state.Artifacts, 2) - assert.Equal(t, pending.ID, state.Artifacts[0].AgentTaskID) - assert.Equal(t, root.ID, state.Artifacts[0].TeamTaskID) - assert.Equal(t, executor.ID, state.Artifacts[0].MemberID) - assert.Equal(t, "hub-server/internal/service/agent_team.go", state.Artifacts[0].Path) - assert.Equal(t, "modified", state.Artifacts[0].Action) - assert.Equal(t, "apply_patch", state.Artifacts[0].ToolName) - assert.Equal(t, "completed", state.Artifacts[0].Status) - assert.Equal(t, state.Artifacts[0].ConflictID, state.Artifacts[1].ConflictID) - require.Len(t, state.Conflicts, 1) - assert.Equal(t, "hub-server/internal/service/agent_team.go", state.Conflicts[0].Path) - assert.Equal(t, "pending", state.Conflicts[0].Status) - assert.ElementsMatch(t, []string{pending.ID, conflictingPending.ID}, state.Conflicts[0].AgentTaskIDs) - assert.ElementsMatch(t, []string{root.ID, conflictingTask.ID}, state.Conflicts[0].TeamTaskIDs) - assert.ElementsMatch(t, []string{executor.ID, reviewer.ID}, state.Conflicts[0].MemberIDs) - assert.ElementsMatch(t, []string{"modified"}, state.Conflicts[0].Actions) - - indexed, err := repository.ListTeamArtifactsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, indexed, 2) - assert.Equal(t, run.ID, indexed[0].TeamRunID) - require.NotNil(t, indexed[0].TeamTaskID) - assert.Equal(t, root.ID, *indexed[0].TeamTaskID) - require.NotNil(t, indexed[0].MemberID) - assert.Equal(t, executor.ID, *indexed[0].MemberID) - require.NotNil(t, indexed[0].AgentTaskID) - assert.Equal(t, pending.ID, *indexed[0].AgentTaskID) - require.NotNil(t, indexed[0].SourceEventID) - assert.NotEmpty(t, *indexed[0].SourceEventID) - assert.Equal(t, "apply_patch", indexed[0].ToolName) - assert.Equal(t, "hub-server/internal/service/agent_team.go", indexed[0].NormalizedPath) - assert.Equal(t, state.Artifacts[0].ConflictID, indexed[0].ConflictID) -} - -func TestAgentTeamService_ResolveConflictAppendsEventAndUpdatesReplay(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - reviewerProfileID := "profile-reviewer" - reviewer := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &reviewerProfileID, - Role: model.TeamMemberRoleReviewer, - } - require.NoError(t, repository.AddTeamMember(db, reviewer)) - firstTaskID := "agent-task-one" - secondTaskID := "agent-task-two" - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusDone, - Objective: "Change shared file", - RunID: &firstTaskID, - })) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: reviewer.ID, - Status: model.TeamTaskStatusDone, - Objective: "Review shared file", - RunID: &secondTaskID, - })) - for _, event := range []model.AgentRunEvent{ - { - TaskID: firstTaskID, - EdgeRunID: "edge-one", - SessionID: run.SessionID, - AgentInstanceID: "agent-one", - EventType: "run.agent.file_change", - Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, - }, - { - TaskID: secondTaskID, - EdgeRunID: "edge-two", - SessionID: run.SessionID, - AgentInstanceID: "agent-two", - EventType: "run.agent.file_change", - Payload: `{"path":"./shared.txt","action":"modified","status":"completed"}`, - }, - } { - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &event)) - } - - conflictID := conflictIDForPath("shared.txt") - resolved, err := svc.ResolveConflict(context.Background(), "user-1", team.ID, run.ID, model.TeamConflictResolution{ - ConflictID: conflictID, - Resolution: model.TeamConflictResolutionAcceptAgentTask, - SelectedAgentTaskID: firstTaskID, - Reason: "Use executor result", - }) - require.NoError(t, err) - require.NotNil(t, resolved) - assert.Equal(t, model.TeamConflictStatusResolved, resolved.Status) - assert.Equal(t, model.TeamConflictResolutionAcceptAgentTask, resolved.Resolution) - assert.Equal(t, firstTaskID, resolved.SelectedTask) - assert.Equal(t, "user-1", resolved.ResolvedBy) - require.NotNil(t, resolved.ResolvedAt) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventConflictResolved, events[0].Type) - assert.Contains(t, events[0].Payload, firstTaskID) - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.Conflicts, 1) - assert.Equal(t, model.TeamConflictStatusResolved, state.Conflicts[0].Status) - assert.Equal(t, model.TeamConflictResolutionAcceptAgentTask, state.Conflicts[0].Resolution) - assert.Equal(t, firstTaskID, state.Conflicts[0].SelectedTask) -} - -func TestAgentTeamService_ResolveConflictRejectsTaskOutsideConflict(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - reviewerProfileID := "profile-reviewer" - reviewer := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &reviewerProfileID, - Role: model.TeamMemberRoleReviewer, - } - require.NoError(t, repository.AddTeamMember(db, reviewer)) - taskID := "agent-task-one" - otherTaskID := "agent-task-two" - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusDone, - Objective: "Change shared file", - RunID: &taskID, - })) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: reviewer.ID, - Status: model.TeamTaskStatusDone, - Objective: "Review shared file", - RunID: &otherTaskID, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: taskID, - EdgeRunID: "edge-one", - SessionID: run.SessionID, - AgentInstanceID: "agent-one", - EventType: "run.agent.file_change", - Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: otherTaskID, - EdgeRunID: "edge-two", - SessionID: run.SessionID, - AgentInstanceID: "agent-two", - EventType: "run.agent.file_change", - Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, - })) - - resolved, err := svc.ResolveConflict(context.Background(), "user-1", team.ID, run.ID, model.TeamConflictResolution{ - ConflictID: conflictIDForPath("shared.txt"), - Resolution: model.TeamConflictResolutionAcceptAgentTask, - SelectedAgentTaskID: "missing-task", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, resolved) -} - -func TestAgentTeamService_DecideApprovalAppendsEventAndUpdatesReplay(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - pending := &model.PendingAgentTask{ - ID: "agent-task-approval", - AgentInstanceID: "agent-executor", - TriggeredByUserID: "user-1", - TriggerMessageID: "msg-approval", - Status: model.TaskStatusRunning, - EdgeRunID: "edge-run-approval", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(pending).Error) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Run gated command", - RunID: &pending.ID, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: pending.EdgeRunID, - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-approval","toolUseId":"tool-approval","toolName":"Bash","status":"pending"}`, - })) - - decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-approval", model.TeamApprovalDecision{ - Decision: "allow", - Reason: "Known safe command", - }) - require.NoError(t, err) - require.NotNil(t, decided) - assert.Equal(t, "req-approval", decided.ApprovalID) - assert.Equal(t, "allow", decided.Status) - assert.Equal(t, "user-1", decided.DecidedBy) - require.NotNil(t, decided.EdgeControl) - assert.Equal(t, pending.EdgeRunID, decided.EdgeControl.RunID) - assert.Equal(t, "req-approval", decided.EdgeControl.RequestID) - assert.Equal(t, "allow", decided.EdgeControl.Decision) - assert.Equal(t, "Known safe command", decided.EdgeControl.Reason) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventApprovalDecided, events[0].Type) - assert.Contains(t, events[0].Payload, `"edge_control"`) - assert.Contains(t, events[0].Payload, `"runId":"edge-run-approval"`) - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.Approvals, 1) - assert.Equal(t, "req-approval", state.Approvals[0].ApprovalID) - assert.Equal(t, "allow", state.Approvals[0].Status) - assert.Equal(t, "Known safe command", state.Approvals[0].Reason) - assert.Equal(t, "user-1", state.Approvals[0].DecidedBy) - require.NotNil(t, state.Approvals[0].DecidedAt) - require.NotNil(t, state.Approvals[0].EdgeControl) - assert.Equal(t, "edge-run-approval", state.Approvals[0].EdgeControl.RunID) -} - -func TestAgentTeamService_DecideApprovalDeliversControlToExactEdgeDevice(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - controlSvc := &mockAgentTeamControlSvc{} - svc := NewAgentTeamService(db, nil, nil) - svc.SetControlService(controlSvc) - team, _, executor, run := seedAgentTeamRun(t, db) - pending := &model.PendingAgentTask{ - ID: "agent-task-control", - AgentInstanceID: "agent-executor", - TriggeredByUserID: "user-1", - TriggerMessageID: "msg-control", - TargetID: "target-local", - Status: model.TaskStatusRunning, - EdgeRunID: "edge-run-control", - EdgeDeviceID: "edge-device-control", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(pending).Error) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Run gated command", - RunID: &pending.ID, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: pending.EdgeRunID, - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-control","toolUseId":"tool-control","toolName":"Bash","status":"pending"}`, - })) - - _, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-control", model.TeamApprovalDecision{ - Decision: "allow", - Reason: "Known safe command", - }) - require.NoError(t, err) - require.Len(t, controlSvc.calls, 1) - call := controlSvc.calls[0] - assert.Equal(t, "user-1", call.userID) - assert.Equal(t, "edge-device-control", call.deviceID) - assert.Equal(t, model.AgentControlKindPermissionDecide, call.payload.Kind) - assert.Equal(t, pending.ID, call.payload.AgentTaskID) - assert.Equal(t, "target-local", call.payload.TargetID) - assert.Equal(t, "edge-device-control", call.payload.EdgeDeviceID) - assert.Equal(t, "req-control", call.payload.ApprovalID) - require.NotNil(t, call.payload.EdgeControl) - assert.Equal(t, "edge-run-control", call.payload.EdgeControl.RunID) - assert.Equal(t, "req-control", call.payload.EdgeControl.RequestID) - assert.Equal(t, "allow", call.payload.EdgeControl.Decision) -} - -func TestAgentTeamService_DecideApprovalRedeliversSameDecisionWithoutDuplicateEvent(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - controlSvc := &mockAgentTeamControlSvc{} - svc := NewAgentTeamService(db, nil, nil) - svc.SetControlService(controlSvc) - team, _, executor, run := seedAgentTeamRun(t, db) - pending := &model.PendingAgentTask{ - ID: "agent-task-redeliver-control", - AgentInstanceID: "agent-executor", - TriggeredByUserID: "user-1", - TriggerMessageID: "msg-redeliver-control", - TargetID: "target-local", - Status: model.TaskStatusRunning, - EdgeRunID: "edge-run-redeliver-control", - EdgeDeviceID: "edge-device-redeliver-control", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(pending).Error) - teamTask := &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Run gated command", - RunID: &pending.ID, - } - require.NoError(t, repository.CreateTeamTask(db, teamTask)) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: pending.EdgeRunID, - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-redeliver","toolUseId":"tool-redeliver","toolName":"Bash","status":"pending"}`, - })) - decidedAt := time.Now().UTC() - record := model.TeamApprovalDecision{ - ApprovalID: "req-redeliver", - AgentTaskID: pending.ID, - TeamTaskID: teamTask.ID, - MemberID: executor.ID, - EdgeRunID: pending.EdgeRunID, - RequestID: "req-redeliver", - ToolName: "Bash", - ToolUseID: "tool-redeliver", - Decision: "allow", - Reason: "Known safe command", - DecidedBy: "user-1", - DecidedAt: decidedAt, - EdgeControl: &model.TeamApprovalEdgeControl{ - RunID: pending.EdgeRunID, - RequestID: "req-redeliver", - Decision: "allow", - Reason: "Known safe command", - }, - } - payload, err := json.Marshal(record) - require.NoError(t, err) - require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ - TeamRunID: run.ID, - Type: model.TeamEventApprovalDecided, - Payload: string(payload), - })) - - decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-redeliver", model.TeamApprovalDecision{ - Decision: "allow", - Reason: "Known safe command", - }) - - require.NoError(t, err) - require.NotNil(t, decided) - assert.Equal(t, "allow", decided.Status) - require.Len(t, controlSvc.calls, 1) - call := controlSvc.calls[0] - assert.Equal(t, "user-1", call.userID) - assert.Equal(t, pending.EdgeDeviceID, call.deviceID) - assert.Equal(t, pending.ID, call.payload.AgentTaskID) - assert.Equal(t, "req-redeliver", call.payload.ApprovalID) - require.NotNil(t, call.payload.EdgeControl) - assert.Equal(t, pending.EdgeRunID, call.payload.EdgeControl.RunID) - assert.Equal(t, "allow", call.payload.EdgeControl.Decision) - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) -} - -func TestAgentTeamService_DecideApprovalRejectsMissingEdgeDeviceForControl(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - controlSvc := &mockAgentTeamControlSvc{} - svc := NewAgentTeamService(db, nil, nil) - svc.SetControlService(controlSvc) - team, _, executor, run := seedAgentTeamRun(t, db) - pending := &model.PendingAgentTask{ - ID: "agent-task-no-device", - AgentInstanceID: "agent-executor", - TriggeredByUserID: "user-1", - TriggerMessageID: "msg-no-device", - Status: model.TaskStatusRunning, - EdgeRunID: "edge-run-no-device", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(pending).Error) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Run gated command", - RunID: &pending.ID, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: pending.EdgeRunID, - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-no-device","toolUseId":"tool-no-device","toolName":"Bash","status":"pending"}`, - })) - - decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-no-device", model.TeamApprovalDecision{ - Decision: "allow", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, decided) - assert.Empty(t, controlSvc.calls) -} - -func TestAgentTeamService_DecideApprovalRejectsAlreadyDecided(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - taskID := "agent-task-decided" - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Run gated command", - RunID: &taskID, - })) - for _, event := range []model.AgentRunEvent{ - { - TaskID: taskID, - EdgeRunID: "edge-run-decided", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-decided","toolUseId":"tool-decided","toolName":"Bash","status":"pending"}`, - }, - { - TaskID: taskID, - EdgeRunID: "edge-run-decided", - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_decided", - Payload: `{"requestId":"req-decided","decision":"deny","reason":"too broad"}`, - }, - } { - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &event)) - } - - decided, err := svc.DecideApproval(context.Background(), "user-1", team.ID, run.ID, "req-decided", model.TeamApprovalDecision{ - Decision: "allow", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - assert.Nil(t, decided) -} - -func TestAgentTeamService_MemberReadableTeamCannotMutateRunDecisions(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - addReadableTeamMemberForUser(t, db, team.ID, "member-user") - - state, err := svc.GetTeamRunState(context.Background(), "member-user", team.ID, run.ID) - require.NoError(t, err) - require.NotNil(t, state) - - reviewerProfileID := "profile-reviewer" - reviewer := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &reviewerProfileID, - Role: model.TeamMemberRoleReviewer, - } - require.NoError(t, repository.AddTeamMember(db, reviewer)) - firstTaskID := "agent-task-conflict-one" - secondTaskID := "agent-task-conflict-two" - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusDone, - Objective: "Change shared file", - RunID: &firstTaskID, - })) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: reviewer.ID, - Status: model.TeamTaskStatusDone, - Objective: "Review shared file", - RunID: &secondTaskID, - })) - for _, event := range []model.AgentRunEvent{ - { - TaskID: firstTaskID, - EdgeRunID: "edge-conflict-one", - SessionID: run.SessionID, - AgentInstanceID: "agent-one", - EventType: "run.agent.file_change", - Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, - }, - { - TaskID: secondTaskID, - EdgeRunID: "edge-conflict-two", - SessionID: run.SessionID, - AgentInstanceID: "agent-two", - EventType: "run.agent.file_change", - Payload: `{"path":"shared.txt","action":"modified","status":"completed"}`, - }, - } { - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &event)) - } - - resolved, err := svc.ResolveConflict(context.Background(), "member-user", team.ID, run.ID, model.TeamConflictResolution{ - ConflictID: conflictIDForPath("shared.txt"), - Resolution: model.TeamConflictResolutionAcceptAgentTask, - SelectedAgentTaskID: firstTaskID, - }) - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) - assert.Nil(t, resolved) - - pending := &model.PendingAgentTask{ - ID: "agent-task-member-approval", - AgentInstanceID: "agent-executor", - TriggeredByUserID: "user-1", - TriggerMessageID: "msg-member-approval", - Status: model.TaskStatusRunning, - EdgeRunID: "edge-run-member-approval", - ExpireAt: time.Now().Add(time.Hour), - } - require.NoError(t, db.Create(pending).Error) - require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ - TeamRunID: run.ID, - AssigneeMemberID: executor.ID, - Status: model.TeamTaskStatusRunning, - Objective: "Run gated command", - RunID: &pending.ID, - })) - require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ - TaskID: pending.ID, - EdgeRunID: pending.EdgeRunID, - SessionID: run.SessionID, - AgentInstanceID: "agent-executor", - EventType: "run.agent.permission_requested", - Payload: `{"requestId":"req-member-approval","toolUseId":"tool-member-approval","toolName":"Bash","status":"pending"}`, - })) - - decided, err := svc.DecideApproval(context.Background(), "member-user", team.ID, run.ID, "req-member-approval", model.TeamApprovalDecision{ - Decision: "allow", - }) - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) - assert.Nil(t, decided) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - assert.Empty(t, events) -} - -func TestAgentTeamService_ListTeamEventsIsOwnerScoped(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, _, run := seedAgentTeamRun(t, db) - require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ - TeamRunID: run.ID, - Type: model.TeamEventRouteRejected, - Payload: `{"reason":"invalid action"}`, - })) - - page, err := svc.ListTeamEvents(context.Background(), "user-1", team.ID, run.ID, 0, 50) - require.NoError(t, err) - require.Len(t, page.Items, 1) - assert.Equal(t, model.TeamEventRouteRejected, page.Items[0].Type) - assert.False(t, page.HasMore) - assert.Equal(t, 1, page.NextSeq) - - _, err = svc.ListTeamEvents(context.Background(), "other-user", team.ID, run.ID, 0, 50) - require.Error(t, err) - assert.Equal(t, errcode.AgentNotFound, err) -} - -func readAgentTeamEvent(t *testing.T, events <-chan bus.Event) bus.Event { - t.Helper() - select { - case event := <-events: - return event - case <-time.After(time.Second): - t.Fatal("agent team event was not published") - } - return bus.Event{} -} - -func setupAgentTeamStateSQLite(t *testing.T) *gorm.DB { - t.Helper() - return setupAgentTeamStateSQLiteDSN(t, ":memory:", 1) -} - -func setupAgentTeamConcurrentSQLite(t *testing.T) *gorm.DB { - t.Helper() - path := filepath.ToSlash(filepath.Join(t.TempDir(), "agentteam-concurrency.db")) - dsn := fmt.Sprintf("file:%s?cache=shared&_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)", path) - return setupAgentTeamStateSQLiteDSN(t, dsn, 8) -} - -func setupAgentTeamStateSQLiteDSN(t *testing.T, dsn string, maxOpenConns int) *gorm.DB { - t.Helper() - db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) - require.NoError(t, err) - sqlDB, err := db.DB() - require.NoError(t, err) - t.Cleanup(func() { _ = sqlDB.Close() }) - sqlDB.SetMaxOpenConns(maxOpenConns) - tables := []string{ - `CREATE TABLE agent_teams ( - id TEXT PRIMARY KEY, - owner_id TEXT NOT NULL, - name TEXT NOT NULL, - description TEXT DEFAULT '', - avatar_url TEXT DEFAULT '', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_members ( - id TEXT PRIMARY KEY, - team_id TEXT NOT NULL, - agent_profile_id TEXT, - role TEXT NOT NULL DEFAULT 'executor', - position INTEGER NOT NULL DEFAULT 0, - created_at DATETIME - )`, - `CREATE TABLE agent_team_runs ( - id TEXT PRIMARY KEY, - team_id TEXT NOT NULL, - session_id TEXT, - trigger_user_id TEXT NOT NULL, - trigger_message TEXT DEFAULT '', - target_id TEXT, - mode TEXT NOT NULL DEFAULT 'supervisor', - status TEXT NOT NULL DEFAULT 'queued', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_assignments ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - from_member_id TEXT NOT NULL, - to_member_id TEXT NOT NULL, - type TEXT NOT NULL DEFAULT 'delegate', - task_prompt TEXT NOT NULL, - context TEXT DEFAULT '', - status TEXT NOT NULL DEFAULT 'pending', - run_id TEXT, - result TEXT DEFAULT '', - depth INTEGER NOT NULL DEFAULT 0, - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_tasks ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - assignment_id TEXT, - assignee_member_id TEXT NOT NULL, - parent_task_id TEXT, - status TEXT NOT NULL DEFAULT 'pending', - objective TEXT NOT NULL, - input_refs TEXT NOT NULL DEFAULT '{}', - run_id TEXT, - attempt INTEGER NOT NULL DEFAULT 1, - risk_level TEXT NOT NULL DEFAULT 'normal', - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE agent_team_events ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - seq INTEGER NOT NULL, - type TEXT NOT NULL, - payload TEXT NOT NULL DEFAULT '{}', - created_at DATETIME - )`, - `CREATE TABLE agent_team_artifacts ( - id TEXT PRIMARY KEY, - team_run_id TEXT NOT NULL, - team_task_id TEXT, - assignment_id TEXT, - member_id TEXT, - agent_task_id TEXT, - edge_run_id TEXT, - source_event_id TEXT, - event_seq INTEGER NOT NULL DEFAULT 0, - path TEXT NOT NULL, - normalized_path TEXT NOT NULL, - action TEXT, - tool_name TEXT, - status TEXT, - conflict_id TEXT, - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE sessions ( - id TEXT PRIMARY KEY, - type TEXT NOT NULL, - name TEXT DEFAULT '', - owner_user_id TEXT, - workspace_id TEXT, - next_seq INTEGER NOT NULL DEFAULT 0, - last_message_at DATETIME, - dissolved BOOLEAN NOT NULL DEFAULT FALSE, - created_at DATETIME - )`, - `CREATE TABLE session_members ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - member_type TEXT NOT NULL, - member_id TEXT NOT NULL, - role TEXT NOT NULL, - pinned BOOLEAN NOT NULL DEFAULT FALSE, - archived BOOLEAN NOT NULL DEFAULT FALSE, - muted BOOLEAN NOT NULL DEFAULT FALSE, - last_read_seq INTEGER NOT NULL DEFAULT 0, - joined_at DATETIME, - left_at DATETIME - )`, - `CREATE TABLE agent_instances ( - id TEXT PRIMARY KEY, - agent_type TEXT NOT NULL, - custom_agent_id TEXT, - session_id TEXT NOT NULL, - inviter_user_id TEXT NOT NULL, - workspace_id TEXT, - display_name TEXT NOT NULL, - created_at DATETIME - )`, - `CREATE TABLE custom_agents ( - id TEXT PRIMARY KEY, - owner_user_id TEXT NOT NULL, - name TEXT NOT NULL, - avatar_url TEXT DEFAULT '', - agent_type TEXT NOT NULL, - system_prompt TEXT DEFAULT '', - capability_tags TEXT DEFAULT '[]', - tool_whitelist TEXT DEFAULT '[]', - model_params TEXT DEFAULT '{}', - output_schema TEXT DEFAULT NULL, - deleted_at DATETIME, - created_at DATETIME, - updated_at DATETIME - )`, - `CREATE TABLE pending_agent_tasks ( - id TEXT PRIMARY KEY, - agent_instance_id TEXT NOT NULL, - triggered_by_user_id TEXT NOT NULL, - trigger_message_id TEXT NOT NULL, - target_id TEXT, - status TEXT NOT NULL, - edge_run_id TEXT DEFAULT '', - edge_device_id TEXT, - error_message TEXT, - model_params TEXT DEFAULT '{}', - created_at DATETIME, - dispatched_at DATETIME, - finished_at DATETIME, - expire_at DATETIME NOT NULL - )`, - `CREATE TABLE messages ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - seq_id INTEGER NOT NULL, - client_msg_id TEXT NOT NULL, - sender_type TEXT NOT NULL, - sender_id TEXT NOT NULL, - content_type TEXT NOT NULL, - content TEXT NOT NULL, - reply_to_message_id TEXT, - recalled BOOLEAN NOT NULL DEFAULT FALSE, - edited BOOLEAN NOT NULL DEFAULT FALSE, - edited_at DATETIME, - created_at DATETIME - )`, - `CREATE UNIQUE INDEX idx_messages_session_client_msg ON messages (session_id, client_msg_id)`, - `CREATE TABLE agent_run_events ( - id TEXT PRIMARY KEY, - task_id TEXT NOT NULL, - edge_run_id TEXT, - session_id TEXT NOT NULL, - agent_instance_id TEXT NOT NULL, - event_seq INTEGER NOT NULL, - event_type TEXT NOT NULL, - payload TEXT NOT NULL, - created_at DATETIME - )`, - } - for _, ddl := range tables { - require.NoError(t, db.Exec(ddl).Error) - } - return db -} - -func seedAgentTeamRun(t *testing.T, db *gorm.DB) (*model.AgentTeam, *model.AgentTeamMember, *model.AgentTeamMember, *model.AgentTeamRun) { - t.Helper() - team := &model.AgentTeam{OwnerID: "user-1", Name: "Route Team"} - require.NoError(t, repository.CreateTeam(db, team)) - - supervisorProfileID := "profile-supervisor" - executorProfileID := "profile-executor" - supervisor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &supervisorProfileID, - Role: model.TeamMemberRoleSupervisor, - } - executor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: &executorProfileID, - Role: model.TeamMemberRoleExecutor, - } - require.NoError(t, repository.AddTeamMember(db, supervisor)) - require.NoError(t, repository.AddTeamMember(db, executor)) - - run := &model.AgentTeamRun{ - TeamID: team.ID, - SessionID: "session-1", - TriggerUserID: "user-1", - TriggerMessage: "ship it", - Status: model.TeamRunStatusRunning, - } - require.NoError(t, repository.CreateTeamRun(db, run)) - return team, supervisor, executor, run -} - -func addReadableTeamMemberForUser(t *testing.T, db *gorm.DB, teamID, userID string) *model.AgentTeamMember { - t.Helper() - agent := &model.CustomAgent{ - OwnerUserID: userID, - Name: "Readable Member Agent", - AgentType: "codex", - SystemPrompt: "Read shared team state", - } - require.NoError(t, repository.CreateCustomAgent(db, agent)) - member := &model.AgentTeamMember{ - TeamID: teamID, - AgentProfileID: &agent.ID, - Role: model.TeamMemberRoleExecutor, - } - require.NoError(t, repository.AddTeamMember(db, member)) - return member -} - -func addTeamSupervisor(t *testing.T, db *gorm.DB, teamID, profileID string) *model.AgentTeamMember { - t.Helper() - member := &model.AgentTeamMember{ - TeamID: teamID, - AgentProfileID: &profileID, - Role: model.TeamMemberRoleSupervisor, - } - require.NoError(t, repository.AddTeamMember(db, member)) - return member -} - -func stringPtr(value string) *string { - return &value -} - -func mustJSON(t *testing.T, value any) string { - t.Helper() - data, err := json.Marshal(value) - require.NoError(t, err) - return string(data) -} - -func seedTeamRunSession(t *testing.T, db *gorm.DB, sessionID, userID string, executor *model.AgentTeamMember) { - t.Helper() - now := time.Now() - require.NoError(t, db.Exec( - `INSERT INTO sessions (id, type, name, owner_user_id, next_seq, dissolved, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, - sessionID, model.SessionTypeGroup, "Team session", userID, 0, false, now, - ).Error) - require.NoError(t, db.Exec( - `INSERT INTO session_members (id, session_id, member_type, member_id, role, joined_at) VALUES (?, ?, ?, ?, ?, ?)`, - "session-member-user", sessionID, model.MemberTypeUser, userID, model.MemberRoleOwner, now, - ).Error) - customAgentID := "" - if executor.AgentProfileID != nil { - customAgentID = *executor.AgentProfileID - } - require.NoError(t, db.Exec( - `INSERT INTO agent_instances (id, agent_type, custom_agent_id, session_id, inviter_user_id, display_name, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, - "agent-executor", "codex", customAgentID, sessionID, userID, "Executor", now, - ).Error) - require.NoError(t, db.Exec( - `INSERT INTO session_members (id, session_id, member_type, member_id, role, joined_at) VALUES (?, ?, ?, ?, ?, ?)`, - "session-member-agent", sessionID, model.MemberTypeAgent, "agent-executor", model.MemberRoleMember, now, - ).Error) -} - -// mockCompeteAggregator implements CompeteAggregator for tests. -type mockCompeteAggregator struct { - summary string -} - -func (m *mockCompeteAggregator) CompareResults(_ context.Context, _ string, entries []model.CompeteSummaryEntry) (string, error) { - if m.summary != "" { - return m.summary, nil - } - return "Comparison: " + strings.Join(entryMemberIDs(entries), " vs "), nil -} - -func entryMemberIDs(entries []model.CompeteSummaryEntry) []string { - ids := make([]string, len(entries)) - for i, e := range entries { - ids[i] = e.MemberID - } - return ids -} - -func TestCompeteMode(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - aggregator := &mockCompeteAggregator{} - svc := NewAgentTeamService(db, nil, nil) - svc.SetCompeteAggregator(aggregator) - - team := &model.AgentTeam{OwnerID: "user-1", Name: "Compete Team"} - require.NoError(t, repository.CreateTeam(db, team)) - - supervisor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-supervisor"), - Role: model.TeamMemberRoleSupervisor, - } - executor1 := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-executor-1"), - Role: model.TeamMemberRoleExecutor, - } - executor2 := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-executor-2"), - Role: model.TeamMemberRoleExecutor, - } - require.NoError(t, repository.AddTeamMember(db, supervisor)) - require.NoError(t, repository.AddTeamMember(db, executor1)) - require.NoError(t, repository.AddTeamMember(db, executor2)) - - run := &model.AgentTeamRun{ - TeamID: team.ID, - SessionID: "session-compete", - TriggerUserID: "user-1", - TriggerMessage: "compare implementations of factorial", - Mode: model.TeamRunModeCompete, - Status: model.TeamRunStatusRunning, - } - require.NoError(t, repository.CreateTeamRun(db, run)) - - // Submit a compete route decision targeting two executors. - workerIDs := executor1.ID + "," + executor2.ID - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "compete", - NextWorker: workerIDs, - Instructions: "Implement factorial in your best style", - Reasoning: "Let's see different approaches", - }) - require.NoError(t, err) - require.NotNil(t, assignment) - assert.Equal(t, model.AssignmentTypeCompete, assignment.Type) - - // Both executors should have assignments. - assignments, err := svc.ListAssignments(context.Background(), "user-1", run.ID) - require.NoError(t, err) - assert.Len(t, assignments, 2) - assert.Equal(t, model.AssignmentTypeCompete, assignments[0].Type) - assert.Equal(t, model.AssignmentTypeCompete, assignments[1].Type) - - // Complete both assignments with results. - as1 := &assignments[0] - as1.Status = model.AssignmentStatusRunning - require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", as1.ID).Update("status", model.AssignmentStatusRunning).Error) - require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", as1.ID, "func fact(n int) int { if n <= 1 { return 1 }; return n * fact(n-1) }")) - - as2 := &assignments[1] - as2.Status = model.AssignmentStatusRunning - require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", as2.ID).Update("status", model.AssignmentStatusRunning).Error) - require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", as2.ID, "def factorial(n): return 1 if n <= 1 else n * factorial(n-1)")) - - // Generate compete summary. - resp, err := svc.GenerateCompeteSummary(context.Background(), "user-1", run.ID, model.CompeteSummaryRequest{}) - require.NoError(t, err) - require.NotNil(t, resp) - assert.Equal(t, run.ID, resp.TeamRunID) - assert.Contains(t, resp.Summary, "Comparison:") - assert.Len(t, resp.Entries, 2) - - // Verify events include compete dispatched and aggregated. - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - hasCompeteDispatched := false - hasCompeteAggregated := false - for _, ev := range events { - if ev.Type == model.TeamEventCompeteDispatched { - hasCompeteDispatched = true - } - if ev.Type == model.TeamEventCompeteAggregated { - hasCompeteAggregated = true - } - } - assert.True(t, hasCompeteDispatched, "expected team.compete.dispatched event") - assert.True(t, hasCompeteAggregated, "expected team.compete.aggregated event") -} - -func TestCompeteModeAutoSelectsExecutors(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - aggregator := &mockCompeteAggregator{} - svc := NewAgentTeamService(db, nil, nil) - svc.SetCompeteAggregator(aggregator) - - team := &model.AgentTeam{OwnerID: "user-1", Name: "Auto Compete Team"} - require.NoError(t, repository.CreateTeam(db, team)) - - supervisor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-supervisor"), - Role: model.TeamMemberRoleSupervisor, - } - executor1 := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-executor-1"), - Role: model.TeamMemberRoleExecutor, - } - executor2 := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-executor-2"), - Role: model.TeamMemberRoleExecutor, - } - require.NoError(t, repository.AddTeamMember(db, supervisor)) - require.NoError(t, repository.AddTeamMember(db, executor1)) - require.NoError(t, repository.AddTeamMember(db, executor2)) - - run := &model.AgentTeamRun{ - TeamID: team.ID, - SessionID: "session-auto-compete", - TriggerUserID: "user-1", - TriggerMessage: "auto compete", - Mode: model.TeamRunModeCompete, - Status: model.TeamRunStatusRunning, - } - require.NoError(t, repository.CreateTeamRun(db, run)) - - // Submit compete decision with empty NextWorker — should pick both executors. - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "compete", - Instructions: "Write hello world", - Reasoning: "Auto-select", - }) - require.NoError(t, err) - require.NotNil(t, assignment) - - assignments, err := svc.ListAssignments(context.Background(), "user-1", run.ID) - require.NoError(t, err) - assert.Len(t, assignments, 2) - // Both should be executors (not supervisor). - for _, a := range assignments { - assert.NotEqual(t, supervisor.ID, a.ToMemberID) - } -} - -func TestCompeteModeRejectsExceedingMaxAgents(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetCompeteMaxAgents(2) - - team := &model.AgentTeam{OwnerID: "user-1", Name: "Max Compete Team"} - require.NoError(t, repository.CreateTeam(db, team)) - - supervisor := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr("profile-supervisor"), - Role: model.TeamMemberRoleSupervisor, - } - require.NoError(t, repository.AddTeamMember(db, supervisor)) - - workerIDs := make([]string, 3) - for i := range 3 { - e := &model.AgentTeamMember{ - TeamID: team.ID, - AgentProfileID: strPtr(fmt.Sprintf("profile-executor-%d", i)), - Role: model.TeamMemberRoleExecutor, - } - require.NoError(t, repository.AddTeamMember(db, e)) - workerIDs[i] = e.ID - } - - run := &model.AgentTeamRun{ - TeamID: team.ID, - SessionID: "session-max-compete", - TriggerUserID: "user-1", - TriggerMessage: "too many", - Mode: model.TeamRunModeCompete, - Status: model.TeamRunStatusRunning, - } - require.NoError(t, repository.CreateTeamRun(db, run)) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "compete", - NextWorker: strings.Join(workerIDs, ","), - Instructions: "Too many workers", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) -} - -func strPtr(s string) *string { - return &s -} - -// ── Human Review Gate (ADR-008) ──────────────────────────────── - -func TestHumanReviewGate(t *testing.T) { - t.Run("disabled by default rejects review API", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - - // HandleRouteDecision should proceed normally (no pending_review). - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - require.NotNil(t, assignment) - - // Verify the run is still in running state, not pending_review. - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) - - // Calling ReviewDagPlan when disabled should fail. - _, err = svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: model.ReviewActionApprove, - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - }) - - t.Run("enabled sets pending_review after route decision", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - require.NotNil(t, assignment) - - // Verify the run is now pending_review. - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusPendingReview, gotRun.Status) - - // Verify review_pending event was recorded. - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - foundPending := false - for _, e := range events { - if e.Type == model.TeamEventReviewPending { - foundPending = true - break - } - } - assert.True(t, foundPending, "expected team.review.pending event") - }) - - t.Run("enabled approve transitions back to running", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - - // Review: approve - state, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: model.ReviewActionApprove, - Comment: "looks good", - }) - require.NoError(t, err) - assert.Equal(t, model.ReviewActionApprove, state.Action) - assert.Equal(t, "looks good", state.Comment) - - // Verify the run is back to running. - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) - - // Verify review_decided event was recorded. - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - foundDecided := false - for _, e := range events { - if e.Type == model.TeamEventReviewDecided { - foundDecided = true - break - } - } - assert.True(t, foundDecided, "expected team.review.decided event") - - // The replay projection must leave the review gate too: before the - // ReviewDecided replay fix the client-facing GetTeamRunState stayed - // pending_review forever even though the DB row was already running. - runState, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusRunning, runState.Status) - require.Len(t, runState.Reviews, 1) - assert.Equal(t, model.ReviewActionApprove, runState.Reviews[0].Action) - }) - - t.Run("decided review restores running in projection for discuss and modify", func(t *testing.T) { - for _, action := range []string{model.ReviewActionDiscuss, model.ReviewActionModify} { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - - _, err = svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: action, - Comment: "not yet", - }) - require.NoError(t, err) - - // Write side sets the DB row back to running for every decided - // action; the projection must agree. - runState, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusRunning, runState.Status, "action=%s", action) - } - }) - - t.Run("enabled discuss cancels pending assignments", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - - // Verify assignment is pending before review. - assignments, err := repository.ListAssignmentsByTeamRun(db, run.ID) - require.NoError(t, err) - require.Len(t, assignments, 1) - assert.Equal(t, model.AssignmentStatusPending, assignments[0].Status) - - // Review: discuss - state, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: model.ReviewActionDiscuss, - Comment: "needs more thought", - }) - require.NoError(t, err) - assert.Equal(t, model.ReviewActionDiscuss, state.Action) - - // Verify the run is back to running (so supervisor can re-plan). - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) - - // Verify assignments were cancelled. - assignments, err = repository.ListAssignmentsByTeamRun(db, run.ID) - require.NoError(t, err) - require.Len(t, assignments, 1) - assert.Equal(t, model.AssignmentStatusCancelled, assignments[0].Status) - }) - - t.Run("enabled modify cancels pending assignments with changes recorded", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - - state, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: model.ReviewActionModify, - Changes: []model.HumanReviewChange{ - {Field: "instructions", Value: "Do the other thing"}, - {Field: "next_worker", Value: "worker-2"}, - }, - Comment: "change the target", - }) - require.NoError(t, err) - assert.Equal(t, model.ReviewActionModify, state.Action) - assert.Len(t, state.Changes, 2) - assert.Equal(t, "instructions", state.Changes[0].Field) - - // Verify the run is back to running. - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusRunning, gotRun.Status) - - // Verify assignments were cancelled. - assignments, err := repository.ListAssignmentsByTeamRun(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.AssignmentStatusCancelled, assignments[0].Status) - }) - - t.Run("review rejects invalid action", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - - _, err = svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: "bogus", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - }) - - t.Run("review rejects non-pending_review status", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - _, _, _, run := seedAgentTeamRun(t, db) // run status is "running" - - _, err := svc.ReviewDagPlan(context.Background(), "user-1", run.ID, model.HumanReviewDecision{ - Action: model.ReviewActionApprove, - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - }) - - t.Run("review rejects non-owner user", func(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - svc.SetHumanReviewEnabled(true) - team, _, executor, run := seedAgentTeamRun(t, db) - - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do the thing", - }) - require.NoError(t, err) - - _, err = svc.ReviewDagPlan(context.Background(), "intruder-user", run.ID, model.HumanReviewDecision{ - Action: model.ReviewActionApprove, - }) - require.Error(t, err) - assert.Equal(t, errcode.AgentTaskNotFound, err) - }) -} - -// ── State machine guards (terminal states & assignment reachability) ── - -// TestAgentTeamService_CompleteAssignmentAcceptsDispatchedViaDispatchPath drives -// an assignment through the real production path (pending -> dispatched -> -// running via DispatchAssignment) and verifies CompleteAssignment succeeds -// without raw SQL status injection. CompleteAssignment also still accepts a -// legacy dispatched row for #1376 compatibility. -func TestAgentTeamService_CompleteAssignmentAcceptsDispatchedViaDispatchPath(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - agentSvc := &mockAgentTeamAgentSvc{returnTaskID: "task-complete-dispatched-1"} - svc := NewAgentTeamService(db, agentSvc, nil) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - seedTeamRunSession(t, db, run.SessionID, "user-1", executor) - - assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor.ID, executor.ID, model.AssignmentTypeDelegate, "Ship the feature", "") - require.NoError(t, err) - - require.NoError(t, svc.DispatchAssignment(context.Background(), "user-1", assignment.ID)) - - var running model.AgentTeamAssignment - require.NoError(t, db.Where("id = ?", assignment.ID).First(&running).Error) - require.Equal(t, model.AssignmentStatusRunning, running.Status) - - require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", assignment.ID, "shipped")) - - var done model.AgentTeamAssignment - require.NoError(t, db.Where("id = ?", assignment.ID).First(&done).Error) - assert.Equal(t, model.AssignmentStatusDone, done.Status) - assert.Equal(t, "shipped", done.Result) - - // Legacy dispatched row (pre-#1384) must still complete. - legacy := &model.AgentTeamAssignment{ - TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, TaskPrompt: "legacy complete", Status: model.AssignmentStatusDispatched, - } - require.NoError(t, repository.CreateAssignment(db, legacy)) - require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", legacy.ID, "legacy shipped")) - var legacyDone model.AgentTeamAssignment - require.NoError(t, db.Where("id = ?", legacy.ID).First(&legacyDone).Error) - assert.Equal(t, model.AssignmentStatusDone, legacyDone.Status) -} - -// TestAgentTeamService_HandleRouteDecisionRejectsTerminalRun verifies that no -// route decision (delegate or a repeated finish) can mutate a run that already -// reached a terminal status, and that each rejection is audited. -func TestAgentTeamService_HandleRouteDecisionRejectsTerminalRun(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, _, executor, run := seedAgentTeamRun(t, db) - require.NoError(t, repository.UpdateTeamRunStatus(db, run.ID, model.TeamRunStatusCompleted)) - - // Delegate on a completed run must be rejected. - _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "delegate", - NextWorker: executor.ID, - Instructions: "Do more work", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - - // A repeated finish with blocked_reason must not downgrade completed to failed. - _, err = svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ - Action: "finish", - BlockedReason: "should not overwrite", - }) - require.Error(t, err) - assert.ErrorIs(t, err, errcode.ErrBadRequest) - - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusCompleted, gotRun.Status) - - assignments, err := repository.ListAssignmentsByTeamRun(db, run.ID) - require.NoError(t, err) - assert.Empty(t, assignments) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - rejected := 0 - for _, e := range events { - if e.Type == model.TeamEventRouteRejected { - rejected++ - } - } - assert.Equal(t, 2, rejected, "expected each terminal-run route decision to append team.route.rejected") -} - -// TestAgentTeamService_FinishRouteDecisionDoesNotDowngradeTerminalRun calls -// finishRouteDecision directly (bypassing the HandleRouteDecision entry guard, -// as a racing finish would) and verifies the conditional status write keeps -// the first terminal outcome and does not emit a conflicting run.failed event. -func TestAgentTeamService_FinishRouteDecisionDoesNotDowngradeTerminalRun(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - _, _, _, run := seedAgentTeamRun(t, db) - - require.NoError(t, svc.finishRouteDecision(run.ID, model.CoordinatorRouteDecision{ - Action: "finish", - Summary: "all done", - })) - gotRun, err := repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - require.Equal(t, model.TeamRunStatusCompleted, gotRun.Status) - - require.NoError(t, svc.finishRouteDecision(run.ID, model.CoordinatorRouteDecision{ - Action: "finish", - BlockedReason: "late failure", - })) - gotRun, err = repository.GetTeamRunByID(db, run.ID) - require.NoError(t, err) - assert.Equal(t, model.TeamRunStatusCompleted, gotRun.Status) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - completedEvents, failedEvents := 0, 0 - for _, e := range events { - switch e.Type { - case model.TeamEventRunCompleted: - completedEvents++ - case model.TeamEventRunFailed: - failedEvents++ - } - } - assert.Equal(t, 1, completedEvents) - assert.Zero(t, failedEvents, "repeated finish must not append team.run.failed") -} - -// TestAgentTeamService_FailAssignmentRejectsTerminalAssignment verifies that a -// done/failed/cancelled assignment cannot be failed again, so repeated fails -// can no longer rewrite the stored result or re-trigger fault escalation. -func TestAgentTeamService_FailAssignmentRejectsTerminalAssignment(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - - for _, status := range []string{ - model.AssignmentStatusDone, - model.AssignmentStatusFailed, - model.AssignmentStatusCancelled, - } { - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, - FromMemberID: supervisor.ID, - ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, - TaskPrompt: "Ship it", - Status: status, - Result: "original outcome", - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - - err := svc.FailAssignment(context.Background(), "user-1", assignment.ID, "[fault_escalation] retries=3 maxRetries=3 error=late") - require.Error(t, err, "status=%s", status) - assert.Equal(t, errcode.ErrBadRequest, err, "status=%s", status) - - var reloaded model.AgentTeamAssignment - require.NoError(t, db.Where("id = ?", assignment.ID).First(&reloaded).Error) - assert.Equal(t, status, reloaded.Status, "terminal status must be preserved") - assert.Equal(t, "original outcome", reloaded.Result) - } - - // The rejected fails must not have appended any team events (no repeated - // escalation side effects). - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - assert.Empty(t, events) -} - -// ── Assignment lifecycle (#1384) ────────────────────────────────────────────── - -func TestAgentTeamService_FailTimedOutAssignmentsTerminatesActiveAndIsIdempotent(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - guardrails := DefaultAgentTeamGuardrails() - guardrails.AssignmentTimeout = time.Minute - svc := NewAgentTeamServiceWithGuardrails(db, nil, nil, guardrails) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - - old := time.Now().Add(-2 * time.Hour) - fresh := time.Now() - timedOut := &model.AgentTeamAssignment{ - TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, TaskPrompt: "stale work", Status: model.AssignmentStatusRunning, - CreatedAt: old, UpdatedAt: old, - } - freshA := &model.AgentTeamAssignment{ - TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, TaskPrompt: "fresh work", Status: model.AssignmentStatusPending, - CreatedAt: fresh, UpdatedAt: fresh, - } - require.NoError(t, repository.CreateAssignment(db, timedOut)) - require.NoError(t, repository.CreateAssignment(db, freshA)) - // Force created_at past the guardrail (GORM autoCreateTime may overwrite on insert). - require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", timedOut.ID).Update("created_at", old).Error) - require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", freshA.ID).Update("created_at", fresh).Error) - - n, err := svc.FailTimedOutAssignments(context.Background()) - require.NoError(t, err) - assert.Equal(t, 1, n) - - reloaded, err := repository.GetAssignmentByID(db, timedOut.ID) - require.NoError(t, err) - assert.Equal(t, model.AssignmentStatusFailed, reloaded.Status) - assert.Equal(t, "assignment timeout reached", reloaded.Result) - - stillFresh, err := repository.GetAssignmentByID(db, freshA.ID) - require.NoError(t, err) - assert.Equal(t, model.AssignmentStatusPending, stillFresh.Status) - - events, err := repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - require.Len(t, events, 1) - assert.Equal(t, model.TeamEventAssignmentFailed, events[0].Type) - - // Second scan must not rewrite or re-emit. - n, err = svc.FailTimedOutAssignments(context.Background()) - require.NoError(t, err) - assert.Equal(t, 0, n) - events, err = repository.ListTeamEventsByRun(db, run.ID) - require.NoError(t, err) - assert.Len(t, events, 1) -} - -func TestAgentTeamService_CompleteAssignmentConcurrentOnlyOneWins(t *testing.T) { - db := setupAgentTeamConcurrentSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - _, supervisor, executor, run := seedAgentTeamRun(t, db) - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, TaskPrompt: "race complete", Status: model.AssignmentStatusRunning, - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - - errs := runConcurrentRouteCalls(2, func() error { - return svc.CompleteAssignment(context.Background(), "user-1", assignment.ID, "winner") - }) - successes, bad := 0, 0 - for _, err := range errs { - switch err { - case nil: - successes++ - case errcode.ErrBadRequest: - bad++ - default: - require.NoError(t, err) - } - } - assert.Equal(t, 1, successes) - assert.Equal(t, 1, bad) - - done, err := repository.GetAssignmentByID(db, assignment.ID) - require.NoError(t, err) - assert.Equal(t, model.AssignmentStatusDone, done.Status) - assert.Equal(t, "winner", done.Result) -} - -func TestAgentTeamService_GetTeamRunStateKeepsTerminalAssignmentOverPendingProjection(t *testing.T) { - db := setupAgentTeamStateSQLite(t) - svc := NewAgentTeamService(db, nil, nil) - team, supervisor, executor, run := seedAgentTeamRun(t, db) - pendingID := "pending-task-terminal-1" - assignment := &model.AgentTeamAssignment{ - TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, - Type: model.AssignmentTypeDelegate, TaskPrompt: "already failed", Status: model.AssignmentStatusFailed, - Result: "assignment timeout reached", RunID: &pendingID, - } - require.NoError(t, repository.CreateAssignment(db, assignment)) - require.NoError(t, db.Exec( - "INSERT INTO pending_agent_tasks (id, agent_instance_id, trigger_message_id, triggered_by_user_id, status, expire_at) VALUES (?, ?, ?, ?, ?, ?)", - pendingID, "agent-executor", "msg-1", "user-1", model.TaskStatusRunning, time.Now().Add(time.Hour), - ).Error) - - state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) - require.NoError(t, err) - require.Len(t, state.Assignments, 1) - assert.Equal(t, model.AssignmentStatusFailed, state.Assignments[0].Status) -} - -func TestIsTerminalTeamRunStatusIncludesCancelledWithoutCancelAPI(t *testing.T) { - // Documented dead write-path: cancelled is a terminal guard token only. - assert.True(t, isTerminalTeamRunStatus(model.TeamRunStatusCancelled)) - assert.True(t, isTerminalTeamRunStatus(model.TeamRunStatusCompleted)) - assert.True(t, isTerminalTeamRunStatus(model.TeamRunStatusFailed)) - assert.False(t, isTerminalTeamRunStatus(model.TeamRunStatusRunning)) -} diff --git a/hub-server/internal/service/agentteam/assignment_lifecycle_test.go b/hub-server/internal/service/agentteam/assignment_lifecycle_test.go new file mode 100644 index 000000000..ded744c84 --- /dev/null +++ b/hub-server/internal/service/agentteam/assignment_lifecycle_test.go @@ -0,0 +1,528 @@ +package agentteam + +import ( + "context" + "testing" + "time" + + "github.com/agenthub/hub-server/internal/bus" + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAgentTeamService_CompleteAssignmentPublishesEvent(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + eventBus, err := bus.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { eventBus.Close(context.Background()) }) + events := make(chan bus.Event, 1) + eventBus.Subscribe(bus.EventTypeTeamAssignmentDone, func(ctx context.Context, event bus.Event) { + events <- event + }) + svc.SetBus(eventBus) + + _, supervisor, executor, run := seedAgentTeamRun(t, db) + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "Ship it", + Status: model.AssignmentStatusRunning, + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + + require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", assignment.ID, "done text")) + + event := readAgentTeamEvent(t, events) + assert.Equal(t, bus.EventTypeTeamAssignmentDone, event.Type) + payload, ok := event.Payload.(map[string]interface{}) + require.True(t, ok) + assert.Equal(t, run.ID, payload["team_run_id"]) + assert.Equal(t, assignment.ID, payload["assignment_id"]) + assert.Equal(t, run.SessionID, payload["session_id"]) + assert.Equal(t, "done text", payload["result"]) +} + +func TestAgentTeamService_FailAssignmentPublishesEvent(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + eventBus, err := bus.New() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { eventBus.Close(context.Background()) }) + events := make(chan bus.Event, 1) + eventBus.Subscribe("team.assignment.failed", func(ctx context.Context, event bus.Event) { + events <- event + }) + svc.SetBus(eventBus) + + _, supervisor, executor, run := seedAgentTeamRun(t, db) + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "Ship it", + Status: model.AssignmentStatusRunning, + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + + require.NoError(t, svc.FailAssignment(context.Background(), "user-1", assignment.ID, "blocked")) + + event := readAgentTeamEvent(t, events) + assert.Equal(t, "team.assignment.failed", event.Type) + payload, ok := event.Payload.(map[string]interface{}) + require.True(t, ok) + assert.Equal(t, run.ID, payload["team_run_id"]) + assert.Equal(t, assignment.ID, payload["assignment_id"]) + assert.Equal(t, run.SessionID, payload["session_id"]) + assert.Equal(t, "blocked", payload["reason"]) +} + +func TestAgentTeamService_CreateAssignmentRejectsDelegationDepthLimit(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + _, supervisor, _, run := seedAgentTeamRun(t, db) + supervisor2 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-2") + supervisor3 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-3") + supervisor4 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-4") + supervisor5 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-5") + + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: supervisor2.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "depth 1", + Status: model.AssignmentStatusDone, + Depth: 1, + })) + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor2.ID, + ToMemberID: supervisor3.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "depth 2", + Status: model.AssignmentStatusDone, + Depth: 2, + })) + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor3.ID, + ToMemberID: supervisor4.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "depth 3", + Status: model.AssignmentStatusDone, + Depth: model.MaxDelegationDepth, + })) + + assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor4.ID, supervisor5.ID, model.AssignmentTypeDelegate, "too deep", "") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) +} + +func TestAgentTeamService_CreateAssignmentRejectsTeamRunTaskLimit(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + + for i := 0; i < model.MaxTasksPerTeamRun; i++ { + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "completed task", + Status: model.AssignmentStatusDone, + Depth: 1, + })) + } + + assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor.ID, executor.ID, model.AssignmentTypeDelegate, "too many", "") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) +} + +func TestAgentTeamService_CreateAssignmentRejectsDelegationCycle(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + _, supervisor, _, run := seedAgentTeamRun(t, db) + supervisor2 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-cycle-2") + supervisor3 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-supervisor-cycle-3") + + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: supervisor2.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "cycle depth 1", + Status: model.AssignmentStatusDone, + Depth: 1, + })) + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor2.ID, + ToMemberID: supervisor3.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "cycle depth 2", + Status: model.AssignmentStatusDone, + Depth: 2, + })) + + assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor3.ID, supervisor.ID, model.AssignmentTypeDelegate, "cycle", "") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) +} + +func TestAgentTeamService_CreateAssignmentUsesConfiguredDelegationDepthLimit(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamServiceWithGuardrails(db, nil, nil, AgentTeamGuardrails{ + MaxDelegationDepth: 1, + }) + _, supervisor, _, run := seedAgentTeamRun(t, db) + supervisor2 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-config-depth-2") + supervisor3 := addTeamSupervisor(t, db, supervisor.TeamID, "profile-config-depth-3") + + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: supervisor2.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "depth 1", + Status: model.AssignmentStatusDone, + Depth: 1, + })) + + assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor2.ID, supervisor3.ID, model.AssignmentTypeDelegate, "too deep for config", "") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) +} + +func TestAgentTeamService_CreateAssignmentUsesConfiguredTeamRunTaskLimit(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamServiceWithGuardrails(db, nil, nil, AgentTeamGuardrails{ + MaxTasksPerTeamRun: 1, + }) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "first task", + Status: model.AssignmentStatusDone, + Depth: 1, + })) + + assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor.ID, executor.ID, model.AssignmentTypeDelegate, "over configured task limit", "") + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) +} + +func TestAgentTeamService_DispatchAssignmentBindsTeamTaskToPendingAgentTask(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + agentSvc := &mockAgentTeamAgentSvc{returnTaskID: "task-dispatch-1"} + svc := NewAgentTeamService(db, agentSvc, nil) + team, supervisor, executor, run := seedAgentTeamRun(t, db) + seedTeamRunSession(t, db, run.SessionID, "user-1", executor) + + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "Implement replay", + Context: "include events", + Status: model.AssignmentStatusPending, + Depth: 1, + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + assignmentID := assignment.ID + teamTask := &model.AgentTeamTask{ + TeamRunID: run.ID, + AssignmentID: &assignmentID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusPending, + Objective: assignment.TaskPrompt, + } + require.NoError(t, repository.CreateTeamTask(db, teamTask)) + + require.NoError(t, svc.DispatchAssignment(context.Background(), "user-1", assignment.ID)) + + var reloadedAssignment model.AgentTeamAssignment + require.NoError(t, db.Where("id = ?", assignment.ID).First(&reloadedAssignment).Error) + require.NotNil(t, reloadedAssignment.RunID) + assert.Equal(t, model.AssignmentStatusRunning, reloadedAssignment.Status) + + // Seed the pending task that the mock agent service would have created. + triggerMsgID := "msg-trigger-dispatch" + require.NoError(t, db.Exec( + "INSERT INTO pending_agent_tasks (id, agent_instance_id, trigger_message_id, triggered_by_user_id, status, expire_at) VALUES (?, ?, ?, ?, ?, ?)", + *reloadedAssignment.RunID, "agent-executor", triggerMsgID, "user-1", model.TaskStatusRunning, time.Now().Add(1*time.Hour), + ).Error) + require.NoError(t, db.Exec( + "INSERT INTO messages (id, session_id, seq_id, client_msg_id, sender_type, sender_id, content_type, content, created_at) VALUES (?, ?, 1, ?, 'user', ?, 'text', ?, ?)", + triggerMsgID, run.SessionID, triggerMsgID, "user-1", "Task: Implement replay\nContext: include events", time.Now(), + ).Error) + + var reloadedTask model.AgentTeamTask + require.NoError(t, db.Where("id = ?", teamTask.ID).First(&reloadedTask).Error) + require.NotNil(t, reloadedTask.RunID) + assert.Equal(t, *reloadedAssignment.RunID, *reloadedTask.RunID) + assert.Equal(t, model.TeamTaskStatusDispatched, reloadedTask.Status) + + var pending model.PendingAgentTask + require.NoError(t, db.Where("id = ?", *reloadedTask.RunID).First(&pending).Error) + assert.Equal(t, "agent-executor", pending.AgentInstanceID) + assert.NotEmpty(t, pending.TriggerMessageID) + + var triggerMessage model.Message + require.NoError(t, db.Where("id = ?", pending.TriggerMessageID).First(&triggerMessage).Error) + assert.Contains(t, triggerMessage.Content, "Implement replay") + assert.Contains(t, triggerMessage.Content, "include events") + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventAssignmentDispatched, events[0].Type) + + require.NoError(t, repository.UpdatePendingTaskStatusWithEdgeRunID(db, pending.ID, model.TaskStatusRunning, "", "edge-run-1")) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: pending.ID, + EdgeRunID: "edge-run-1", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: model.RunEventTypeOutputBatch, + Payload: `{"content":"runtime output"}`, + })) + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.Assignments, 1) + assert.Equal(t, model.AssignmentStatusRunning, state.Assignments[0].Status) + assert.Equal(t, pending.ID, state.Assignments[0].AgentTaskID) + assert.Equal(t, "edge-run-1", state.Assignments[0].EdgeRunID) + require.Len(t, state.Tasks, 1) + assert.Equal(t, model.TeamTaskStatusRunning, state.Tasks[0].Status) + assert.Equal(t, pending.ID, state.Tasks[0].AgentTaskID) + assert.Equal(t, "edge-run-1", state.Tasks[0].EdgeRunID) + require.Len(t, state.RunEvents, 1) + assert.Equal(t, pending.ID, state.RunEvents[0].AgentTaskID) + assert.Equal(t, "edge-run-1", state.RunEvents[0].EdgeRunID) + assert.Equal(t, model.RunEventTypeOutputBatch, state.RunEvents[0].EventType) + assert.JSONEq(t, `{"content":"runtime output"}`, state.RunEvents[0].Payload) +} + +func TestAgentTeamService_DispatchAssignmentPassesTeamRunTargetID(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + agentSvc := &mockAgentTeamAgentSvc{returnTaskID: "task-dispatch-1"} + svc := NewAgentTeamService(db, agentSvc, nil) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + targetID := "target-local-edge-1" + run.TargetID = &targetID + require.NoError(t, db.Model(&model.AgentTeamRun{}).Where("id = ?", run.ID).Update("target_id", targetID).Error) + seedTeamRunSession(t, db, run.SessionID, "user-1", executor) + + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "Implement replay", + Context: "include events", + Status: model.AssignmentStatusPending, + Depth: 1, + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + + require.NoError(t, svc.DispatchAssignment(context.Background(), "user-1", assignment.ID)) + + assert.Equal(t, "target-local-edge-1", agentSvc.targetID) + var reloadedAssignment model.AgentTeamAssignment + require.NoError(t, db.Where("id = ?", assignment.ID).First(&reloadedAssignment).Error) + require.NotNil(t, reloadedAssignment.RunID) + assert.Equal(t, "task-dispatch-1", *reloadedAssignment.RunID) +} + +// ── State machine guards (terminal states & assignment reachability) ── + +// TestAgentTeamService_CompleteAssignmentAcceptsDispatchedViaDispatchPath drives +// an assignment through the real production path (pending -> dispatched -> +// running via DispatchAssignment) and verifies CompleteAssignment succeeds +// without raw SQL status injection. CompleteAssignment also still accepts a +// legacy dispatched row for #1376 compatibility. +func TestAgentTeamService_CompleteAssignmentAcceptsDispatchedViaDispatchPath(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + agentSvc := &mockAgentTeamAgentSvc{returnTaskID: "task-complete-dispatched-1"} + svc := NewAgentTeamService(db, agentSvc, nil) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + seedTeamRunSession(t, db, run.SessionID, "user-1", executor) + + assignment, err := svc.CreateAssignment(context.Background(), "user-1", run.ID, supervisor.ID, executor.ID, model.AssignmentTypeDelegate, "Ship the feature", "") + require.NoError(t, err) + + require.NoError(t, svc.DispatchAssignment(context.Background(), "user-1", assignment.ID)) + + var running model.AgentTeamAssignment + require.NoError(t, db.Where("id = ?", assignment.ID).First(&running).Error) + require.Equal(t, model.AssignmentStatusRunning, running.Status) + + require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", assignment.ID, "shipped")) + + var done model.AgentTeamAssignment + require.NoError(t, db.Where("id = ?", assignment.ID).First(&done).Error) + assert.Equal(t, model.AssignmentStatusDone, done.Status) + assert.Equal(t, "shipped", done.Result) + + // Legacy dispatched row (pre-#1384) must still complete. + legacy := &model.AgentTeamAssignment{ + TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, TaskPrompt: "legacy complete", Status: model.AssignmentStatusDispatched, + } + require.NoError(t, repository.CreateAssignment(db, legacy)) + require.NoError(t, svc.CompleteAssignment(context.Background(), "user-1", legacy.ID, "legacy shipped")) + var legacyDone model.AgentTeamAssignment + require.NoError(t, db.Where("id = ?", legacy.ID).First(&legacyDone).Error) + assert.Equal(t, model.AssignmentStatusDone, legacyDone.Status) +} + +// TestAgentTeamService_FailAssignmentRejectsTerminalAssignment verifies that a +// done/failed/cancelled assignment cannot be failed again, so repeated fails +// can no longer rewrite the stored result or re-trigger fault escalation. +func TestAgentTeamService_FailAssignmentRejectsTerminalAssignment(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + + for _, status := range []string{ + model.AssignmentStatusDone, + model.AssignmentStatusFailed, + model.AssignmentStatusCancelled, + } { + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "Ship it", + Status: status, + Result: "original outcome", + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + + err := svc.FailAssignment(context.Background(), "user-1", assignment.ID, "[fault_escalation] retries=3 maxRetries=3 error=late") + require.Error(t, err, "status=%s", status) + assert.Equal(t, errcode.ErrBadRequest, err, "status=%s", status) + + var reloaded model.AgentTeamAssignment + require.NoError(t, db.Where("id = ?", assignment.ID).First(&reloaded).Error) + assert.Equal(t, status, reloaded.Status, "terminal status must be preserved") + assert.Equal(t, "original outcome", reloaded.Result) + } + + // The rejected fails must not have appended any team events (no repeated + // escalation side effects). + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + assert.Empty(t, events) +} + +// ── Assignment lifecycle (#1384) ────────────────────────────────────────────── + +func TestAgentTeamService_FailTimedOutAssignmentsTerminatesActiveAndIsIdempotent(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + guardrails := DefaultAgentTeamGuardrails() + guardrails.AssignmentTimeout = time.Minute + svc := NewAgentTeamServiceWithGuardrails(db, nil, nil, guardrails) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + + old := time.Now().Add(-2 * time.Hour) + fresh := time.Now() + timedOut := &model.AgentTeamAssignment{ + TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, TaskPrompt: "stale work", Status: model.AssignmentStatusRunning, + CreatedAt: old, UpdatedAt: old, + } + freshA := &model.AgentTeamAssignment{ + TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, TaskPrompt: "fresh work", Status: model.AssignmentStatusPending, + CreatedAt: fresh, UpdatedAt: fresh, + } + require.NoError(t, repository.CreateAssignment(db, timedOut)) + require.NoError(t, repository.CreateAssignment(db, freshA)) + // Force created_at past the guardrail (GORM autoCreateTime may overwrite on insert). + require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", timedOut.ID).Update("created_at", old).Error) + require.NoError(t, db.Model(&model.AgentTeamAssignment{}).Where("id = ?", freshA.ID).Update("created_at", fresh).Error) + + n, err := svc.FailTimedOutAssignments(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, n) + + reloaded, err := repository.GetAssignmentByID(db, timedOut.ID) + require.NoError(t, err) + assert.Equal(t, model.AssignmentStatusFailed, reloaded.Status) + assert.Equal(t, "assignment timeout reached", reloaded.Result) + + stillFresh, err := repository.GetAssignmentByID(db, freshA.ID) + require.NoError(t, err) + assert.Equal(t, model.AssignmentStatusPending, stillFresh.Status) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventAssignmentFailed, events[0].Type) + + // Second scan must not rewrite or re-emit. + n, err = svc.FailTimedOutAssignments(context.Background()) + require.NoError(t, err) + assert.Equal(t, 0, n) + events, err = repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + assert.Len(t, events, 1) +} + +func TestAgentTeamService_CompleteAssignmentConcurrentOnlyOneWins(t *testing.T) { + db := setupAgentTeamConcurrentSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + _, supervisor, executor, run := seedAgentTeamRun(t, db) + assignment := &model.AgentTeamAssignment{ + TeamRunID: run.ID, FromMemberID: supervisor.ID, ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, TaskPrompt: "race complete", Status: model.AssignmentStatusRunning, + } + require.NoError(t, repository.CreateAssignment(db, assignment)) + + errs := runConcurrentRouteCalls(2, func() error { + return svc.CompleteAssignment(context.Background(), "user-1", assignment.ID, "winner") + }) + successes, bad := 0, 0 + for _, err := range errs { + switch err { + case nil: + successes++ + case errcode.ErrBadRequest: + bad++ + default: + require.NoError(t, err) + } + } + assert.Equal(t, 1, successes) + assert.Equal(t, 1, bad) + + done, err := repository.GetAssignmentByID(db, assignment.ID) + require.NoError(t, err) + assert.Equal(t, model.AssignmentStatusDone, done.Status) + assert.Equal(t, "winner", done.Result) +} diff --git a/hub-server/internal/service/agentteam/helpers_test.go b/hub-server/internal/service/agentteam/helpers_test.go index 2c37cc445..7c1a2aa3b 100644 --- a/hub-server/internal/service/agentteam/helpers_test.go +++ b/hub-server/internal/service/agentteam/helpers_test.go @@ -1,12 +1,20 @@ package agentteam import ( + "context" "database/sql" + "encoding/json" "fmt" + "path/filepath" "strings" "testing" + "time" "github.com/DATA-DOG/go-sqlmock" + "github.com/agenthub/hub-server/internal/bus" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" "gorm.io/driver/postgres" "gorm.io/gorm" @@ -32,3 +40,377 @@ func newMockDBAgent(t *testing.T) (*gorm.DB, sqlmock.Sqlmock, *sql.DB) { require.NoError(t, err) return gormDB, mock, sqlDB } + +// newMockAgentTeamDB creates a sqlmock-backed gorm.DB for agent team tests. +func newMockAgentTeamDB(t *testing.T) (*gorm.DB, sqlmock.Sqlmock) { + t.Helper() + db, mock, _ := newMockDBAgent(t) + return db, mock +} + +// mockAgentTeamAgentSvc implements agentTeamAgentSvc for tests. +type mockAgentTeamAgentSvc struct { + triggerMessageID string + targetID string + modelParams string + returnTaskID string +} + +func (m *mockAgentTeamAgentSvc) AddAgentToSession(ctx context.Context, userID, sessionID, agentType, customAgentID, displayName string) (*model.AgentInstance, error) { + return &model.AgentInstance{}, nil +} + +func (m *mockAgentTeamAgentSvc) TriggerAgentTask(ctx context.Context, userID, triggerMessageID, targetAgentInstanceID, targetAgentType, targetCustomAgentID, modelParams, targetID string) (*model.PendingAgentTask, error) { + m.triggerMessageID = triggerMessageID + m.targetID = targetID + m.modelParams = modelParams + taskID := m.returnTaskID + if taskID == "" { + taskID = "task-1" + } + return &model.PendingAgentTask{ID: taskID}, nil +} + +type mockAgentTeamControlSvc struct { + calls []agentTeamControlCall +} + +type agentTeamControlCall struct { + userID string + deviceID string + payload model.AgentControlPayload +} + +func (m *mockAgentTeamControlSvc) DeliverToDesktopDevice(ctx context.Context, userID, deviceID string, payload model.AgentControlPayload) error { + m.calls = append(m.calls, agentTeamControlCall{ + userID: userID, + deviceID: deviceID, + payload: payload, + }) + return nil +} + +func readAgentTeamEvent(t *testing.T, events <-chan bus.Event) bus.Event { + t.Helper() + select { + case event := <-events: + return event + case <-time.After(time.Second): + t.Fatal("agent team event was not published") + } + return bus.Event{} +} + +func setupAgentTeamStateSQLite(t *testing.T) *gorm.DB { + t.Helper() + return setupAgentTeamStateSQLiteDSN(t, ":memory:", 1) +} + +func setupAgentTeamConcurrentSQLite(t *testing.T) *gorm.DB { + t.Helper() + path := filepath.ToSlash(filepath.Join(t.TempDir(), "agentteam-concurrency.db")) + dsn := fmt.Sprintf("file:%s?cache=shared&_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)", path) + return setupAgentTeamStateSQLiteDSN(t, dsn, 8) +} + +func setupAgentTeamStateSQLiteDSN(t *testing.T, dsn string, maxOpenConns int) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) + require.NoError(t, err) + sqlDB, err := db.DB() + require.NoError(t, err) + t.Cleanup(func() { _ = sqlDB.Close() }) + sqlDB.SetMaxOpenConns(maxOpenConns) + tables := []string{ + `CREATE TABLE agent_teams ( + id TEXT PRIMARY KEY, + owner_id TEXT NOT NULL, + name TEXT NOT NULL, + description TEXT DEFAULT '', + avatar_url TEXT DEFAULT '', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_members ( + id TEXT PRIMARY KEY, + team_id TEXT NOT NULL, + agent_profile_id TEXT, + role TEXT NOT NULL DEFAULT 'executor', + position INTEGER NOT NULL DEFAULT 0, + created_at DATETIME + )`, + `CREATE TABLE agent_team_runs ( + id TEXT PRIMARY KEY, + team_id TEXT NOT NULL, + session_id TEXT, + trigger_user_id TEXT NOT NULL, + trigger_message TEXT DEFAULT '', + target_id TEXT, + mode TEXT NOT NULL DEFAULT 'supervisor', + status TEXT NOT NULL DEFAULT 'queued', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_assignments ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + from_member_id TEXT NOT NULL, + to_member_id TEXT NOT NULL, + type TEXT NOT NULL DEFAULT 'delegate', + task_prompt TEXT NOT NULL, + context TEXT DEFAULT '', + status TEXT NOT NULL DEFAULT 'pending', + run_id TEXT, + result TEXT DEFAULT '', + depth INTEGER NOT NULL DEFAULT 0, + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_tasks ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + assignment_id TEXT, + assignee_member_id TEXT NOT NULL, + parent_task_id TEXT, + status TEXT NOT NULL DEFAULT 'pending', + objective TEXT NOT NULL, + input_refs TEXT NOT NULL DEFAULT '{}', + run_id TEXT, + attempt INTEGER NOT NULL DEFAULT 1, + risk_level TEXT NOT NULL DEFAULT 'normal', + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE agent_team_events ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + seq INTEGER NOT NULL, + type TEXT NOT NULL, + payload TEXT NOT NULL DEFAULT '{}', + created_at DATETIME + )`, + `CREATE TABLE agent_team_artifacts ( + id TEXT PRIMARY KEY, + team_run_id TEXT NOT NULL, + team_task_id TEXT, + assignment_id TEXT, + member_id TEXT, + agent_task_id TEXT, + edge_run_id TEXT, + source_event_id TEXT, + event_seq INTEGER NOT NULL DEFAULT 0, + path TEXT NOT NULL, + normalized_path TEXT NOT NULL, + action TEXT, + tool_name TEXT, + status TEXT, + conflict_id TEXT, + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE sessions ( + id TEXT PRIMARY KEY, + type TEXT NOT NULL, + name TEXT DEFAULT '', + owner_user_id TEXT, + workspace_id TEXT, + next_seq INTEGER NOT NULL DEFAULT 0, + last_message_at DATETIME, + dissolved BOOLEAN NOT NULL DEFAULT FALSE, + created_at DATETIME + )`, + `CREATE TABLE session_members ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + member_type TEXT NOT NULL, + member_id TEXT NOT NULL, + role TEXT NOT NULL, + pinned BOOLEAN NOT NULL DEFAULT FALSE, + archived BOOLEAN NOT NULL DEFAULT FALSE, + muted BOOLEAN NOT NULL DEFAULT FALSE, + last_read_seq INTEGER NOT NULL DEFAULT 0, + joined_at DATETIME, + left_at DATETIME + )`, + `CREATE TABLE agent_instances ( + id TEXT PRIMARY KEY, + agent_type TEXT NOT NULL, + custom_agent_id TEXT, + session_id TEXT NOT NULL, + inviter_user_id TEXT NOT NULL, + workspace_id TEXT, + display_name TEXT NOT NULL, + created_at DATETIME + )`, + `CREATE TABLE custom_agents ( + id TEXT PRIMARY KEY, + owner_user_id TEXT NOT NULL, + name TEXT NOT NULL, + avatar_url TEXT DEFAULT '', + agent_type TEXT NOT NULL, + system_prompt TEXT DEFAULT '', + capability_tags TEXT DEFAULT '[]', + tool_whitelist TEXT DEFAULT '[]', + model_params TEXT DEFAULT '{}', + output_schema TEXT DEFAULT NULL, + deleted_at DATETIME, + created_at DATETIME, + updated_at DATETIME + )`, + `CREATE TABLE pending_agent_tasks ( + id TEXT PRIMARY KEY, + agent_instance_id TEXT NOT NULL, + triggered_by_user_id TEXT NOT NULL, + trigger_message_id TEXT NOT NULL, + target_id TEXT, + status TEXT NOT NULL, + edge_run_id TEXT DEFAULT '', + edge_device_id TEXT, + error_message TEXT, + model_params TEXT DEFAULT '{}', + created_at DATETIME, + dispatched_at DATETIME, + finished_at DATETIME, + expire_at DATETIME NOT NULL + )`, + `CREATE TABLE messages ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + seq_id INTEGER NOT NULL, + client_msg_id TEXT NOT NULL, + sender_type TEXT NOT NULL, + sender_id TEXT NOT NULL, + content_type TEXT NOT NULL, + content TEXT NOT NULL, + reply_to_message_id TEXT, + recalled BOOLEAN NOT NULL DEFAULT FALSE, + edited BOOLEAN NOT NULL DEFAULT FALSE, + edited_at DATETIME, + created_at DATETIME + )`, + `CREATE UNIQUE INDEX idx_messages_session_client_msg ON messages (session_id, client_msg_id)`, + `CREATE TABLE agent_run_events ( + id TEXT PRIMARY KEY, + task_id TEXT NOT NULL, + edge_run_id TEXT, + session_id TEXT NOT NULL, + agent_instance_id TEXT NOT NULL, + event_seq INTEGER NOT NULL, + event_type TEXT NOT NULL, + payload TEXT NOT NULL, + created_at DATETIME + )`, + } + for _, ddl := range tables { + require.NoError(t, db.Exec(ddl).Error) + } + return db +} + +func seedAgentTeamRun(t *testing.T, db *gorm.DB) (*model.AgentTeam, *model.AgentTeamMember, *model.AgentTeamMember, *model.AgentTeamRun) { + t.Helper() + team := &model.AgentTeam{OwnerID: "user-1", Name: "Route Team"} + require.NoError(t, repository.CreateTeam(db, team)) + + supervisorProfileID := "profile-supervisor" + executorProfileID := "profile-executor" + supervisor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &supervisorProfileID, + Role: model.TeamMemberRoleSupervisor, + } + executor := &model.AgentTeamMember{ + TeamID: team.ID, + AgentProfileID: &executorProfileID, + Role: model.TeamMemberRoleExecutor, + } + require.NoError(t, repository.AddTeamMember(db, supervisor)) + require.NoError(t, repository.AddTeamMember(db, executor)) + + run := &model.AgentTeamRun{ + TeamID: team.ID, + SessionID: "session-1", + TriggerUserID: "user-1", + TriggerMessage: "ship it", + Status: model.TeamRunStatusRunning, + } + require.NoError(t, repository.CreateTeamRun(db, run)) + return team, supervisor, executor, run +} + +func addReadableTeamMemberForUser(t *testing.T, db *gorm.DB, teamID, userID string) *model.AgentTeamMember { + t.Helper() + agent := &model.CustomAgent{ + OwnerUserID: userID, + Name: "Readable Member Agent", + AgentType: "codex", + SystemPrompt: "Read shared team state", + } + require.NoError(t, repository.CreateCustomAgent(db, agent)) + member := &model.AgentTeamMember{ + TeamID: teamID, + AgentProfileID: &agent.ID, + Role: model.TeamMemberRoleExecutor, + } + require.NoError(t, repository.AddTeamMember(db, member)) + return member +} + +func addTeamSupervisor(t *testing.T, db *gorm.DB, teamID, profileID string) *model.AgentTeamMember { + t.Helper() + member := &model.AgentTeamMember{ + TeamID: teamID, + AgentProfileID: &profileID, + Role: model.TeamMemberRoleSupervisor, + } + require.NoError(t, repository.AddTeamMember(db, member)) + return member +} + +func stringPtr(value string) *string { + return &value +} + +func mustJSON(t *testing.T, value any) string { + t.Helper() + data, err := json.Marshal(value) + require.NoError(t, err) + return string(data) +} + +func seedTeamRunSession(t *testing.T, db *gorm.DB, sessionID, userID string, executor *model.AgentTeamMember) { + t.Helper() + now := time.Now() + require.NoError(t, db.Exec( + `INSERT INTO sessions (id, type, name, owner_user_id, next_seq, dissolved, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, + sessionID, model.SessionTypeGroup, "Team session", userID, 0, false, now, + ).Error) + require.NoError(t, db.Exec( + `INSERT INTO session_members (id, session_id, member_type, member_id, role, joined_at) VALUES (?, ?, ?, ?, ?, ?)`, + "session-member-user", sessionID, model.MemberTypeUser, userID, model.MemberRoleOwner, now, + ).Error) + customAgentID := "" + if executor.AgentProfileID != nil { + customAgentID = *executor.AgentProfileID + } + require.NoError(t, db.Exec( + `INSERT INTO agent_instances (id, agent_type, custom_agent_id, session_id, inviter_user_id, display_name, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)`, + "agent-executor", "codex", customAgentID, sessionID, userID, "Executor", now, + ).Error) + require.NoError(t, db.Exec( + `INSERT INTO session_members (id, session_id, member_type, member_id, role, joined_at) VALUES (?, ?, ?, ?, ?, ?)`, + "session-member-agent", sessionID, model.MemberTypeAgent, "agent-executor", model.MemberRoleMember, now, + ).Error) +} + +func entryMemberIDs(entries []model.CompeteSummaryEntry) []string { + ids := make([]string, len(entries)) + for i, e := range entries { + ids[i] = e.MemberID + } + return ids +} + +func strPtr(s string) *string { + return &s +} diff --git a/hub-server/internal/service/agentteam/route_decision_test.go b/hub-server/internal/service/agentteam/route_decision_test.go new file mode 100644 index 000000000..262822fc7 --- /dev/null +++ b/hub-server/internal/service/agentteam/route_decision_test.go @@ -0,0 +1,351 @@ +package agentteam + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/agenthub/hub-server/internal/errcode" + "github.com/agenthub/hub-server/internal/model" + "github.com/agenthub/hub-server/internal/repository" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAgentTeamService_HandleRouteDecisionCreatesAssignmentAndAuditEvents(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, supervisor, executor, run := seedAgentTeamRun(t, db) + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Implement the replay UI", + Reasoning: "executor owns UI work", + Reason: "worker owns fixture implementation", + Context: "state endpoint is ready", + AgentID: executor.ID, + ParentTaskID: "parent-task-1", + CorrelationID: "corr-route-1", + }) + require.NoError(t, err) + require.NotNil(t, assignment) + assert.Equal(t, supervisor.ID, assignment.FromMemberID) + assert.Equal(t, executor.ID, assignment.ToMemberID) + assert.Equal(t, model.AssignmentTypeDelegate, assignment.Type) + assert.Equal(t, model.AssignmentStatusPending, assignment.Status) + assert.Equal(t, 1, assignment.Depth) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 3) + assert.Equal(t, model.TeamEventRouteDecided, events[0].Type) + assert.Equal(t, model.TeamEventAssignmentCreated, events[1].Type) + assert.Equal(t, model.TeamEventTaskCreated, events[2].Type) + var accepted model.CoordinatorRouteDecision + require.NoError(t, json.Unmarshal([]byte(events[0].Payload), &accepted)) + require.NotEmpty(t, accepted.SubtaskID) + assert.True(t, accepted.Accepted) + assert.Equal(t, executor.ID, accepted.AgentID) + assert.Equal(t, "parent-task-1", accepted.ParentTaskID) + assert.Equal(t, "worker owns fixture implementation", accepted.Reason) + assert.Equal(t, "corr-route-1", accepted.CorrelationID) + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.RouteLog, 1) + assert.Equal(t, "delegate", state.RouteLog[0].Action) + assert.Equal(t, "corr-route-1", state.RouteLog[0].CorrelationID) + require.Len(t, state.RouteAuditLog, 1) + assert.Equal(t, "accepted", state.RouteAuditLog[0].Status) + assert.Equal(t, "corr-route-1", state.RouteAuditLog[0].CorrelationID) + assert.Equal(t, accepted.SubtaskID, state.RouteAuditLog[0].SubtaskID) + assert.Equal(t, "parent-task-1", state.RouteAuditLog[0].ParentTaskID) + assert.Equal(t, executor.ID, state.RouteAuditLog[0].AgentID) + assert.Equal(t, "worker owns fixture implementation", state.RouteAuditLog[0].Reason) + require.Len(t, state.Assignments, 1) + assert.Equal(t, assignment.ID, state.Assignments[0].AssignmentID) + assert.Equal(t, 1, state.Members[1].ActiveTasks) + require.Len(t, state.Tasks, 1) + assert.Equal(t, assignment.ID, state.Tasks[0].AssignmentID) + assert.Equal(t, executor.ID, state.Tasks[0].AssigneeMemberID) + assert.Equal(t, "parent-task-1", state.Tasks[0].ParentTaskID) + assert.Equal(t, "Implement the replay UI", state.Tasks[0].Objective) + assert.Equal(t, model.TeamTaskStatusPending, state.Tasks[0].Status) +} + +func TestAgentTeamService_HandleRouteDecisionRejectsMissingWorkerWithAuditEvent(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, _, run := seedAgentTeamRun(t, db) + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: "missing-member", + Instructions: "Do work", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) + assert.Contains(t, events[0].Payload, "next_worker") + + state, err := svc.GetTeamRunState(context.Background(), "user-1", team.ID, run.ID) + require.NoError(t, err) + require.Len(t, state.RouteAuditLog, 1) + assert.Equal(t, "rejected", state.RouteAuditLog[0].Status) + assert.Equal(t, "missing-member", state.RouteAuditLog[0].AgentID) + assert.Equal(t, "next_worker is not a team member", state.RouteAuditLog[0].Reason) +} + +func TestAgentTeamService_HandleRouteDecisionRejectsWhenTaskLimitReached(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, supervisor, executor, run := seedAgentTeamRun(t, db) + for i := 0; i < model.MaxTasksPerTeamRun; i++ { + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "existing task", + Status: model.AssignmentStatusDone, + Depth: 1, + })) + } + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "one more task", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) + assert.Contains(t, events[0].Payload, "task limit") +} + +func TestAgentTeamService_HandleRouteDecisionRejectsWhenActiveSubagentLimitReached(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, supervisor, executor, run := seedAgentTeamRun(t, db) + for i := 0; i < model.MaxActiveSubAgentsPerRun; i++ { + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "active task", + Status: model.AssignmentStatusRunning, + Depth: 1, + })) + } + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "blocked by active limit", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) + assert.Contains(t, events[0].Payload, "active subagent limit") +} + +func TestAgentTeamService_HandleRouteDecisionRejectsWhenRouteRepeatLimitReached(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + repeated := model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "repeat me", + Reasoning: "same route", + } + for i := 0; i < model.MaxRouteRepeats; i++ { + require.NoError(t, repository.AppendTeamEvent(db, &model.AgentTeamEvent{ + TeamRunID: run.ID, + Type: model.TeamEventRouteDecided, + Payload: mustJSON(t, repeated), + })) + } + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, repeated) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, model.MaxRouteRepeats+1) + assert.Equal(t, model.TeamEventRouteRejected, events[len(events)-1].Type) + assert.Contains(t, events[len(events)-1].Payload, "route repeat limit") +} + +func TestAgentTeamService_HandleRouteDecisionRejectsTimedOutAssignment(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, supervisor, executor, run := seedAgentTeamRun(t, db) + require.NoError(t, repository.CreateAssignment(db, &model.AgentTeamAssignment{ + TeamRunID: run.ID, + FromMemberID: supervisor.ID, + ToMemberID: executor.ID, + Type: model.AssignmentTypeDelegate, + TaskPrompt: "stale task", + Status: model.AssignmentStatusRunning, + Depth: 1, + CreatedAt: time.Now().Add(-model.DefaultAssignmentTimeout - time.Minute), + })) + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "blocked by timeout", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) + assert.Contains(t, events[0].Payload, "assignment timeout") +} + +func TestAgentTeamService_HandleRouteDecisionRejectsBudgetExceeded(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + agentTaskID := "budget-task-1" + require.NoError(t, repository.CreateTeamTask(db, &model.AgentTeamTask{ + TeamRunID: run.ID, + AssigneeMemberID: executor.ID, + Status: model.TeamTaskStatusDone, + Objective: "spent budget", + RunID: &agentTaskID, + })) + require.NoError(t, repository.CreateAgentRunEventWithNextSeq(db, &model.AgentRunEvent{ + TaskID: agentTaskID, + EdgeRunID: "edge-budget", + SessionID: run.SessionID, + AgentInstanceID: "agent-executor", + EventType: "run.agent.result", + Payload: `{"success":true,"usage":{"input_tokens":600,"output_tokens":400},"tokenLimit":1000}`, + })) + + assignment, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "blocked by budget", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + assert.Nil(t, assignment) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + require.Len(t, events, 1) + assert.Equal(t, model.TeamEventRouteRejected, events[0].Type) + assert.Contains(t, events[0].Payload, "budget exceeded") +} + +// TestAgentTeamService_HandleRouteDecisionRejectsTerminalRun verifies that no +// route decision (delegate or a repeated finish) can mutate a run that already +// reached a terminal status, and that each rejection is audited. +func TestAgentTeamService_HandleRouteDecisionRejectsTerminalRun(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + team, _, executor, run := seedAgentTeamRun(t, db) + require.NoError(t, repository.UpdateTeamRunStatus(db, run.ID, model.TeamRunStatusCompleted)) + + // Delegate on a completed run must be rejected. + _, err := svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "delegate", + NextWorker: executor.ID, + Instructions: "Do more work", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + + // A repeated finish with blocked_reason must not downgrade completed to failed. + _, err = svc.HandleRouteDecision(context.Background(), "user-1", team.ID, run.ID, model.CoordinatorRouteDecision{ + Action: "finish", + BlockedReason: "should not overwrite", + }) + require.Error(t, err) + assert.ErrorIs(t, err, errcode.ErrBadRequest) + + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusCompleted, gotRun.Status) + + assignments, err := repository.ListAssignmentsByTeamRun(db, run.ID) + require.NoError(t, err) + assert.Empty(t, assignments) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + rejected := 0 + for _, e := range events { + if e.Type == model.TeamEventRouteRejected { + rejected++ + } + } + assert.Equal(t, 2, rejected, "expected each terminal-run route decision to append team.route.rejected") +} + +// TestAgentTeamService_FinishRouteDecisionDoesNotDowngradeTerminalRun calls +// finishRouteDecision directly (bypassing the HandleRouteDecision entry guard, +// as a racing finish would) and verifies the conditional status write keeps +// the first terminal outcome and does not emit a conflicting run.failed event. +func TestAgentTeamService_FinishRouteDecisionDoesNotDowngradeTerminalRun(t *testing.T) { + db := setupAgentTeamStateSQLite(t) + svc := NewAgentTeamService(db, nil, nil) + _, _, _, run := seedAgentTeamRun(t, db) + + require.NoError(t, svc.finishRouteDecision(run.ID, model.CoordinatorRouteDecision{ + Action: "finish", + Summary: "all done", + })) + gotRun, err := repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + require.Equal(t, model.TeamRunStatusCompleted, gotRun.Status) + + require.NoError(t, svc.finishRouteDecision(run.ID, model.CoordinatorRouteDecision{ + Action: "finish", + BlockedReason: "late failure", + })) + gotRun, err = repository.GetTeamRunByID(db, run.ID) + require.NoError(t, err) + assert.Equal(t, model.TeamRunStatusCompleted, gotRun.Status) + + events, err := repository.ListTeamEventsByRun(db, run.ID) + require.NoError(t, err) + completedEvents, failedEvents := 0, 0 + for _, e := range events { + switch e.Type { + case model.TeamEventRunCompleted: + completedEvents++ + case model.TeamEventRunFailed: + failedEvents++ + } + } + assert.Equal(t, 1, completedEvents) + assert.Zero(t, failedEvents, "repeated finish must not append team.run.failed") +} diff --git a/hub-server/internal/service/agentteam/route_helpers_test.go b/hub-server/internal/service/agentteam/route_helpers_test.go index da536b22e..c6e60c8ba 100644 --- a/hub-server/internal/service/agentteam/route_helpers_test.go +++ b/hub-server/internal/service/agentteam/route_helpers_test.go @@ -264,3 +264,11 @@ func TestSupervisorRouteModelParams(t *testing.T) { assert.Contains(t, params["structured_output_schema"], `"action"`) assert.Contains(t, params["append_system_prompt"], "supervisor mode") } + +func TestIsTerminalTeamRunStatusIncludesCancelledWithoutCancelAPI(t *testing.T) { + // Documented dead write-path: cancelled is a terminal guard token only. + assert.True(t, isTerminalTeamRunStatus(model.TeamRunStatusCancelled)) + assert.True(t, isTerminalTeamRunStatus(model.TeamRunStatusCompleted)) + assert.True(t, isTerminalTeamRunStatus(model.TeamRunStatusFailed)) + assert.False(t, isTerminalTeamRunStatus(model.TeamRunStatusRunning)) +}