feat(coderd/database): add ListTasks query (#20282)

Relates to https://github.com/coder/internal/issues/981

Adds a `ListTasks` query that allows filtering by OwnerID and OrganizationID.
This commit is contained in:
Cian Johnston
2025-10-14 17:33:30 +01:00
committed by GitHub
parent 06db58771f
commit 9f229370e7
8 changed files with 231 additions and 0 deletions
+5
View File
@@ -4485,6 +4485,11 @@ func (q *querier) ListProvisionerKeysByOrganizationExcludeReserved(ctx context.C
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.ListProvisionerKeysByOrganizationExcludeReserved)(ctx, organizationID)
}
func (q *querier) ListTasks(ctx context.Context, arg database.ListTasksParams) ([]database.Task, error) {
// TODO(Cian): replace this with a sql filter for improved performance. https://github.com/coder/internal/issues/1061
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.ListTasks)(ctx, arg)
}
func (q *querier) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.UserSecret, error) {
obj := rbac.ResourceUserSecret.WithOwner(userID.String())
if err := q.authorizeContext(ctx, policy.ActionRead, obj); err != nil {
+12
View File
@@ -2392,6 +2392,18 @@ func (s *MethodTestSuite) TestTasks() {
dbm.EXPECT().GetTaskByWorkspaceID(gomock.Any(), task.WorkspaceID.UUID).Return(task, nil).AnyTimes()
check.Args(task.WorkspaceID.UUID).Asserts(task, policy.ActionRead).Returns(task)
}))
s.Run("ListTasks", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
u1 := testutil.Fake(s.T(), faker, database.User{})
u2 := testutil.Fake(s.T(), faker, database.User{})
org1 := testutil.Fake(s.T(), faker, database.Organization{})
org2 := testutil.Fake(s.T(), faker, database.Organization{})
_ = testutil.Fake(s.T(), faker, database.OrganizationMember{UserID: u1.ID, OrganizationID: org1.ID})
_ = testutil.Fake(s.T(), faker, database.OrganizationMember{UserID: u2.ID, OrganizationID: org2.ID})
t1 := testutil.Fake(s.T(), faker, database.Task{OwnerID: u1.ID})
t2 := testutil.Fake(s.T(), faker, database.Task{OwnerID: u2.ID})
dbm.EXPECT().ListTasks(gomock.Any(), gomock.Any()).Return([]database.Task{t1, t2}, nil).AnyTimes()
check.Args(database.ListTasksParams{}).Asserts(t1, policy.ActionRead, t2, policy.ActionRead).Returns([]database.Task{t1, t2})
}))
}
func (s *MethodTestSuite) TestProvisionerKeys() {
@@ -2735,6 +2735,13 @@ func (m queryMetricsStore) ListProvisionerKeysByOrganizationExcludeReserved(ctx
return r0, r1
}
func (m queryMetricsStore) ListTasks(ctx context.Context, arg database.ListTasksParams) ([]database.Task, error) {
start := time.Now()
r0, r1 := m.s.ListTasks(ctx, arg)
m.queryLatencies.WithLabelValues("ListTasks").Observe(time.Since(start).Seconds())
return r0, r1
}
func (m queryMetricsStore) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.UserSecret, error) {
start := time.Now()
r0, r1 := m.s.ListUserSecrets(ctx, userID)
+15
View File
@@ -5848,6 +5848,21 @@ func (mr *MockStoreMockRecorder) ListProvisionerKeysByOrganizationExcludeReserve
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListProvisionerKeysByOrganizationExcludeReserved", reflect.TypeOf((*MockStore)(nil).ListProvisionerKeysByOrganizationExcludeReserved), ctx, organizationID)
}
// ListTasks mocks base method.
func (m *MockStore) ListTasks(ctx context.Context, arg database.ListTasksParams) ([]database.Task, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ListTasks", ctx, arg)
ret0, _ := ret[0].([]database.Task)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// ListTasks indicates an expected call of ListTasks.
func (mr *MockStoreMockRecorder) ListTasks(ctx, arg any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListTasks", reflect.TypeOf((*MockStore)(nil).ListTasks), ctx, arg)
}
// ListUserSecrets mocks base method.
func (m *MockStore) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.UserSecret, error) {
m.ctrl.T.Helper()
+1
View File
@@ -595,6 +595,7 @@ type sqlcQuerier interface {
ListAIBridgeUserPromptsByInterceptionIDs(ctx context.Context, interceptionIds []uuid.UUID) ([]AIBridgeUserPrompt, error)
ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error)
ListProvisionerKeysByOrganizationExcludeReserved(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error)
ListTasks(ctx context.Context, arg ListTasksParams) ([]Task, error)
ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]UserSecret, error)
ListWorkspaceAgentPortShares(ctx context.Context, workspaceID uuid.UUID) ([]WorkspaceAgentPortShare, error)
MarkAllInboxNotificationsAsRead(ctx context.Context, arg MarkAllInboxNotificationsAsReadParams) error
+136
View File
@@ -7314,3 +7314,139 @@ func TestUsageEventsTrigger(t *testing.T) {
require.Len(t, rows, 0)
})
}
func TestListTasks(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
// Given: two organizations and two users, one of which is a member of both
org1 := dbgen.Organization(t, db, database.Organization{})
org2 := dbgen.Organization(t, db, database.Organization{})
user1 := dbgen.User(t, db, database.User{})
user2 := dbgen.User(t, db, database.User{})
_ = dbgen.OrganizationMember(t, db, database.OrganizationMember{
OrganizationID: org1.ID,
UserID: user1.ID,
})
_ = dbgen.OrganizationMember(t, db, database.OrganizationMember{
OrganizationID: org2.ID,
UserID: user2.ID,
})
// Given: a template with an active version
tv := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user1.ID,
OrganizationID: org1.ID,
})
tpl := dbgen.Template(t, db, database.Template{
CreatedBy: user1.ID,
OrganizationID: org1.ID,
ActiveVersionID: tv.ID,
})
// Helper function to create a task
createTask := func(orgID, ownerID uuid.UUID) database.TaskTable {
ws := dbgen.Workspace(t, db, database.WorkspaceTable{
OrganizationID: orgID,
OwnerID: ownerID,
TemplateID: tpl.ID,
})
pj := dbgen.ProvisionerJob(t, db, ps, database.ProvisionerJob{})
sidebarAppID := uuid.New()
_ = dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
JobID: pj.ID,
TemplateVersionID: tv.ID,
WorkspaceID: ws.ID,
})
wr := dbgen.WorkspaceResource(t, db, database.WorkspaceResource{
JobID: pj.ID,
})
agt := dbgen.WorkspaceAgent(t, db, database.WorkspaceAgent{
ResourceID: wr.ID,
})
wa := dbgen.WorkspaceApp(t, db, database.WorkspaceApp{
ID: sidebarAppID,
AgentID: agt.ID,
})
tsk := dbgen.Task(t, db, database.TaskTable{
OrganizationID: orgID,
OwnerID: ownerID,
Prompt: testutil.GetRandomName(t),
TemplateVersionID: tv.ID,
WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true},
})
_ = dbgen.TaskWorkspaceApp(t, db, database.TaskWorkspaceApp{
TaskID: tsk.ID,
WorkspaceAgentID: uuid.NullUUID{Valid: true, UUID: agt.ID},
WorkspaceAppID: uuid.NullUUID{Valid: true, UUID: wa.ID},
})
t.Logf("task_id:%s owner_id:%s org_id:%s", tsk.ID, ownerID, orgID)
return tsk
}
// Given: user1 has one task, user2 has one task, user3 has two tasks (one in each org)
task1 := createTask(org1.ID, user1.ID)
task2 := createTask(org1.ID, user2.ID)
task3 := createTask(org2.ID, user2.ID)
// Then: run various filters and assert expected results
for _, tc := range []struct {
name string
filter database.ListTasksParams
expectIDs []uuid.UUID
}{
{
name: "no filter",
filter: database.ListTasksParams{
OwnerID: uuid.Nil,
OrganizationID: uuid.Nil,
},
expectIDs: []uuid.UUID{task3.ID, task2.ID, task1.ID},
},
{
name: "filter by user ID",
filter: database.ListTasksParams{
OwnerID: user1.ID,
OrganizationID: uuid.Nil,
},
expectIDs: []uuid.UUID{task1.ID},
},
{
name: "filter by organization ID",
filter: database.ListTasksParams{
OwnerID: uuid.Nil,
OrganizationID: org1.ID,
},
expectIDs: []uuid.UUID{task2.ID, task1.ID},
},
{
name: "filter by user and organization ID",
filter: database.ListTasksParams{
OwnerID: user2.ID,
OrganizationID: org2.ID,
},
expectIDs: []uuid.UUID{task3.ID},
},
{
name: "no results",
filter: database.ListTasksParams{
OwnerID: user1.ID,
OrganizationID: org2.ID,
},
expectIDs: nil,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
tasks, err := db.ListTasks(ctx, tc.filter)
if assert.NoError(t, err) {
require.Len(t, tasks, len(tc.expectIDs))
for idx, eid := range tc.expectIDs {
assert.Equal(t, eid.String(), tasks[idx].ID.String())
}
}
})
}
}
+48
View File
@@ -12599,6 +12599,54 @@ func (q *sqlQuerier) InsertTask(ctx context.Context, arg InsertTaskParams) (Task
return i, err
}
const listTasks = `-- name: ListTasks :many
SELECT id, organization_id, owner_id, name, workspace_id, template_version_id, template_parameters, prompt, created_at, deleted_at, status FROM tasks_with_status tws
WHERE tws.deleted_at IS NULL
AND CASE WHEN $1::UUID != '00000000-0000-0000-0000-000000000000' THEN tws.owner_id = $1::UUID ELSE TRUE END
AND CASE WHEN $2::UUID != '00000000-0000-0000-0000-000000000000' THEN tws.organization_id = $2::UUID ELSE TRUE END
ORDER BY tws.created_at DESC
`
type ListTasksParams struct {
OwnerID uuid.UUID `db:"owner_id" json:"owner_id"`
OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"`
}
func (q *sqlQuerier) ListTasks(ctx context.Context, arg ListTasksParams) ([]Task, error) {
rows, err := q.db.QueryContext(ctx, listTasks, arg.OwnerID, arg.OrganizationID)
if err != nil {
return nil, err
}
defer rows.Close()
var items []Task
for rows.Next() {
var i Task
if err := rows.Scan(
&i.ID,
&i.OrganizationID,
&i.OwnerID,
&i.Name,
&i.WorkspaceID,
&i.TemplateVersionID,
&i.TemplateParameters,
&i.Prompt,
&i.CreatedAt,
&i.DeletedAt,
&i.Status,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const upsertTaskWorkspaceApp = `-- name: UpsertTaskWorkspaceApp :one
INSERT INTO task_workspace_apps
(task_id, workspace_build_number, workspace_agent_id, workspace_app_id)
+7
View File
@@ -21,3 +21,10 @@ SELECT * FROM tasks_with_status WHERE id = @id::uuid;
-- name: GetTaskByWorkspaceID :one
SELECT * FROM tasks_with_status WHERE workspace_id = @workspace_id::uuid;
-- name: ListTasks :many
SELECT * FROM tasks_with_status tws
WHERE tws.deleted_at IS NULL
AND CASE WHEN @owner_id::UUID != '00000000-0000-0000-0000-000000000000' THEN tws.owner_id = @owner_id::UUID ELSE TRUE END
AND CASE WHEN @organization_id::UUID != '00000000-0000-0000-0000-000000000000' THEN tws.organization_id = @organization_id::UUID ELSE TRUE END
ORDER BY tws.created_at DESC;