refactor: consolidate template and workspace acl validation (#19192)

This commit is contained in:
ケイラ
2025-08-07 10:14:58 -06:00
committed by GitHub
parent 02de067d46
commit 26458cd6f0
14 changed files with 535 additions and 103 deletions
+2 -1
View File
@@ -1415,7 +1415,8 @@ func New(options *Options) *API {
r.Get("/timings", api.workspaceTimings)
r.Route("/acl", func(r chi.Router) {
r.Use(
httpmw.RequireExperiment(api.Experiments, codersdk.ExperimentWorkspaceSharing))
httpmw.RequireExperiment(api.Experiments, codersdk.ExperimentWorkspaceSharing),
)
r.Patch("/", api.patchWorkspaceACL)
})
+20
View File
@@ -5376,6 +5376,26 @@ func (q *querier) UpsertWorkspaceAppAuditSession(ctx context.Context, arg databa
return q.db.UpsertWorkspaceAppAuditSession(ctx, arg)
}
func (q *querier) ValidateGroupIDs(ctx context.Context, groupIDs []uuid.UUID) (database.ValidateGroupIDsRow, error) {
// This check is probably overly restrictive, but the "correct" check isn't
// necessarily obvious. It's only used as a verification check for ACLs right
// now, which are performed as system.
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil {
return database.ValidateGroupIDsRow{}, err
}
return q.db.ValidateGroupIDs(ctx, groupIDs)
}
func (q *querier) ValidateUserIDs(ctx context.Context, userIDs []uuid.UUID) (database.ValidateUserIDsRow, error) {
// This check is probably overly restrictive, but the "correct" check isn't
// necessarily obvious. It's only used as a verification check for ACLs right
// now, which are performed as system.
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil {
return database.ValidateUserIDsRow{}, err
}
return q.db.ValidateUserIDs(ctx, userIDs)
}
func (q *querier) GetAuthorizedTemplates(ctx context.Context, arg database.GetTemplatesWithFilterParams, _ rbac.PreparedAuthorized) ([]database.Template, error) {
// TODO Delete this function, all GetTemplates should be authorized. For now just call getTemplates on the authz querier.
return q.GetTemplatesWithFilter(ctx, arg)
+9
View File
@@ -623,6 +623,11 @@ func (s *MethodTestSuite) TestGroup() {
ID: g.ID,
}).Asserts(g, policy.ActionUpdate)
}))
s.Run("ValidateGroupIDs", s.Subtest(func(db database.Store, check *expects) {
o := dbgen.Organization(s.T(), db, database.Organization{})
g := dbgen.Group(s.T(), db, database.Group{OrganizationID: o.ID})
check.Args([]uuid.UUID{g.ID}).Asserts(rbac.ResourceSystem, policy.ActionRead)
}))
}
func (s *MethodTestSuite) TestProvisionerJob() {
@@ -2077,6 +2082,10 @@ func (s *MethodTestSuite) TestUser() {
Interval: int32((time.Hour * 24).Seconds()),
}).Asserts(rbac.ResourceUser, policy.ActionRead)
}))
s.Run("ValidateUserIDs", s.Subtest(func(db database.Store, check *expects) {
u := dbgen.User(s.T(), db, database.User{})
check.Args([]uuid.UUID{u.ID}).Asserts(rbac.ResourceSystem, policy.ActionRead)
}))
}
func (s *MethodTestSuite) TestWorkspace() {
+14
View File
@@ -3372,6 +3372,20 @@ func (m queryMetricsStore) UpsertWorkspaceAppAuditSession(ctx context.Context, a
return r0, r1
}
func (m queryMetricsStore) ValidateGroupIDs(ctx context.Context, groupIds []uuid.UUID) (database.ValidateGroupIDsRow, error) {
start := time.Now()
r0, r1 := m.s.ValidateGroupIDs(ctx, groupIds)
m.queryLatencies.WithLabelValues("ValidateGroupIDs").Observe(time.Since(start).Seconds())
return r0, r1
}
func (m queryMetricsStore) ValidateUserIDs(ctx context.Context, userIds []uuid.UUID) (database.ValidateUserIDsRow, error) {
start := time.Now()
r0, r1 := m.s.ValidateUserIDs(ctx, userIds)
m.queryLatencies.WithLabelValues("ValidateUserIDs").Observe(time.Since(start).Seconds())
return r0, r1
}
func (m queryMetricsStore) GetAuthorizedTemplates(ctx context.Context, arg database.GetTemplatesWithFilterParams, prepared rbac.PreparedAuthorized) ([]database.Template, error) {
start := time.Now()
templates, err := m.s.GetAuthorizedTemplates(ctx, arg, prepared)
+30
View File
@@ -7159,6 +7159,36 @@ func (mr *MockStoreMockRecorder) UpsertWorkspaceAppAuditSession(ctx, arg any) *g
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertWorkspaceAppAuditSession", reflect.TypeOf((*MockStore)(nil).UpsertWorkspaceAppAuditSession), ctx, arg)
}
// ValidateGroupIDs mocks base method.
func (m *MockStore) ValidateGroupIDs(ctx context.Context, groupIds []uuid.UUID) (database.ValidateGroupIDsRow, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ValidateGroupIDs", ctx, groupIds)
ret0, _ := ret[0].(database.ValidateGroupIDsRow)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// ValidateGroupIDs indicates an expected call of ValidateGroupIDs.
func (mr *MockStoreMockRecorder) ValidateGroupIDs(ctx, groupIds any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateGroupIDs", reflect.TypeOf((*MockStore)(nil).ValidateGroupIDs), ctx, groupIds)
}
// ValidateUserIDs mocks base method.
func (m *MockStore) ValidateUserIDs(ctx context.Context, userIds []uuid.UUID) (database.ValidateUserIDsRow, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ValidateUserIDs", ctx, userIds)
ret0, _ := ret[0].(database.ValidateUserIDsRow)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// ValidateUserIDs indicates an expected call of ValidateUserIDs.
func (mr *MockStoreMockRecorder) ValidateUserIDs(ctx, userIds any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateUserIDs", reflect.TypeOf((*MockStore)(nil).ValidateUserIDs), ctx, userIds)
}
// Wrappers mocks base method.
func (m *MockStore) Wrappers() []string {
m.ctrl.T.Helper()
+2
View File
@@ -689,6 +689,8 @@ type sqlcQuerier interface {
// was started. This means that a new row was inserted (no previous session) or
// the updated_at is older than stale interval.
UpsertWorkspaceAppAuditSession(ctx context.Context, arg UpsertWorkspaceAppAuditSessionParams) (bool, error)
ValidateGroupIDs(ctx context.Context, groupIds []uuid.UUID) (ValidateGroupIDsRow, error)
ValidateUserIDs(ctx context.Context, userIds []uuid.UUID) (ValidateUserIDsRow, error)
}
var _ sqlcQuerier = (*sqlQuerier)(nil)
+64
View File
@@ -2869,6 +2869,37 @@ func (q *sqlQuerier) UpdateGroupByID(ctx context.Context, arg UpdateGroupByIDPar
return i, err
}
const validateGroupIDs = `-- name: ValidateGroupIDs :one
WITH input AS (
SELECT
unnest($1::uuid[]) AS id
)
SELECT
array_agg(input.id)::uuid[] as invalid_group_ids,
COUNT(*) = 0 as ok
FROM
-- Preserve rows where there is not a matching left (groups) row for each
-- right (input) row...
groups
RIGHT JOIN input ON groups.id = input.id
WHERE
-- ...so that we can retain exactly those rows where an input ID does not
-- match an existing group.
groups.id IS NULL
`
type ValidateGroupIDsRow struct {
InvalidGroupIds []uuid.UUID `db:"invalid_group_ids" json:"invalid_group_ids"`
Ok bool `db:"ok" json:"ok"`
}
func (q *sqlQuerier) ValidateGroupIDs(ctx context.Context, groupIds []uuid.UUID) (ValidateGroupIDsRow, error) {
row := q.db.QueryRowContext(ctx, validateGroupIDs, pq.Array(groupIds))
var i ValidateGroupIDsRow
err := row.Scan(pq.Array(&i.InvalidGroupIds), &i.Ok)
return i, err
}
const getTemplateAppInsights = `-- name: GetTemplateAppInsights :many
WITH
-- Create a list of all unique apps by template, this is used to
@@ -14792,6 +14823,39 @@ func (q *sqlQuerier) UpdateUserThemePreference(ctx context.Context, arg UpdateUs
return i, err
}
const validateUserIDs = `-- name: ValidateUserIDs :one
WITH input AS (
SELECT
unnest($1::uuid[]) AS id
)
SELECT
array_agg(input.id)::uuid[] as invalid_user_ids,
COUNT(*) = 0 as ok
FROM
-- Preserve rows where there is not a matching left (users) row for each
-- right (input) row...
users
RIGHT JOIN input ON users.id = input.id
WHERE
-- ...so that we can retain exactly those rows where an input ID does not
-- match an existing user...
users.id IS NULL OR
-- ...or that only matches a user that was deleted.
users.deleted = true
`
type ValidateUserIDsRow struct {
InvalidUserIds []uuid.UUID `db:"invalid_user_ids" json:"invalid_user_ids"`
Ok bool `db:"ok" json:"ok"`
}
func (q *sqlQuerier) ValidateUserIDs(ctx context.Context, userIds []uuid.UUID) (ValidateUserIDsRow, error) {
row := q.db.QueryRowContext(ctx, validateUserIDs, pq.Array(userIds))
var i ValidateUserIDsRow
err := row.Scan(pq.Array(&i.InvalidUserIds), &i.Ok)
return i, err
}
const getWorkspaceAgentDevcontainersByAgentID = `-- name: GetWorkspaceAgentDevcontainersByAgentID :many
SELECT
id, workspace_agent_id, created_at, workspace_folder, config_path, name
+18
View File
@@ -8,6 +8,24 @@ WHERE
LIMIT
1;
-- name: ValidateGroupIDs :one
WITH input AS (
SELECT
unnest(@group_ids::uuid[]) AS id
)
SELECT
array_agg(input.id)::uuid[] as invalid_group_ids,
COUNT(*) = 0 as ok
FROM
-- Preserve rows where there is not a matching left (groups) row for each
-- right (input) row...
groups
RIGHT JOIN input ON groups.id = input.id
WHERE
-- ...so that we can retain exactly those rows where an input ID does not
-- match an existing group.
groups.id IS NULL;
-- name: GetGroupByOrgAndName :one
SELECT
*
+20
View File
@@ -25,6 +25,26 @@ WHERE
LIMIT
1;
-- name: ValidateUserIDs :one
WITH input AS (
SELECT
unnest(@user_ids::uuid[]) AS id
)
SELECT
array_agg(input.id)::uuid[] as invalid_user_ids,
COUNT(*) = 0 as ok
FROM
-- Preserve rows where there is not a matching left (users) row for each
-- right (input) row...
users
RIGHT JOIN input ON users.id = input.id
WHERE
-- ...so that we can retain exactly those rows where an input ID does not
-- match an existing user...
users.id IS NULL OR
-- ...or that only matches a user that was deleted.
users.deleted = true;
-- name: GetUsersByIDs :many
-- This shouldn't check for deleted, because it's frequently used
-- to look up references to actions. eg. a user could build a workspace
+130
View File
@@ -0,0 +1,130 @@
package acl
import (
"context"
"fmt"
"github.com/google/uuid"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
"github.com/coder/coder/v2/codersdk"
)
type UpdateValidator[Role codersdk.WorkspaceRole | codersdk.TemplateRole] interface {
// Users should return a map from user UUIDs (as strings) to the role they
// are being assigned. Additionally, it should return a string that will be
// used as the field name for the ValidationErrors returned from Validate.
Users() (map[string]Role, string)
// Groups should return a map from group UUIDs (as strings) to the role they
// are being assigned. Additionally, it should return a string that will be
// used as the field name for the ValidationErrors returned from Validate.
Groups() (map[string]Role, string)
// ValidateRole should return an error that will be used in the
// ValidationError if the role is invalid for the corresponding resource type.
ValidateRole(role Role) error
}
func Validate[Role codersdk.WorkspaceRole | codersdk.TemplateRole](
ctx context.Context,
db database.Store,
v UpdateValidator[Role],
) []codersdk.ValidationError {
// nolint:gocritic // Validate requires full read access to users and groups
ctx = dbauthz.AsSystemRestricted(ctx)
var validErrs []codersdk.ValidationError
groupRoles, groupsField := v.Groups()
groupIDs := make([]uuid.UUID, 0, len(groupRoles))
for idStr, role := range groupRoles {
// Validate the provided role names
if err := v.ValidateRole(role); err != nil {
validErrs = append(validErrs, codersdk.ValidationError{
Field: groupsField,
Detail: err.Error(),
})
}
// Validate that the IDs are UUIDs
id, err := uuid.Parse(idStr)
if err != nil {
validErrs = append(validErrs, codersdk.ValidationError{
Field: groupsField,
Detail: fmt.Sprintf("%v is not a valid UUID.", idStr),
})
continue
}
// Don't check if the ID exists when setting the role to
// WorkspaceRoleDeleted or TemplateRoleDeleted. They might've existing at
// some point and got deleted. If we report that as an error here then they
// can't be removed.
if string(role) == "" {
continue
}
groupIDs = append(groupIDs, id)
}
// Validate that the groups exist
groupValidation, err := db.ValidateGroupIDs(ctx, groupIDs)
if err != nil {
validErrs = append(validErrs, codersdk.ValidationError{
Field: groupsField,
Detail: fmt.Sprintf("failed to validate group IDs: %v", err.Error()),
})
}
if !groupValidation.Ok {
for _, id := range groupValidation.InvalidGroupIds {
validErrs = append(validErrs, codersdk.ValidationError{
Field: groupsField,
Detail: fmt.Sprintf("group with ID %v does not exist", id),
})
}
}
userRoles, usersField := v.Users()
userIDs := make([]uuid.UUID, 0, len(userRoles))
for idStr, role := range userRoles {
// Validate the provided role names
if err := v.ValidateRole(role); err != nil {
validErrs = append(validErrs, codersdk.ValidationError{
Field: usersField,
Detail: err.Error(),
})
}
// Validate that the IDs are UUIDs
id, err := uuid.Parse(idStr)
if err != nil {
validErrs = append(validErrs, codersdk.ValidationError{
Field: usersField,
Detail: fmt.Sprintf("%v is not a valid UUID.", idStr),
})
continue
}
// Don't check if the ID exists when setting the role to
// WorkspaceRoleDeleted or TemplateRoleDeleted. They might've existing at
// some point and got deleted. If we report that as an error here then they
// can't be removed.
if string(role) == "" {
continue
}
userIDs = append(userIDs, id)
}
// Validate that the groups exist
userValidation, err := db.ValidateUserIDs(ctx, userIDs)
if err != nil {
validErrs = append(validErrs, codersdk.ValidationError{
Field: usersField,
Detail: fmt.Sprintf("failed to validate user IDs: %v", err.Error()),
})
}
if !userValidation.Ok {
for _, id := range userValidation.InvalidUserIds {
validErrs = append(validErrs, codersdk.ValidationError{
Field: usersField,
Detail: fmt.Sprintf("user with ID %v does not exist", id),
})
}
}
return validErrs
}
+91
View File
@@ -0,0 +1,91 @@
package acl_test
import (
"testing"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/coderd"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/rbac/acl"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
)
func TestOK(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
o := dbgen.Organization(t, db, database.Organization{})
g := dbgen.Group(t, db, database.Group{OrganizationID: o.ID})
u := dbgen.User(t, db, database.User{})
ctx := testutil.Context(t, testutil.WaitShort)
update := codersdk.UpdateWorkspaceACL{
UserRoles: map[string]codersdk.WorkspaceRole{
u.ID.String(): codersdk.WorkspaceRoleAdmin,
// An unknown ID is allowed if and only if the specified role is either
// codersdk.WorkspaceRoleDeleted or codersdk.TemplateRoleDeleted.
uuid.NewString(): codersdk.WorkspaceRoleDeleted,
},
GroupRoles: map[string]codersdk.WorkspaceRole{
g.ID.String(): codersdk.WorkspaceRoleAdmin,
// An unknown ID is allowed if and only if the specified role is either
// codersdk.WorkspaceRoleDeleted or codersdk.TemplateRoleDeleted.
uuid.NewString(): codersdk.WorkspaceRoleDeleted,
},
}
errors := acl.Validate(ctx, db, coderd.WorkspaceACLUpdateValidator(update))
require.Empty(t, errors)
}
func TestDeniesUnknownIDs(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
update := codersdk.UpdateWorkspaceACL{
UserRoles: map[string]codersdk.WorkspaceRole{
uuid.NewString(): codersdk.WorkspaceRoleAdmin,
},
GroupRoles: map[string]codersdk.WorkspaceRole{
uuid.NewString(): codersdk.WorkspaceRoleAdmin,
},
}
errors := acl.Validate(ctx, db, coderd.WorkspaceACLUpdateValidator(update))
require.Len(t, errors, 2)
require.Equal(t, errors[0].Field, "group_roles")
require.ErrorContains(t, errors[0], "does not exist")
require.Equal(t, errors[1].Field, "user_roles")
require.ErrorContains(t, errors[1], "does not exist")
}
func TestDeniesUnknownRolesAndInvalidIDs(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
update := codersdk.UpdateWorkspaceACL{
UserRoles: map[string]codersdk.WorkspaceRole{
"Quifrey": "level 5",
},
GroupRoles: map[string]codersdk.WorkspaceRole{
"apprentices": "level 2",
},
}
errors := acl.Validate(ctx, db, coderd.WorkspaceACLUpdateValidator(update))
require.Len(t, errors, 4)
require.Equal(t, errors[0].Field, "group_roles")
require.ErrorContains(t, errors[0], "role \"level 2\" is not a valid workspace role")
require.Equal(t, errors[1].Field, "group_roles")
require.ErrorContains(t, errors[1], "not a valid UUID")
require.Equal(t, errors[2].Field, "user_roles")
require.ErrorContains(t, errors[2], "role \"level 5\" is not a valid workspace role")
require.Equal(t, errors[3].Field, "user_roles")
require.ErrorContains(t, errors[3], "not a valid UUID")
}
+18 -46
View File
@@ -32,6 +32,7 @@ import (
"github.com/coder/coder/v2/coderd/notifications"
"github.com/coder/coder/v2/coderd/prebuilds"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/rbac/acl"
"github.com/coder/coder/v2/coderd/rbac/policy"
"github.com/coder/coder/v2/coderd/schedule"
"github.com/coder/coder/v2/coderd/schedule/cron"
@@ -2086,17 +2087,10 @@ func (api *API) patchWorkspaceACL(rw http.ResponseWriter, r *http.Request) {
return
}
validErrs := validateWorkspaceACLPerms(ctx, api.Database, req.UserRoles, "user_roles")
validErrs = append(validErrs, validateWorkspaceACLPerms(
ctx,
api.Database,
req.GroupRoles,
"group_roles",
)...)
validErrs := acl.Validate(ctx, api.Database, WorkspaceACLUpdateValidator(req))
if len(validErrs) > 0 {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid request to update template metadata!",
Message: "Invalid request to update workspace ACL",
Validations: validErrs,
})
return
@@ -2492,50 +2486,28 @@ func (api *API) publishWorkspaceAgentLogsUpdate(ctx context.Context, workspaceAg
}
}
func validateWorkspaceACLPerms(ctx context.Context, db database.Store, perms map[string]codersdk.WorkspaceRole, field string) []codersdk.ValidationError {
// nolint:gocritic // Validate requires full read access to users and groups
ctx = dbauthz.AsSystemRestricted(ctx)
var validErrs []codersdk.ValidationError
for idStr, role := range perms {
if err := validateWorkspaceRole(role); err != nil {
validErrs = append(validErrs, codersdk.ValidationError{Field: field, Detail: err.Error()})
continue
}
type WorkspaceACLUpdateValidator codersdk.UpdateWorkspaceACL
id, err := uuid.Parse(idStr)
if err != nil {
validErrs = append(validErrs, codersdk.ValidationError{Field: field, Detail: idStr + "is not a valid UUID."})
continue
}
var (
workspaceACLUpdateUsersFieldName = "user_roles"
workspaceACLUpdateGroupsFieldName = "group_roles"
)
switch field {
case "user_roles":
// TODO(lilac): put this back after Kirby button shenanigans are over
// This could get slow if we get a ton of user perm updates.
// _, err = db.GetUserByID(ctx, id)
// if err != nil {
// validErrs = append(validErrs, codersdk.ValidationError{Field: field, Detail: fmt.Sprintf("Failed to find resource with ID %q: %v", idStr, err.Error())})
// continue
// }
case "group_roles":
// This could get slow if we get a ton of group perm updates.
_, err = db.GetGroupByID(ctx, id)
if err != nil {
validErrs = append(validErrs, codersdk.ValidationError{Field: field, Detail: fmt.Sprintf("Failed to find resource with ID %q: %v", idStr, err.Error())})
continue
}
default:
validErrs = append(validErrs, codersdk.ValidationError{Field: field, Detail: "invalid field"})
}
}
// WorkspaceACLUpdateValidator implements acl.UpdateValidator[codersdk.WorkspaceRole]
var _ acl.UpdateValidator[codersdk.WorkspaceRole] = WorkspaceACLUpdateValidator{}
return validErrs
func (w WorkspaceACLUpdateValidator) Users() (map[string]codersdk.WorkspaceRole, string) {
return w.UserRoles, workspaceACLUpdateUsersFieldName
}
func validateWorkspaceRole(role codersdk.WorkspaceRole) error {
func (w WorkspaceACLUpdateValidator) Groups() (map[string]codersdk.WorkspaceRole, string) {
return w.GroupRoles, workspaceACLUpdateGroupsFieldName
}
func (WorkspaceACLUpdateValidator) ValidateRole(role codersdk.WorkspaceRole) error {
actions := db2sdk.WorkspaceRoleActions(role)
if len(actions) == 0 && role != codersdk.WorkspaceRoleDeleted {
return xerrors.Errorf("role %q is not a valid Workspace role", role)
return xerrors.Errorf("role %q is not a valid workspace role", role)
}
return nil