From 939dcdc0ac23204580b3a1e71754798e2d5a26d1 Mon Sep 17 00:00:00 2001 From: Presley Pizzo Date: Mon, 17 Oct 2022 19:42:32 +0000 Subject: [PATCH] Write interface and database fake --- coderd/database/databasefake/databasefake.go | 150 +++++++++++++++++++ coderd/database/modelqueries.go | 51 +++++++ 2 files changed, 201 insertions(+) diff --git a/coderd/database/databasefake/databasefake.go b/coderd/database/databasefake/databasefake.go index bdb4c8e0f0..a25f35e604 100644 --- a/coderd/database/databasefake/databasefake.go +++ b/coderd/database/databasefake/databasefake.go @@ -711,6 +711,156 @@ func (q *fakeQuerier) GetAuthorizedWorkspaces(ctx context.Context, arg database. return workspaces, nil } +func (q *fakeQuerier) GetWorkspaceCount(ctx context.Context, arg database.GetWorkspaceCountParams) (int64, error) { + count, err := q.GetAuthorizedWorkspaceCount(ctx, arg, nil) + return count, err +} + +//nolint:gocyclo +func (q *fakeQuerier) GetAuthorizedWorkspaceCount(ctx context.Context, arg database.GetWorkspacesParams, authorizedFilter rbac.AuthorizeFilter) (int64, error) { + q.mutex.RLock() + defer q.mutex.RUnlock() + + workspaces := make([]database.Workspace, 0) + for _, workspace := range q.workspaces { + if arg.OwnerID != uuid.Nil && workspace.OwnerID != arg.OwnerID { + continue + } + + if arg.OwnerUsername != "" { + owner, err := q.GetUserByID(ctx, workspace.OwnerID) + if err == nil && !strings.EqualFold(arg.OwnerUsername, owner.Username) { + continue + } + } + + if arg.TemplateName != "" { + template, err := q.GetTemplateByID(ctx, workspace.TemplateID) + if err == nil && !strings.EqualFold(arg.TemplateName, template.Name) { + continue + } + } + + if !arg.Deleted && workspace.Deleted { + continue + } + + if arg.Name != "" && !strings.Contains(strings.ToLower(workspace.Name), strings.ToLower(arg.Name)) { + continue + } + + if arg.Status != "" { + build, err := q.GetLatestWorkspaceBuildByWorkspaceID(ctx, workspace.ID) + if err != nil { + return 0, xerrors.Errorf("get latest build: %w", err) + } + + job, err := q.GetProvisionerJobByID(ctx, build.JobID) + if err != nil { + return 0, xerrors.Errorf("get provisioner job: %w", err) + } + + switch arg.Status { + case "pending": + if !job.StartedAt.Valid { + continue + } + + case "starting": + if !job.StartedAt.Valid && + !job.CanceledAt.Valid && + job.CompletedAt.Valid && + time.Since(job.UpdatedAt) > 30*time.Second || + build.Transition != database.WorkspaceTransitionStart { + continue + } + + case "running": + if !job.CompletedAt.Valid && + job.CanceledAt.Valid && + job.Error.Valid || + build.Transition != database.WorkspaceTransitionStart { + continue + } + + case "stopping": + if !job.StartedAt.Valid && + !job.CanceledAt.Valid && + job.CompletedAt.Valid && + time.Since(job.UpdatedAt) > 30*time.Second || + build.Transition != database.WorkspaceTransitionStop { + continue + } + + case "stopped": + if !job.CompletedAt.Valid && + job.CanceledAt.Valid && + job.Error.Valid || + build.Transition != database.WorkspaceTransitionStop { + continue + } + + case "failed": + if (!job.CanceledAt.Valid && !job.Error.Valid) || + (!job.CompletedAt.Valid && !job.Error.Valid) { + continue + } + + case "canceling": + if !job.CanceledAt.Valid && job.CompletedAt.Valid { + continue + } + + case "canceled": + if !job.CanceledAt.Valid && !job.CompletedAt.Valid { + continue + } + + case "deleted": + if !job.StartedAt.Valid && + job.CanceledAt.Valid && + !job.CompletedAt.Valid && + time.Since(job.UpdatedAt) > 30*time.Second || + build.Transition != database.WorkspaceTransitionDelete { + continue + } + + case "deleting": + if !job.CompletedAt.Valid && + job.CanceledAt.Valid && + job.Error.Valid && + build.Transition != database.WorkspaceTransitionDelete { + continue + } + + default: + return 0, xerrors.Errorf("unknown workspace status in filter: %q", arg.Status) + } + } + + if len(arg.TemplateIds) > 0 { + match := false + for _, id := range arg.TemplateIds { + if workspace.TemplateID == id { + match = true + break + } + } + if !match { + continue + } + } + + // If the filter exists, ensure the object is authorized. + if authorizedFilter != nil && !authorizedFilter.Eval(workspace.RBACObject()) { + continue + } + workspaces = append(workspaces, workspace) + } + + return int64(len(workspaces)), nil +} + func (q *fakeQuerier) GetWorkspaceByID(_ context.Context, id uuid.UUID) (database.Workspace, error) { q.mutex.RLock() defer q.mutex.RUnlock() diff --git a/coderd/database/modelqueries.go b/coderd/database/modelqueries.go index 3383b6af96..827287de1b 100644 --- a/coderd/database/modelqueries.go +++ b/coderd/database/modelqueries.go @@ -159,6 +159,7 @@ func (q *sqlQuerier) GetTemplateGroupRoles(ctx context.Context, id uuid.UUID) ([ type workspaceQuerier interface { GetAuthorizedWorkspaces(ctx context.Context, arg GetWorkspacesParams, authorizedFilter rbac.AuthorizeFilter) ([]Workspace, error) + GetAuthorizedWorkspaceCount(ctx context.Context, arg GetWorkspaceCountParams, authorizedFilter rbac.AuthorizeFilter) (int64, error) } // GetAuthorizedWorkspaces returns all workspaces that the user is authorized to access. @@ -213,3 +214,53 @@ func (q *sqlQuerier) GetAuthorizedWorkspaces(ctx context.Context, arg GetWorkspa } return items, nil } + +func (q *sqlQuerier) GetAuthorizedWorkspaceCount(ctx context.Context, arg GetWorkspacesParams, authorizedFilter rbac.AuthorizeFilter) (int64, error) { + // In order to properly use ORDER BY, OFFSET, and LIMIT, we need to inject the + // authorizedFilter between the end of the where clause and those statements. + filter := strings.Replace(getWorkspaces, "-- @authorize_filter", fmt.Sprintf(" AND %s", authorizedFilter.SQLString(rbac.NoACLConfig())), 1) + // The name comment is for metric tracking + query := fmt.Sprintf("-- name: GetAuthorizedWorkspaces :many\n%s", filter) + rows, err := q.db.QueryContext(ctx, query, + arg.Deleted, + arg.Status, + arg.OwnerID, + arg.OwnerUsername, + arg.TemplateName, + pq.Array(arg.TemplateIds), + arg.Name, + arg.Offset, + arg.Limit, + ) + if err != nil { + return 0, xerrors.Errorf("get authorized workspaces: %w", err) + } + defer rows.Close() + var items []Workspace + for rows.Next() { + var i Workspace + if err := rows.Scan( + &i.ID, + &i.CreatedAt, + &i.UpdatedAt, + &i.OwnerID, + &i.OrganizationID, + &i.TemplateID, + &i.Deleted, + &i.Name, + &i.AutostartSchedule, + &i.Ttl, + &i.LastUsedAt, + ); err != nil { + return 0, err + } + items = append(items, i) + } + if err := rows.Close(); err != nil { + return 0, err + } + if err := rows.Err(); err != nil { + return 0, err + } + return int64(len(items)), nil +}