diff --git a/coderd/apidoc/docs.go b/coderd/apidoc/docs.go index 1c1f1946b91..1f0454cabdd 100644 --- a/coderd/apidoc/docs.go +++ b/coderd/apidoc/docs.go @@ -20125,6 +20125,7 @@ const docTemplate = `{ "workspace-build-updates", "nats_pubsub", "minimum-implicit-member", + "workspace-capable-licensing", "ai-gateway-cost-control", "chat-advisor", "chat-virtual-desktop" @@ -20141,6 +20142,7 @@ const docTemplate = `{ "ExperimentNotifications": "Sends notifications via SMTP and webhooks following certain events.", "ExperimentOAuth2": "Enables OAuth2 provider functionality.", "ExperimentWorkspaceBuildUpdates": "Enables publishing workspace build updates to the all builds pubsub channel.", + "ExperimentWorkspaceCapableLicensing": "Counts only users holding the workspace-create permission toward the license seat limit.", "ExperimentWorkspaceUsage": "Enables the new workspace usage tracking." }, "x-enum-descriptions": [ @@ -20153,6 +20155,7 @@ const docTemplate = `{ "Enables publishing workspace build updates to the all builds pubsub channel.", "Enables embedded NATS pubsub.", "Allows organizations to deviate from the default organization-member roles, in support of Gateway Accounts.", + "Counts only users holding the workspace-create permission toward the license seat limit.", "Enables AI Gateway cost control functionality.", "Enables the advisor tool for root agent chats.", "Enables virtual desktop and computer use provider for agents." @@ -20167,6 +20170,7 @@ const docTemplate = `{ "ExperimentWorkspaceBuildUpdates", "ExperimentNATSPubsub", "ExperimentMinimumImplicitMember", + "ExperimentWorkspaceCapableLicensing", "ExperimentAIGatewayCostControl", "ExperimentChatAdvisor", "ExperimentChatVirtualDesktop" diff --git a/coderd/apidoc/swagger.json b/coderd/apidoc/swagger.json index 78a0c965ebf..adc9f99cafb 100644 --- a/coderd/apidoc/swagger.json +++ b/coderd/apidoc/swagger.json @@ -18277,6 +18277,7 @@ "workspace-build-updates", "nats_pubsub", "minimum-implicit-member", + "workspace-capable-licensing", "ai-gateway-cost-control", "chat-advisor", "chat-virtual-desktop" @@ -18293,6 +18294,7 @@ "ExperimentNotifications": "Sends notifications via SMTP and webhooks following certain events.", "ExperimentOAuth2": "Enables OAuth2 provider functionality.", "ExperimentWorkspaceBuildUpdates": "Enables publishing workspace build updates to the all builds pubsub channel.", + "ExperimentWorkspaceCapableLicensing": "Counts only users holding the workspace-create permission toward the license seat limit.", "ExperimentWorkspaceUsage": "Enables the new workspace usage tracking." }, "x-enum-descriptions": [ @@ -18305,6 +18307,7 @@ "Enables publishing workspace build updates to the all builds pubsub channel.", "Enables embedded NATS pubsub.", "Allows organizations to deviate from the default organization-member roles, in support of Gateway Accounts.", + "Counts only users holding the workspace-create permission toward the license seat limit.", "Enables AI Gateway cost control functionality.", "Enables the advisor tool for root agent chats.", "Enables virtual desktop and computer use provider for agents." @@ -18319,6 +18322,7 @@ "ExperimentWorkspaceBuildUpdates", "ExperimentNATSPubsub", "ExperimentMinimumImplicitMember", + "ExperimentWorkspaceCapableLicensing", "ExperimentAIGatewayCostControl", "ExperimentChatAdvisor", "ExperimentChatVirtualDesktop" diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 6d24264d470..861083e8ff4 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -2955,6 +2955,13 @@ func (q *querier) GetActiveUserCount(ctx context.Context, includeSystem bool) (i return q.db.GetActiveUserCount(ctx, includeSystem) } +func (q *querier) GetActiveUsersAuthorizationRoles(ctx context.Context) ([]database.GetActiveUsersAuthorizationRolesRow, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { + return nil, err + } + return q.db.GetActiveUsersAuthorizationRoles(ctx) +} + func (q *querier) GetActiveWorkspaceBuildsByTemplateID(ctx context.Context, templateID uuid.UUID) ([]database.WorkspaceBuild, error) { // This is a system-only function. if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 8f98fd7d7bc..a212ae933d6 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -5038,6 +5038,10 @@ func (s *MethodTestSuite) TestSystemFunctions() { dbm.EXPECT().GetActiveUserCount(gomock.Any(), false).Return(int64(0), nil).AnyTimes() check.Args(false).Asserts(rbac.ResourceSystem, policy.ActionRead).Returns(int64(0)) })) + s.Run("GetActiveUsersAuthorizationRoles", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + dbm.EXPECT().GetActiveUsersAuthorizationRoles(gomock.Any()).Return([]database.GetActiveUsersAuthorizationRolesRow{}, nil).AnyTimes() + check.Args().Asserts(rbac.ResourceSystem, policy.ActionRead).Returns([]database.GetActiveUsersAuthorizationRolesRow{}) + })) s.Run("GetAuthorizationUserRoles", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { u := testutil.Fake(s.T(), faker, database.User{}) dbm.EXPECT().GetAuthorizationUserRoles(gomock.Any(), u.ID).Return(database.GetAuthorizationUserRolesRow{}, nil).AnyTimes() diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 07d33217470..119764c5816 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1281,6 +1281,14 @@ func (m queryMetricsStore) GetActiveUserCount(ctx context.Context, includeSystem return r0, r1 } +func (m queryMetricsStore) GetActiveUsersAuthorizationRoles(ctx context.Context) ([]database.GetActiveUsersAuthorizationRolesRow, error) { + start := time.Now() + r0, r1 := m.s.GetActiveUsersAuthorizationRoles(ctx) + m.queryLatencies.WithLabelValues("GetActiveUsersAuthorizationRoles").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetActiveUsersAuthorizationRoles").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetActiveWorkspaceBuildsByTemplateID(ctx context.Context, templateID uuid.UUID) ([]database.WorkspaceBuild, error) { start := time.Now() r0, r1 := m.s.GetActiveWorkspaceBuildsByTemplateID(ctx, templateID) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 9a255e2df74..b16a20f91f0 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -2233,6 +2233,21 @@ func (mr *MockStoreMockRecorder) GetActiveUserCount(ctx, includeSystem any) *gom return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveUserCount", reflect.TypeOf((*MockStore)(nil).GetActiveUserCount), ctx, includeSystem) } +// GetActiveUsersAuthorizationRoles mocks base method. +func (m *MockStore) GetActiveUsersAuthorizationRoles(ctx context.Context) ([]database.GetActiveUsersAuthorizationRolesRow, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetActiveUsersAuthorizationRoles", ctx) + ret0, _ := ret[0].([]database.GetActiveUsersAuthorizationRolesRow) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetActiveUsersAuthorizationRoles indicates an expected call of GetActiveUsersAuthorizationRoles. +func (mr *MockStoreMockRecorder) GetActiveUsersAuthorizationRoles(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveUsersAuthorizationRoles", reflect.TypeOf((*MockStore)(nil).GetActiveUsersAuthorizationRoles), ctx) +} + // GetActiveWorkspaceBuildsByTemplateID mocks base method. func (m *MockStore) GetActiveWorkspaceBuildsByTemplateID(ctx context.Context, templateID uuid.UUID) ([]database.WorkspaceBuild, error) { m.ctrl.T.Helper() diff --git a/coderd/database/modelmethods.go b/coderd/database/modelmethods.go index 76ca27166ba..3576c98276a 100644 --- a/coderd/database/modelmethods.go +++ b/coderd/database/modelmethods.go @@ -876,6 +876,18 @@ func (r GetAuthorizationUserRolesRow) RoleNames() ([]rbac.RoleIdentifier, error) return names, nil } +func (r GetActiveUsersAuthorizationRolesRow) RoleNames() ([]rbac.RoleIdentifier, error) { + names := make([]rbac.RoleIdentifier, 0, len(r.Roles)) + for _, role := range r.Roles { + value, err := rbac.RoleNameFromString(role) + if err != nil { + return nil, xerrors.Errorf("convert role %q: %w", role, err) + } + names = append(names, value) + } + return names, nil +} + func (k CryptoKey) ExpiresAt(keyDuration time.Duration) time.Time { return k.StartsAt.Add(keyDuration).UTC() } diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 7579a25c7dc..92e1438d37c 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -341,6 +341,13 @@ type sqlcQuerier interface { GetActiveChatsByAgentID(ctx context.Context, agentID uuid.UUID) ([]Chat, error) GetActivePresetPrebuildSchedules(ctx context.Context) ([]TemplateVersionPresetPrebuildSchedule, error) GetActiveUserCount(ctx context.Context, includeSystem bool) (int64, error) + // Returns the authorization roles (site and org-scoped, including implied + // member roles and organization default roles) and the group memberships + // for every active, non-deleted user who is neither a system user nor a + // service account, matching the GetActiveUserCount population. + // Must stay semantically in sync with GetAuthorizationUserRoles; + // TestGetActiveUsersAuthorizationRolesParity enforces this. + GetActiveUsersAuthorizationRoles(ctx context.Context) ([]GetActiveUsersAuthorizationRolesRow, error) GetActiveWorkspaceBuildsByTemplateID(ctx context.Context, templateID uuid.UUID) ([]WorkspaceBuild, error) // For PG Coordinator HTMLDebug GetAllTailnetCoordinators(ctx context.Context) ([]TailnetCoordinator, error) @@ -364,6 +371,9 @@ type sqlcQuerier interface { GetAuthenticatedWorkspaceAgentAndBuildByAuthToken(ctx context.Context, authToken uuid.UUID) (GetAuthenticatedWorkspaceAgentAndBuildByAuthTokenRow, error) // This function returns roles for authorization purposes. Implied member roles // are included. + // Must stay semantically in sync with GetActiveUsersAuthorizationRoles + // (implied member roles, org default roles, groups); + // TestGetActiveUsersAuthorizationRolesParity enforces this. GetAuthorizationUserRoles(ctx context.Context, userID uuid.UUID) (GetAuthorizationUserRolesRow, error) // Returns read-only root chat candidates for state-machine-backed // auto-archive. Activity is computed across the root family. The query diff --git a/coderd/database/querier_test.go b/coderd/database/querier_test.go index 305cf43a1cb..dc5e7afecde 100644 --- a/coderd/database/querier_test.go +++ b/coderd/database/querier_test.go @@ -17723,3 +17723,92 @@ func requireAIGatewayKeysViolation( require.FailNow(t, "test case must expect a constraint error") } } + +// TestGetActiveUsersAuthorizationRolesParity verifies that the bulk +// GetActiveUsersAuthorizationRoles query returns, for every eligible +// user, the same roles and groups as the per-user +// GetAuthorizationUserRoles query. The two queries encode the implied +// member roles, organization default roles, and group memberships +// independently and must not drift. +func TestGetActiveUsersAuthorizationRolesParity(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitLong) + + orgA := dbgen.Organization(t, db, database.Organization{}) + orgB := dbgen.Organization(t, db, database.Organization{}) + + activeUser := func(seed database.User) database.User { + seed.Status = database.UserStatusActive + return dbgen.User(t, db, seed) + } + member := func(orgID uuid.UUID, user database.User, roles ...string) { + dbgen.OrganizationMember(t, db, database.OrganizationMember{ + OrganizationID: orgID, + UserID: user.ID, + Roles: roles, + }) + } + + // Site-wide role, zero org memberships. + owner := activeUser(database.User{RBACRoles: []string{rbac.RoleOwner().Name}}) + + // Plain single-org member; effective roles come from the implied + // member role plus the org's default member roles. + plain := activeUser(database.User{}) + member(orgA.ID, plain) + + // Explicit org roles across two organizations. + multiOrg := activeUser(database.User{}) + member(orgA.ID, multiOrg, rbac.RoleOrgAdmin()) + member(orgB.ID, multiOrg) + + // Custom org role. + customRole, err := db.InsertCustomRole(ctx, database.InsertCustomRoleParams{ + Name: "parity-role", + DisplayName: "Parity Role", + OrganizationID: uuid.NullUUID{UUID: orgA.ID, Valid: true}, + OrgPermissions: []database.CustomRolePermission{{ + ResourceType: rbac.ResourceWorkspace.Type, + Action: policy.ActionCreate, + }}, + }) + require.NoError(t, err) + custom := activeUser(database.User{}) + member(orgA.ID, custom, customRole.Name) + + // Group memberships. + grouped := activeUser(database.User{}) + member(orgA.ID, grouped) + for range 2 { + group := dbgen.Group(t, db, database.Group{OrganizationID: orgA.ID}) + dbgen.GroupMember(t, db, database.GroupMemberTable{ + UserID: grouped.ID, + GroupID: group.ID, + }) + } + + // Excluded from the bulk query: service accounts and non-active + // users. + sa := activeUser(database.User{IsServiceAccount: true}) + member(orgA.ID, sa) + suspended := dbgen.User(t, db, database.User{Status: database.UserStatusSuspended}) + member(orgA.ID, suspended) + + rows, err := db.GetActiveUsersAuthorizationRoles(ctx) + require.NoError(t, err) + + gotIDs := make([]uuid.UUID, 0, len(rows)) + for _, row := range rows { + gotIDs = append(gotIDs, row.ID) + } + require.ElementsMatch(t, []uuid.UUID{owner.ID, plain.ID, multiOrg.ID, custom.ID, grouped.ID}, gotIDs) + + for _, row := range rows { + single, err := db.GetAuthorizationUserRoles(ctx, row.ID) + require.NoError(t, err) + require.ElementsMatch(t, single.Roles, row.Roles, "roles diverged for user %s", row.ID) + require.ElementsMatch(t, single.Groups, row.Groups, "groups diverged for user %s", row.ID) + } +} diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index f4ae41eb3c2..70ffa52d094 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -30175,6 +30175,97 @@ func (q *sqlQuerier) GetActiveUserCount(ctx context.Context, includeSystem bool) return count, err } +const getActiveUsersAuthorizationRoles = `-- name: GetActiveUsersAuthorizationRoles :many +WITH org_roles AS ( + SELECT + organization_members.user_id, + -- The roles are returned as a flat array, org scoped and site side. + -- Concatenating the organization id scopes the organization roles. + array_agg(org_role || ':' || organization_members.organization_id::text) AS roles + FROM + organization_members + JOIN organizations ON organizations.id = organization_members.organization_id, + -- All org members get an implied organization-member role for + -- their orgs. Memberships of service accounts are aggregated here + -- too, but their rows never survive the join against the outer + -- WHERE, so the organization-service-account case does not apply. + -- + -- organizations.default_org_member_roles applies to every member + -- but is not materialized on membership rows, so it is unioned in + -- here. + unnest( + array_cat( + array_append(organization_members.roles, 'organization-member'), + organizations.default_org_member_roles + ) + ) AS org_role + GROUP BY + organization_members.user_id +), +user_groups AS ( + SELECT + group_members.user_id, + array_agg(group_members.group_id :: text) AS groups + FROM + group_members + GROUP BY + group_members.user_id +) +SELECT + users.id, + array_cat( + -- All users are members + array_append(users.rbac_roles, 'member'), + -- Users with no org memberships have no org_roles row. + coalesce(org_roles.roles, ARRAY[]::text[]) + ) :: text[] AS roles, + coalesce(user_groups.groups, ARRAY[]::text[]) :: text[] AS groups +FROM + users + LEFT JOIN org_roles ON org_roles.user_id = users.id + LEFT JOIN user_groups ON user_groups.user_id = users.id +WHERE + users.status = 'active'::user_status + AND users.deleted = false + AND users.is_system = false + AND users.is_service_account = false +` + +type GetActiveUsersAuthorizationRolesRow struct { + ID uuid.UUID `db:"id" json:"id"` + Roles []string `db:"roles" json:"roles"` + Groups []string `db:"groups" json:"groups"` +} + +// Returns the authorization roles (site and org-scoped, including implied +// member roles and organization default roles) and the group memberships +// for every active, non-deleted user who is neither a system user nor a +// service account, matching the GetActiveUserCount population. +// Must stay semantically in sync with GetAuthorizationUserRoles; +// TestGetActiveUsersAuthorizationRolesParity enforces this. +func (q *sqlQuerier) GetActiveUsersAuthorizationRoles(ctx context.Context) ([]GetActiveUsersAuthorizationRolesRow, error) { + rows, err := q.db.QueryContext(ctx, getActiveUsersAuthorizationRoles) + if err != nil { + return nil, err + } + defer rows.Close() + var items []GetActiveUsersAuthorizationRolesRow + for rows.Next() { + var i GetActiveUsersAuthorizationRolesRow + if err := rows.Scan(&i.ID, pq.Array(&i.Roles), pq.Array(&i.Groups)); 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 getAuthorizationUserRoles = `-- name: GetAuthorizationUserRoles :one SELECT -- username and email are returned just to help for logging purposes @@ -30247,6 +30338,9 @@ type GetAuthorizationUserRolesRow struct { // This function returns roles for authorization purposes. Implied member roles // are included. +// Must stay semantically in sync with GetActiveUsersAuthorizationRoles +// (implied member roles, org default roles, groups); +// TestGetActiveUsersAuthorizationRolesParity enforces this. func (q *sqlQuerier) GetAuthorizationUserRoles(ctx context.Context, userID uuid.UUID) (GetAuthorizationUserRolesRow, error) { row := q.db.QueryRowContext(ctx, getAuthorizationUserRoles, userID) var i GetAuthorizationUserRolesRow diff --git a/coderd/database/queries/users.sql b/coderd/database/queries/users.sql index 3c79b405225..e283b43c6be 100644 --- a/coderd/database/queries/users.sql +++ b/coderd/database/queries/users.sql @@ -594,6 +594,9 @@ WHERE -- name: GetAuthorizationUserRoles :one -- This function returns roles for authorization purposes. Implied member roles -- are included. +-- Must stay semantically in sync with GetActiveUsersAuthorizationRoles +-- (implied member roles, org default roles, groups); +-- TestGetActiveUsersAuthorizationRolesParity enforces this. SELECT -- username and email are returned just to help for logging purposes -- status is used to enforce 'suspended' users, as all roles are ignored @@ -653,6 +656,67 @@ FROM WHERE users.id = @user_id; +-- name: GetActiveUsersAuthorizationRoles :many +-- Returns the authorization roles (site and org-scoped, including implied +-- member roles and organization default roles) and the group memberships +-- for every active, non-deleted user who is neither a system user nor a +-- service account, matching the GetActiveUserCount population. +-- Must stay semantically in sync with GetAuthorizationUserRoles; +-- TestGetActiveUsersAuthorizationRolesParity enforces this. +WITH org_roles AS ( + SELECT + organization_members.user_id, + -- The roles are returned as a flat array, org scoped and site side. + -- Concatenating the organization id scopes the organization roles. + array_agg(org_role || ':' || organization_members.organization_id::text) AS roles + FROM + organization_members + JOIN organizations ON organizations.id = organization_members.organization_id, + -- All org members get an implied organization-member role for + -- their orgs. Memberships of service accounts are aggregated here + -- too, but their rows never survive the join against the outer + -- WHERE, so the organization-service-account case does not apply. + -- + -- organizations.default_org_member_roles applies to every member + -- but is not materialized on membership rows, so it is unioned in + -- here. + unnest( + array_cat( + array_append(organization_members.roles, 'organization-member'), + organizations.default_org_member_roles + ) + ) AS org_role + GROUP BY + organization_members.user_id +), +user_groups AS ( + SELECT + group_members.user_id, + array_agg(group_members.group_id :: text) AS groups + FROM + group_members + GROUP BY + group_members.user_id +) +SELECT + users.id, + array_cat( + -- All users are members + array_append(users.rbac_roles, 'member'), + -- Users with no org memberships have no org_roles row. + coalesce(org_roles.roles, ARRAY[]::text[]) + ) :: text[] AS roles, + coalesce(user_groups.groups, ARRAY[]::text[]) :: text[] AS groups +FROM + users + LEFT JOIN org_roles ON org_roles.user_id = users.id + LEFT JOIN user_groups ON user_groups.user_id = users.id +WHERE + users.status = 'active'::user_status + AND users.deleted = false + AND users.is_system = false + AND users.is_service_account = false; + -- name: UpdateUserQuietHoursSchedule :one UPDATE users diff --git a/coderd/rbac/rolestore/rolestore.go b/coderd/rbac/rolestore/rolestore.go index 9f95c1870a8..233df7ae527 100644 --- a/coderd/rbac/rolestore/rolestore.go +++ b/coderd/rbac/rolestore/rolestore.go @@ -32,6 +32,34 @@ func CustomRoleCacheContext(ctx context.Context) context.Context { return context.WithValue(ctx, customRoleCtxKey{}, syncmap.New[string, rbac.Role]()) } +// PrefetchCustomRoles fetches every custom role in a single query and +// stores them in the returned context's role cache, so Expand calls on +// that context resolve custom roles without further database lookups. +// Roles deleted after the prefetch are still absent from the cache and +// fall back to an individual lookup on Expand. +func PrefetchCustomRoles(ctx context.Context, db database.Store) (context.Context, error) { + ctx = CustomRoleCacheContext(ctx) + cache := roleCache(ctx) + + dbroles, err := db.CustomRoles(ctx, database.CustomRolesParams{ + LookupRoles: nil, + ExcludeOrgRoles: false, + OrganizationID: uuid.Nil, + IncludeSystemRoles: true, + }) + if err != nil { + return ctx, xerrors.Errorf("fetch custom roles: %w", err) + } + for _, dbrole := range dbroles { + converted, err := ConvertDBRole(dbrole) + if err != nil { + return ctx, xerrors.Errorf("convert db role %q: %w", dbrole.Name, err) + } + cache.Store(dbrole.RoleIdentifier().String(), converted) + } + return ctx, nil +} + func roleCache(ctx context.Context) *syncmap.Map[string, rbac.Role] { c, ok := ctx.Value(customRoleCtxKey{}).(*syncmap.Map[string, rbac.Role]) if !ok { diff --git a/coderd/rbac/rolestore/rolestore_test.go b/coderd/rbac/rolestore/rolestore_test.go index 80b6fb40f4c..78a2c3de626 100644 --- a/coderd/rbac/rolestore/rolestore_test.go +++ b/coderd/rbac/rolestore/rolestore_test.go @@ -1,15 +1,19 @@ package rolestore_test import ( + "context" "database/sql" "testing" "github.com/google/uuid" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + "golang.org/x/xerrors" "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbgen" + "github.com/coder/coder/v2/coderd/database/dbmock" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/rbac/rolestore" @@ -42,6 +46,63 @@ func TestExpandCustomRoleRoles(t *testing.T) { require.Len(t, roles, 1, "role found") } +func TestPrefetchCustomRoles(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + + orgID := uuid.New() + prefetched := database.CustomRole{ + Name: "prefetched", + DisplayName: "Prefetched", + OrganizationID: uuid.NullUUID{UUID: orgID, Valid: true}, + } + // The mock permits exactly one CustomRoles call: the unfiltered + // prefetch. A cache miss in Expand below would fail the test with an + // unexpected second call. + mDB.EXPECT().CustomRoles(gomock.Any(), database.CustomRolesParams{ + LookupRoles: nil, + ExcludeOrgRoles: false, + OrganizationID: uuid.Nil, + IncludeSystemRoles: true, + }).Times(1).Return([]database.CustomRole{prefetched}, nil) + + ctx, err := rolestore.PrefetchCustomRoles(context.Background(), mDB) + require.NoError(t, err) + + roles, err := rolestore.Expand(ctx, mDB, []rbac.RoleIdentifier{{Name: "prefetched", OrganizationID: orgID}}) + require.NoError(t, err) + require.Len(t, roles, 1) + require.Equal(t, "prefetched", roles[0].Identifier.Name) +} + +func TestPrefetchCustomRolesErrors(t *testing.T) { + t.Parallel() + + t.Run("FetchError", func(t *testing.T) { + t.Parallel() + mDB := dbmock.NewMockStore(gomock.NewController(t)) + mDB.EXPECT().CustomRoles(gomock.Any(), gomock.Any()).Return(nil, xerrors.New("boom")) + + _, err := rolestore.PrefetchCustomRoles(context.Background(), mDB) + require.ErrorContains(t, err, "fetch custom roles") + }) + + t.Run("ConvertError", func(t *testing.T) { + t.Parallel() + // Org permissions without an organization ID cannot be converted. + mDB := dbmock.NewMockStore(gomock.NewController(t)) + mDB.EXPECT().CustomRoles(gomock.Any(), gomock.Any()).Return([]database.CustomRole{{ + Name: "broken", + OrgPermissions: []database.CustomRolePermission{{ResourceType: "workspace", Action: "create"}}, + }}, nil) + + _, err := rolestore.PrefetchCustomRoles(context.Background(), mDB) + require.ErrorContains(t, err, `convert db role "broken"`) + }) +} + func TestReconcileSystemRole(t *testing.T) { t.Parallel() diff --git a/codersdk/deployment.go b/codersdk/deployment.go index 4a749c21fc5..4fbc7cc388e 100644 --- a/codersdk/deployment.go +++ b/codersdk/deployment.go @@ -5250,18 +5250,19 @@ type Experiment string const ( // Add new experiments here! - ExperimentExample Experiment = "example" // This isn't used for anything. - ExperimentAutoFillParameters Experiment = "auto-fill-parameters" // This should not be taken out of experiments until we have redesigned the feature. - ExperimentNotifications Experiment = "notifications" // Sends notifications via SMTP and webhooks following certain events. - ExperimentWorkspaceUsage Experiment = "workspace-usage" // Enables the new workspace usage tracking. - ExperimentOAuth2 Experiment = "oauth2" // Enables OAuth2 provider functionality. - ExperimentMCPServerHTTP Experiment = "mcp-server-http" // Enables the MCP HTTP server functionality. - ExperimentWorkspaceBuildUpdates Experiment = "workspace-build-updates" // Enables publishing workspace build updates to the all builds pubsub channel. - ExperimentNATSPubsub Experiment = "nats_pubsub" // Enables embedded NATS pubsub. - ExperimentMinimumImplicitMember Experiment = "minimum-implicit-member" // Allows organizations to deviate from the default organization-member roles, in support of Gateway Accounts. - ExperimentAIGatewayCostControl Experiment = "ai-gateway-cost-control" // Enables AI Gateway cost control functionality. - ExperimentChatAdvisor Experiment = "chat-advisor" // Enables the advisor tool for root agent chats. - ExperimentChatVirtualDesktop Experiment = "chat-virtual-desktop" // Enables virtual desktop and computer use provider for agents. + ExperimentExample Experiment = "example" // This isn't used for anything. + ExperimentAutoFillParameters Experiment = "auto-fill-parameters" // This should not be taken out of experiments until we have redesigned the feature. + ExperimentNotifications Experiment = "notifications" // Sends notifications via SMTP and webhooks following certain events. + ExperimentWorkspaceUsage Experiment = "workspace-usage" // Enables the new workspace usage tracking. + ExperimentOAuth2 Experiment = "oauth2" // Enables OAuth2 provider functionality. + ExperimentMCPServerHTTP Experiment = "mcp-server-http" // Enables the MCP HTTP server functionality. + ExperimentWorkspaceBuildUpdates Experiment = "workspace-build-updates" // Enables publishing workspace build updates to the all builds pubsub channel. + ExperimentNATSPubsub Experiment = "nats_pubsub" // Enables embedded NATS pubsub. + ExperimentMinimumImplicitMember Experiment = "minimum-implicit-member" // Allows organizations to deviate from the default organization-member roles, in support of Gateway Accounts. + ExperimentWorkspaceCapableLicensing Experiment = "workspace-capable-licensing" // Counts only users holding the workspace-create permission toward the license seat limit. + ExperimentAIGatewayCostControl Experiment = "ai-gateway-cost-control" // Enables AI Gateway cost control functionality. + ExperimentChatAdvisor Experiment = "chat-advisor" // Enables the advisor tool for root agent chats. + ExperimentChatVirtualDesktop Experiment = "chat-virtual-desktop" // Enables virtual desktop and computer use provider for agents. ) func (e Experiment) DisplayName() string { @@ -5284,6 +5285,8 @@ func (e Experiment) DisplayName() string { return "NATS Pubsub" case ExperimentMinimumImplicitMember: return "Gateway Accounts (minimum implicit member)" + case ExperimentWorkspaceCapableLicensing: + return "Workspace-Capable Licensing" case ExperimentAIGatewayCostControl: return "AI Gateway Cost Control" case ExperimentChatAdvisor: @@ -5309,6 +5312,7 @@ var ExperimentsKnown = Experiments{ ExperimentNATSPubsub, ExperimentWorkspaceBuildUpdates, ExperimentMinimumImplicitMember, + ExperimentWorkspaceCapableLicensing, ExperimentAIGatewayCostControl, ExperimentChatAdvisor, ExperimentChatVirtualDesktop, diff --git a/docs/reference/api/schemas.md b/docs/reference/api/schemas.md index c000a0cbffe..8c61cab1468 100644 --- a/docs/reference/api/schemas.md +++ b/docs/reference/api/schemas.md @@ -7242,9 +7242,9 @@ CreateWorkspaceRequest provides options for creating a new workspace. Only one o #### Enumerated Values -| Value(s) | -|--------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| -| `ai-gateway-cost-control`, `auto-fill-parameters`, `chat-advisor`, `chat-virtual-desktop`, `example`, `mcp-server-http`, `minimum-implicit-member`, `nats_pubsub`, `notifications`, `oauth2`, `workspace-build-updates`, `workspace-usage` | +| Value(s) | +|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| `ai-gateway-cost-control`, `auto-fill-parameters`, `chat-advisor`, `chat-virtual-desktop`, `example`, `mcp-server-http`, `minimum-implicit-member`, `nats_pubsub`, `notifications`, `oauth2`, `workspace-build-updates`, `workspace-capable-licensing`, `workspace-usage` | ## codersdk.ExternalAPIKeyScopes diff --git a/enterprise/coderd/coderd.go b/enterprise/coderd/coderd.go index a38b7ee830c..857161b4edf 100644 --- a/enterprise/coderd/coderd.go +++ b/enterprise/coderd/coderd.go @@ -937,7 +937,7 @@ func (api *API) updateEntitlements(ctx context.Context) error { } reloadedEntitlements, err := license.Entitlements( - ctx, api.Database, + ctx, api.Logger, api.Database, len(agedReplicas), len(api.ExternalAuthConfigs), api.LicenseKeys, map[codersdk.FeatureName]bool{ codersdk.FeatureAuditLog: api.AuditLogging, codersdk.FeatureConnectionLog: api.ConnectionLogging, @@ -953,7 +953,10 @@ func (api *API) updateEntitlements(ctx context.Context) error { codersdk.FeatureAccessControl: true, codersdk.FeatureControlSharedPorts: true, codersdk.FeatureAIBridge: api.DeploymentValues.AI.BridgeConfig.Enabled.Value(), - }) + }, + api.AGPL.HTTPAuth.Authorizer, + api.AGPL.Experiments, + ) if err != nil { return codersdk.Entitlements{}, err } diff --git a/enterprise/coderd/license/license.go b/enterprise/coderd/license/license.go index 8092e5f6258..8ee9ecd6a41 100644 --- a/enterprise/coderd/license/license.go +++ b/enterprise/coderd/license/license.go @@ -13,11 +13,19 @@ import ( "github.com/golang-jwt/jwt/v4" "golang.org/x/xerrors" + "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/codersdk" ) +// Exceeding this timeout fails the entitlements computation; the caller +// keeps serving the previous entitlements. The count normally completes +// in well under a second, but its cost scales with the number of unique +// role sets and it runs on a context with no deadline of its own. +const workspaceCapableUserCountTimeout = 60 * time.Second + // Entitlements processes licenses to return whether features are enabled or not. // TODO(@deansheather): This function and the related LicensesEntitlements // function should be refactored into smaller functions that: @@ -26,11 +34,14 @@ import ( // 3. generate warnings related to usage func Entitlements( ctx context.Context, + logger slog.Logger, db database.Store, replicaCount int, externalAuthCount int, keys map[string]ed25519.PublicKey, enablements map[codersdk.FeatureName]bool, + authorizer rbac.Authorizer, + experiments codersdk.Experiments, ) (codersdk.Entitlements, error) { now := time.Now() @@ -46,6 +57,24 @@ func Entitlements( return codersdk.Entitlements{}, xerrors.Errorf("query active user count: %w", err) } + // Workspace-capable licensing counts only users the RBAC engine + // authorizes to create workspaces. The mode alone decides whether the + // counting function below is invoked. + // + // TODO: when the workspace-capable-licensing experiment is removed, a + // nil authorizer must become a hard dev error rather than a silent + // fallback to active-user counting. Tests already pass a real + // authorizer; only the dedicated nil-fallback tests rely on this + // branch. + countingMode := UserCountingModeActive + if experiments.Enabled(codersdk.ExperimentWorkspaceCapableLicensing) { + if authorizer == nil { + logger.Warn(ctx, "workspace-capable licensing experiment is enabled but no authorizer is configured, counting all active users") + } else { + countingMode = UserCountingModeWorkspaceCapable + } + } + // nolint:gocritic // Getting active AI seat count is a system function. activeAISeatCount, err := db.GetActiveAISeatCount(dbauthz.AsSystemRestricted(ctx)) if err != nil { @@ -69,6 +98,12 @@ func Entitlements( ReplicaCount: replicaCount, ExternalAuthCount: externalAuthCount, ExternalTemplateCount: int64(len(externalTemplates)), + UserCountingMode: countingMode, + WorkspaceCapableUserCountFn: func(ctx context.Context) (int64, error) { + ctx, cancel := context.WithTimeout(ctx, workspaceCapableUserCountTimeout) + defer cancel() + return CountWorkspaceCapableUsers(ctx, logger, db, authorizer) + }, ManagedAgentCountFn: func(ctx context.Context, startTime time.Time, endTime time.Time) (int64, error) { // This is not super accurate, as the start and end times will be // truncated to the date in UTC timezone. This is an optimization @@ -103,10 +138,179 @@ type FeatureArguments struct { // state of the world, but a count between two points in time determined by // the licenses. ManagedAgentCountFn ManagedAgentCountFn + // UserCountingMode selects the count that FeatureUserLimit candidates + // from AI Governance addon licenses are evaluated against. Under + // UserCountingModeWorkspaceCapable they use WorkspaceCapableUserCountFn's + // count; under any other value, including the zero value, every + // candidate uses ActiveUserCount. + UserCountingMode UserCountingMode + // WorkspaceCapableUserCountFn returns the number of active users the + // RBAC engine authorizes to create workspaces. It is invoked only + // under UserCountingModeWorkspaceCapable, and only when a valid + // license carries both the AI Governance addon and a FeatureUserLimit + // claim; the result then applies to that license's FeatureUserLimit + // candidate, and replaces ActiveUserCount when such a candidate is + // selected for enforcement. May be nil under UserCountingModeActive; + // leaving it nil when the workspace-capable mode would invoke it is a + // dev error. + WorkspaceCapableUserCountFn WorkspaceCapableUserCountFn } +// UserCountingMode selects how license seats are counted for +// FeatureUserLimit candidates from AI Governance addon licenses. +type UserCountingMode string + +const ( + // UserCountingModeActive evaluates every FeatureUserLimit candidate + // against the active user count. + UserCountingModeActive UserCountingMode = "active_users" + // UserCountingModeWorkspaceCapable evaluates addon-carrying candidates + // against the workspace-capable user count. + UserCountingModeWorkspaceCapable UserCountingMode = "workspace_capable_users" +) + type ManagedAgentCountFn func(ctx context.Context, from time.Time, to time.Time) (int64, error) +type WorkspaceCapableUserCountFn func(ctx context.Context) (int64, error) + +// userLimitCandidate is one license's FeatureUserLimit terms: its seat limit, +// its entitlement, and the counting mode implied by whether the license +// carries the AI Governance addon (workspace-capable counting of +// workspace-capable users vs. counting all active users). +type userLimitCandidate struct { + limit int64 + entitlement codersdk.Entitlement + aiGovernanceAddon bool +} + +// resolvedCandidate pairs a candidate with the count its counting mode +// implies: the workspace-capable count for addon candidates when +// workspace-capable counting is active, the active user count otherwise. +type resolvedCandidate struct { + userLimitCandidate + count int64 +} + +// betterUserLimit reports whether candidate a is more favorable than b. +// Ordering mirrors Feature.Compare: a candidate whose count is within its +// limit beats one whose count is not, then higher entitlement, then +// higher limit; the addon mode breaks remaining ties since its count is +// never larger than the active user count. +func betterUserLimit(a, b resolvedCandidate) bool { + compliantA := a.count <= a.limit + compliantB := b.count <= b.limit + if compliantA != compliantB { + return compliantA + } + if a.entitlement.Weight() != b.entitlement.Weight() { + return a.entitlement.Weight() > b.entitlement.Weight() + } + if a.limit != b.limit { + return a.limit > b.limit + } + return a.aiGovernanceAddon && !b.aiGovernanceAddon +} + +// userLimitSelection reports how the enforced FeatureUserLimit was chosen. +type userLimitSelection struct { + // workspaceCapable is true when the selected candidate counts + // workspace-capable users rather than all active users. + workspaceCapable bool + // addonEntitled is true when at least one addon-carrying candidate is + // fully valid rather than in its grace period. + addonEntitled bool +} + +// selectUserLimit picks the most favorable FeatureUserLimit candidate and +// applies its terms to the entitlements. Every candidate is evaluated +// against the count its own license's mode implies (the workspace-capable +// count for workspace-capable candidates, the active user count +// otherwise), so one license's limit is never combined with another +// license's counting mode. A candidate satisfied by its count wins over +// any unsatisfied one. +// +// For example, a deployment holding a 200-seat non-addon license and a +// 100-seat AI Governance license: +// +// active | capable | 200-seat license | 100-seat addon license | selected +// 250 | 90 | over | satisfied | addon: 90/100 +// 180 | 150 | satisfied | over | non-addon: 180/200 +// +// Neither license's limit is ever paired with the other's count: 90 +// capable users against the 200-seat limit, or 180 active users against +// the 100-seat limit, are not considered. +// +// With no candidates the entitlements are left untouched. On a count +// failure the entitlements computation must be aborted. +func selectUserLimit( + ctx context.Context, + entitlements *codersdk.Entitlements, + featureArguments FeatureArguments, + candidates []userLimitCandidate, +) (userLimitSelection, error) { + var sel userLimitSelection + if len(candidates) == 0 { + return sel, nil + } + + hasAddonCandidate := false + for _, c := range candidates { + if c.aiGovernanceAddon { + hasAddonCandidate = true + if c.entitlement == codersdk.EntitlementEntitled { + sel.addonEntitled = true + } + } + } + + var capableCount *int64 + if hasAddonCandidate && featureArguments.UserCountingMode == UserCountingModeWorkspaceCapable { + if featureArguments.WorkspaceCapableUserCountFn == nil { + return sel, xerrors.New("dev error: workspace-capable user count function is not set") + } + count, err := featureArguments.WorkspaceCapableUserCountFn(ctx) + if err != nil { + // A failed seat count is deliberately a hard failure rather + // than a recorded entitlement error: continuing with + // ActiveUserCount would silently change what FeatureUserLimit + // measures. The caller keeps the previous entitlements, so a + // failure yields a stale count rather than a different one. + return sel, xerrors.Errorf("count workspace capable users: %w", err) + } + capableCount = &count + } + + resolved := make([]resolvedCandidate, len(candidates)) + for i, c := range candidates { + resolved[i] = resolvedCandidate{userLimitCandidate: c, count: featureArguments.ActiveUserCount} + if c.aiGovernanceAddon && capableCount != nil { + resolved[i].count = *capableCount + } + } + + best := resolved[0] + for _, c := range resolved[1:] { + if betterUserLimit(c, best) { + best = c + } + } + + if best.aiGovernanceAddon && capableCount != nil { + sel.workspaceCapable = true + } + + // AddFeature merged limits and entitlements across licenses without + // pairing them to counting modes; overwrite the merged terms with the + // selected candidate's. Actual is replaced wholesale, so the merged + // feature's alias of the caller's ActiveUserCount no longer matters. + userLimit := entitlements.Features[codersdk.FeatureUserLimit] + userLimit.Limit = &best.limit + userLimit.Entitlement = best.entitlement + userLimit.Actual = &best.count + entitlements.Features[codersdk.FeatureUserLimit] = userLimit + return sel, nil +} + // LicensesEntitlements returns the entitlements for licenses. Entitlements are // merged from all licenses and the highest entitlement is used for each feature. // Arguments: @@ -131,6 +335,13 @@ func LicensesEntitlements( // suppress the soft warning for AI Bridge GA. hasExplicitAIBridgeEntitlement := false + // Each valid license's FeatureUserLimit claim forms a candidate pairing of + // seat limit and counting mode: licenses carrying the AI Governance + // addon count workspace-capable users, others count all active users. + // The most favorable candidate is selected once all licenses are + // processed. + var userLimitCandidates []userLimitCandidate + // Default all entitlements to be disabled. entitlements := codersdk.Entitlements{ Features: map[codersdk.FeatureName]codersdk.Feature{ @@ -357,6 +568,7 @@ func LicensesEntitlements( } addonFeatures := make(map[codersdk.FeatureName]codersdk.Feature) + licenseHasAIGovernanceAddon := false // Finally, add all features from the addons. We do this last so that // any dependencies of an addon are validated against the calculated @@ -372,6 +584,9 @@ func LicensesEntitlements( // Ignore the addon and don't add any features. continue } + if addon == codersdk.AddonAIGovernance { + licenseHasAIGovernanceAddon = true + } for _, featureName := range addon.Features() { if _, exists := addonFeatures[featureName]; !exists { addonFeatures[featureName] = codersdk.Feature{ @@ -384,6 +599,21 @@ func LicensesEntitlements( for featureName, feature := range addonFeatures { entitlements.AddFeature(featureName, feature) } + + if limit := claims.Features[codersdk.FeatureUserLimit]; limit > 0 { + userLimitCandidates = append(userLimitCandidates, userLimitCandidate{ + limit: limit, + entitlement: entitlement, + aiGovernanceAddon: licenseHasAIGovernanceAddon, + }) + } + } + + // The FeatureUserLimit feature's final terms come from best-pair selection + // across the candidates rather than the AddFeature merge. + userLimitSel, err := selectUserLimit(ctx, &entitlements, featureArguments, userLimitCandidates) + if err != nil { + return entitlements, err } // Now the license specific warnings and errors are added to the entitlements. @@ -480,14 +710,33 @@ func LicensesEntitlements( if entitlements.HasLicense { userLimit := entitlements.Features[codersdk.FeatureUserLimit] - if userLimit.Limit != nil && featureArguments.ActiveUserCount > *userLimit.Limit { + // The enforced count and its meaning come from the selected + // candidate: userLimit.Actual is the count the limit was evaluated + // against, and the noun names what it counted. + userLimitActual := featureArguments.ActiveUserCount + if userLimit.Actual != nil { + userLimitActual = *userLimit.Actual + } + userNoun := "active users" + if userLimitSel.workspaceCapable { + userNoun = "workspace-capable users" + } + if userLimit.Limit != nil && userLimitActual > *userLimit.Limit { entitlements.Warnings = append(entitlements.Warnings, fmt.Sprintf( - "Your deployment has %d active users but is only licensed for %d.", - featureArguments.ActiveUserCount, *userLimit.Limit)) + "Your deployment has %d %s but is only licensed for %d.", + userLimitActual, userNoun, *userLimit.Limit)) } else if userLimit.Limit != nil && userLimit.Entitlement == codersdk.EntitlementGracePeriod { entitlements.Warnings = append(entitlements.Warnings, fmt.Sprintf( - "Your deployment has %d active users but the license with the limit %d is expired.", - featureArguments.ActiveUserCount, *userLimit.Limit)) + "Your deployment has %d %s but the license with the limit %d is expired.", + userLimitActual, userNoun, *userLimit.Limit)) + } + // The addon exists only on grace-period licenses: warn that + // workspace-capable counting stops at the end of the grace period, + // at which point every active user counts. + if userLimitSel.workspaceCapable && !userLimitSel.addonEntitled { + entitlements.Warnings = append(entitlements.Warnings, fmt.Sprintf( + "Your license with the AI Governance addon is expired. When it fully expires, all %d active users will count toward the user limit instead of the %d workspace-capable users.", + featureArguments.ActiveUserCount, userLimitActual)) } if featureArguments.ActiveAISeatCount > 0 { actual := featureArguments.ActiveAISeatCount diff --git a/enterprise/coderd/license/license_test.go b/enterprise/coderd/license/license_test.go index 10dea231cd8..f41a0ee3c9b 100644 --- a/enterprise/coderd/license/license_test.go +++ b/enterprise/coderd/license/license_test.go @@ -8,6 +8,7 @@ import ( "time" "github.com/google/uuid" + "github.com/prometheus/client_golang/prometheus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" @@ -16,11 +17,18 @@ import ( "github.com/coder/coder/v2/coderd/database/dbmock" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/database/dbtime" + "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" "github.com/coder/coder/v2/enterprise/coderd/license" + "github.com/coder/coder/v2/testutil" ) +// testAuthorizer satisfies Entitlements' expectation of a non-nil +// authorizer. The callers below never enable the workspace-capable +// licensing experiment, so it is never asked to authorize anything. +var testAuthorizer = rbac.NewCachingAuthorizer(prometheus.NewRegistry()) + func TestEntitlements(t *testing.T) { t.Parallel() all := make(map[codersdk.FeatureName]bool) @@ -33,7 +41,7 @@ func TestEntitlements(t *testing.T) { t.Run("Defaults", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.False(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -45,7 +53,7 @@ func TestEntitlements(t *testing.T) { t.Run("Always return the current user count", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.False(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -58,7 +66,7 @@ func TestEntitlements(t *testing.T) { JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{}), Exp: dbtime.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -86,7 +94,7 @@ func TestEntitlements(t *testing.T) { }), Exp: dbtime.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -110,7 +118,7 @@ func TestEntitlements(t *testing.T) { }), Exp: dbtime.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -137,7 +145,7 @@ func TestEntitlements(t *testing.T) { Exp: dbtime.Now().AddDate(0, 0, 5), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -166,7 +174,7 @@ func TestEntitlements(t *testing.T) { Exp: time.Now().AddDate(0, 0, 5), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -202,7 +210,7 @@ func TestEntitlements(t *testing.T) { require.NoError(t, err) // Warning should be generated. - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -230,7 +238,7 @@ func TestEntitlements(t *testing.T) { require.NoError(t, err) // Warning should be suppressed. - entitlements, err = license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err = license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -261,7 +269,7 @@ func TestEntitlements(t *testing.T) { require.NoError(t, err) // Should generate a warning. - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -289,7 +297,7 @@ func TestEntitlements(t *testing.T) { require.NoError(t, err) // Warning should still be generated. - entitlements, err = license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err = license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -315,7 +323,7 @@ func TestEntitlements(t *testing.T) { Exp: dbtime.Now().AddDate(0, 0, 5), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -344,7 +352,7 @@ func TestEntitlements(t *testing.T) { Exp: dbtime.Now().AddDate(0, 0, 5), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -364,7 +372,7 @@ func TestEntitlements(t *testing.T) { JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{}), Exp: time.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -434,7 +442,7 @@ func TestEntitlements(t *testing.T) { }), Exp: time.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.Contains(t, entitlements.Warnings, "Your deployment has 2 active users but is only licensed for 1.") @@ -462,7 +470,7 @@ func TestEntitlements(t *testing.T) { }), Exp: time.Now().Add(60 * 24 * time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.Empty(t, entitlements.Warnings) @@ -485,7 +493,7 @@ func TestEntitlements(t *testing.T) { }), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -501,7 +509,7 @@ func TestEntitlements(t *testing.T) { }), }) require.NoError(t, err) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -549,7 +557,7 @@ func TestEntitlements(t *testing.T) { JWT: coderdenttest.GenerateLicense(t, licenseOptions), }) require.NoError(t, err) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -600,7 +608,7 @@ func TestEntitlements(t *testing.T) { }), }) require.NoError(t, err) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -621,7 +629,7 @@ func TestEntitlements(t *testing.T) { AllFeatures: true, }), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -654,7 +662,7 @@ func TestEntitlements(t *testing.T) { AllFeatures: true, }), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -688,7 +696,7 @@ func TestEntitlements(t *testing.T) { ExpiresAt: dbtime.Now().Add(time.Hour), }), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.False(t, entitlements.Trial) @@ -714,7 +722,7 @@ func TestEntitlements(t *testing.T) { t.Run("MultipleReplicasNoLicense", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) - entitlements, err := license.Entitlements(context.Background(), db, 2, 1, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 2, 1, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.False(t, entitlements.HasLicense) require.Len(t, entitlements.Errors, 1) @@ -732,9 +740,9 @@ func TestEntitlements(t *testing.T) { }, }), }) - entitlements, err := license.Entitlements(context.Background(), db, 2, 1, coderdenttest.Keys, map[codersdk.FeatureName]bool{ + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 2, 1, coderdenttest.Keys, map[codersdk.FeatureName]bool{ codersdk.FeatureHighAvailability: true, - }) + }, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.Len(t, entitlements.Errors, 1) @@ -755,9 +763,9 @@ func TestEntitlements(t *testing.T) { }), Exp: time.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 2, 1, coderdenttest.Keys, map[codersdk.FeatureName]bool{ + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 2, 1, coderdenttest.Keys, map[codersdk.FeatureName]bool{ codersdk.FeatureHighAvailability: true, - }) + }, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.Len(t, entitlements.Warnings, 1) @@ -767,7 +775,7 @@ func TestEntitlements(t *testing.T) { t.Run("MultipleGitAuthNoLicense", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) - entitlements, err := license.Entitlements(context.Background(), db, 1, 2, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 2, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.False(t, entitlements.HasLicense) require.Len(t, entitlements.Errors, 1) @@ -785,9 +793,9 @@ func TestEntitlements(t *testing.T) { }, }), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 2, coderdenttest.Keys, map[codersdk.FeatureName]bool{ + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 2, coderdenttest.Keys, map[codersdk.FeatureName]bool{ codersdk.FeatureMultipleExternalAuth: true, - }) + }, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.Len(t, entitlements.Errors, 1) @@ -808,9 +816,9 @@ func TestEntitlements(t *testing.T) { }), Exp: time.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 2, coderdenttest.Keys, map[codersdk.FeatureName]bool{ + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 2, coderdenttest.Keys, map[codersdk.FeatureName]bool{ codersdk.FeatureMultipleExternalAuth: true, - }) + }, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) require.Len(t, entitlements.Warnings, 1) @@ -875,7 +883,7 @@ func TestEntitlements(t *testing.T) { GetTemplatesWithFilter(gomock.Any(), gomock.Any()). Return([]database.Template{}, nil) - entitlements, err := license.Entitlements(context.Background(), mDB, 1, 0, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), mDB, 1, 0, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -993,7 +1001,7 @@ func TestEntitlements(t *testing.T) { GetTemplatesWithFilter(gomock.Any(), gomock.Any()). Return([]database.Template{}, nil) - entitlements, err := license.Entitlements(context.Background(), mDB, 1, 0, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), mDB, 1, 0, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -1063,7 +1071,7 @@ func TestEntitlements(t *testing.T) { codersdk.FeatureAIGovernanceUserLimit: true, } - entitlements, err := license.Entitlements(context.Background(), mDB, 1, 0, coderdenttest.Keys, enablements) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), mDB, 1, 0, coderdenttest.Keys, enablements, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -1126,7 +1134,7 @@ func TestEntitlements(t *testing.T) { codersdk.FeatureAIGovernanceUserLimit: true, } - entitlements, err := license.Entitlements(context.Background(), mDB, 1, 0, coderdenttest.Keys, enablements) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), mDB, 1, 0, coderdenttest.Keys, enablements, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -1186,7 +1194,7 @@ func TestEntitlements(t *testing.T) { GetTemplatesWithFilter(gomock.Any(), gomock.Any()). Return([]database.Template{}, nil) - entitlements, err := license.Entitlements(context.Background(), mDB, 1, 0, coderdenttest.Keys, all) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), mDB, 1, 0, coderdenttest.Keys, all, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -2163,7 +2171,7 @@ func TestAIGovernanceAddon(t *testing.T) { codersdk.FeatureAIBridge: true, codersdk.FeatureBoundary: true, } - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, enablements) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -2195,7 +2203,7 @@ func TestAIGovernanceAddon(t *testing.T) { codersdk.FeatureAIBridge: true, codersdk.FeatureBoundary: true, } - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, enablements) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -2233,7 +2241,7 @@ func TestAIGovernanceAddon(t *testing.T) { codersdk.FeatureAIBridge: true, codersdk.FeatureBoundary: true, } - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, enablements) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -2264,7 +2272,7 @@ func TestAIGovernanceAddon(t *testing.T) { Exp: dbtime.Now().Add(time.Hour), }) - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, empty) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, empty, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) @@ -2297,7 +2305,7 @@ func TestAIGovernanceAddon(t *testing.T) { codersdk.FeatureAIBridge: true, codersdk.FeatureBoundary: true, } - entitlements, err := license.Entitlements(context.Background(), db, 1, 1, coderdenttest.Keys, enablements) + entitlements, err := license.Entitlements(context.Background(), testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, testAuthorizer, nil) require.NoError(t, err) require.True(t, entitlements.HasLicense) diff --git a/enterprise/coderd/license/usercount.go b/enterprise/coderd/license/usercount.go new file mode 100644 index 00000000000..3699456d6d8 --- /dev/null +++ b/enterprise/coderd/license/usercount.go @@ -0,0 +1,154 @@ +package license + +import ( + "context" + "crypto/sha256" + "encoding/json" + "slices" + "strings" + "time" + + "github.com/google/uuid" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/coderd/rbac" + "github.com/coder/coder/v2/coderd/rbac/policy" + "github.com/coder/coder/v2/coderd/rbac/rolestore" +) + +// countingSubjectID replaces the real user ID in every evaluated subject +// and on the object owner. The substitution is safe because the policy +// only ever compares the subject ID to the object owner, and the +// evaluated object is synthetic with no user or group ACL lists, so no +// other rule can reference a real ID. Subjects with equal roles and +// groups are therefore byte-identical. +var countingSubjectID = uuid.MustParse("ad966897-b805-4a2c-8dab-3cfcbba0a683").String() + +// CountWorkspaceCapableUsers returns the number of active users the RBAC +// engine authorizes to create a workspace, either in one of the +// organizations they belong to or in any organization via a site-wide +// role such as owner. System users and service accounts are excluded by +// the underlying query, matching GetActiveUserCount. +func CountWorkspaceCapableUsers(ctx context.Context, logger slog.Logger, db database.Store, authorizer rbac.Authorizer) (int64, error) { + if authorizer == nil { + return 0, xerrors.New("dev error: authorizer is required") + } + + start := time.Now() + + // All custom roles are prefetched into the context's role cache in a + // single query; role expansion below then resolves both builtin and + // custom roles without per-role-set database lookups. + //nolint:gocritic // Counting licensed seats is a system function. + ctx, err := rolestore.PrefetchCustomRoles(dbauthz.AsSystemRestricted(ctx), db) + if err != nil { + return 0, xerrors.Errorf("prefetch custom roles: %w", err) + } + + //nolint:gocritic // Counting licensed seats is a system function. + rows, err := db.GetActiveUsersAuthorizationRoles(dbauthz.AsSystemRestricted(ctx)) + if err != nil { + return 0, xerrors.Errorf("get active users authorization roles: %w", err) + } + + // Users with equivalent canonical subjects share one authorization + // verdict, so evaluation cost scales with unique subjects, not users. + capableBySignature := make(map[[sha256.Size]byte]bool) + var count int64 + for _, row := range rows { + roleNames, err := row.RoleNames() + if err != nil { + // A stored role string that fails to parse grants nothing: + // authorization fails closed on it, so this user cannot + // create a workspace. Treat the user as not capable instead + // of failing the entire count. + logger.Warn(ctx, "user has an unparsable role, counting them as not workspace-capable for license seats", + slog.F("user_id", row.ID), + slog.Error(err), + ) + continue + } + subject := countingSubject(roleNames, row.Groups) + sig, err := authorizationSignature(subject) + if err != nil { + return 0, xerrors.Errorf("compute authorization signature for user %s: %w", row.ID, err) + } + capable, ok := capableBySignature[sig] + if !ok { + capable, err = canCreateWorkspace(ctx, db, authorizer, subject) + if err != nil { + return 0, xerrors.Errorf("evaluate workspace-create for user %s: %w", row.ID, err) + } + capableBySignature[sig] = capable + } + if capable { + count++ + } + } + + // Emitted only when workspace-capable counting runs, so the line's + // presence identifies the counting mode (workspace-capable vs. all + // active users). + logger.Info(ctx, "counted workspace-capable users for license seats", + slog.F("workspace_capable_users", count), + slog.F("active_users", len(rows)), + slog.F("unique_subjects", len(capableBySignature)), + slog.F("elapsed", time.Since(start)), + ) + return count, nil +} + +// countingSubject builds the canonical evaluation subject for a user. +func countingSubject(roleNames []rbac.RoleIdentifier, groups []string) rbac.Subject { + slices.SortFunc(roleNames, func(a, b rbac.RoleIdentifier) int { + return strings.Compare(a.String(), b.String()) + }) + roleNames = slices.CompactFunc(roleNames, func(a, b rbac.RoleIdentifier) bool { + return a == b + }) + groups = slices.Clone(groups) + slices.Sort(groups) + groups = slices.Compact(groups) + return rbac.Subject{ + Type: rbac.SubjectTypeUser, + ID: countingSubjectID, + Roles: rbac.RoleIdentifiers(roleNames), + Groups: groups, + Scope: rbac.ScopeAll, + } +} + +// authorizationSignature returns a hash of the subject's JSON form; +// every subject field, including any added later, is part of the key. +func authorizationSignature(subject rbac.Subject) ([sha256.Size]byte, error) { + var sig [sha256.Size]byte + hash := sha256.New() + if err := json.NewEncoder(hash).Encode(subject); err != nil { + return sig, xerrors.Errorf("encode subject: %w", err) + } + copy(sig[:], hash.Sum(nil)) + return sig, nil +} + +// canCreateWorkspace reports whether the RBAC engine authorizes the +// subject to create a workspace they own in any organization: via +// membership grants or via a site-wide role that applies regardless of +// membership. +func canCreateWorkspace(ctx context.Context, db database.Store, authorizer rbac.Authorizer, subject rbac.Subject) (bool, error) { + //nolint:gocritic // Expanding custom roles requires system access. + roles, err := rolestore.Expand(dbauthz.AsSystemRestricted(ctx), db, subject.SafeRoleNames()) + if err != nil { + return false, xerrors.Errorf("expand roles: %w", err) + } + subject.Roles = roles + subject = subject.WithCachedASTValue() + + // The any-organization form allows exactly when some per-organization + // check would (the policy takes the maximum vote across the subject's + // memberships), and also covers users who belong to zero organizations. + return authorizer.Authorize(ctx, subject, policy.ActionCreate, + rbac.ResourceWorkspace.AnyOrganization().WithOwner(subject.ID)) == nil, nil +} diff --git a/enterprise/coderd/license/usercount_bench_test.go b/enterprise/coderd/license/usercount_bench_test.go new file mode 100644 index 00000000000..4e635d04ef0 --- /dev/null +++ b/enterprise/coderd/license/usercount_bench_test.go @@ -0,0 +1,226 @@ +package license_test + +import ( + "context" + "database/sql" + "fmt" + "testing" + + "github.com/google/uuid" + "github.com/prometheus/client_golang/prometheus" + "github.com/stretchr/testify/require" + + "cdr.dev/slog/v3" + "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" + "github.com/coder/coder/v2/enterprise/coderd/license" +) + +// BenchmarkCountWorkspaceCapableUsers measures how workspace-capable seat +// counting scales along its two cost axes: the number of eligible active +// users (row fetch and per-row signature work) and the number of unique +// role sets (role expansion and rego evaluation, one per set). +// +// Scenarios (users/orgs/roles are seeded in bulk via SQL): +// +// - Uniform: one org, half gateways, half workspace-capable. Unique +// role sets stay constant, so this isolates per-user row cost. +// - ManyOrgs: users spread evenly across orgs, plain members. Unique +// role sets scale with org count. +// - UniquePairs: every user belongs to a distinct pair of orgs, so +// every user is a unique role set. Worst-case rego evaluation with +// builtin roles only. +// - CustomRoles: users hold org-scoped custom roles round-robin. +// Exercises the custom-role prefetch and expansion path. +// +// Run with: +// +// go test ./enterprise/coderd/license/ -bench BenchmarkCountWorkspaceCapableUsers -benchtime 5x -run '^$' -v +func BenchmarkCountWorkspaceCapableUsers(b *testing.B) { + // Workspace-create flows only through explicit grants under + // MinimumImplicitMember, so capability actually varies between users. + rbac.ReloadBuiltinRoles(&rbac.RoleOptions{MinimumImplicitMember: true}) + b.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + ctx := context.Background() + authorizer := rbac.NewCachingAuthorizer(prometheus.NewRegistry()) + // Discard logs: the per-count Info line and its fields are not what + // is being measured. + logger := slog.Make() + + for _, scenario := range []benchScenario{ + {name: "Uniform/1k", users: 1_000, orgs: 1}, + {name: "Uniform/10k", users: 10_000, orgs: 1}, + {name: "Uniform/50k", users: 50_000, orgs: 1}, + {name: "ManyOrgs/10k-100orgs", users: 10_000, orgs: 100}, + {name: "UniquePairs/10k", users: 10_000, orgs: 100, uniquePairs: true}, + {name: "CustomRoles/10k-1000roles", users: 10_000, orgs: 10, customRolesPerOrg: 100}, + } { + b.Run(scenario.name, func(b *testing.B) { + db, _, sqlDB := dbtestutil.NewDBWithSQLDB(b) + seedBenchUsers(ctx, b, db, sqlDB, scenario) + + b.ResetTimer() + var count int64 + for i := 0; i < b.N; i++ { + var err error + count, err = license.CountWorkspaceCapableUsers(ctx, logger, db, authorizer) + require.NoError(b, err) + } + b.StopTimer() + require.NotZero(b, count, "scenario must produce capable users") + b.ReportMetric(float64(scenario.users), "users") + b.ReportMetric(float64(count), "capable") + }) + } +} + +type benchScenario struct { + name string + users int + orgs int + // uniquePairs gives every user memberships in a distinct pair of + // orgs, making every user a unique role set. + uniquePairs bool + // customRolesPerOrg grants each user one org-scoped custom role, + // assigned round-robin. + customRolesPerOrg int +} + +// seedBenchUsers bulk-inserts active users and their org memberships. +// Deterministic UUIDs (zero-prefixed, numbered) let memberships be +// generated from the same series without returning inserted rows. +func seedBenchUsers(ctx context.Context, b *testing.B, db database.Store, sqlDB *sql.DB, s benchScenario) { + b.Helper() + + orgIDs := make([]uuid.UUID, s.orgs) + for i := range orgIDs { + org := dbgen.Organization(b, db, database.Organization{}) + emptyOrgDefaultRoles(ctx, b, db, org) + orgIDs[i] = org.ID + } + + // Deterministic user IDs let membership rows be generated from the + // same series without returning inserted rows. + _, err := sqlDB.ExecContext(ctx, ` + CREATE OR REPLACE FUNCTION benchUserID(i bigint) RETURNS uuid AS $$ + SELECT ('00000000-0000-0000-0000-' || lpad(i::text, 12, '0'))::uuid + $$ LANGUAGE sql IMMUTABLE; + `) + require.NoError(b, err) + + _, err = sqlDB.ExecContext(ctx, ` + INSERT INTO users (id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type) + SELECT + benchUserID(i), + 'bench-' || i || '@example.com', + 'bench-' || i, + '\x'::bytea, + now(), now(), + 'active'::user_status, + '{}'::text[], + 'password'::login_type + FROM generate_series(1, $1) AS i; + `, s.users) + require.NoError(b, err) + + orgIDText := make([]string, len(orgIDs)) + for i, id := range orgIDs { + orgIDText[i] = id.String() + } + + switch { + case s.uniquePairs: + // Membership in orgs (i mod K) and (i/K mod K): distinct pairs, + // hence distinct role sets, for up to K^2 users. Even users hold + // the workspace-access grant in their first org so capability + // varies across the population. + require.GreaterOrEqual(b, s.orgs*s.orgs, s.users, "not enough org pairs for unique role sets") + _, err = sqlDB.ExecContext(ctx, fmt.Sprintf(` + INSERT INTO organization_members (user_id, organization_id, created_at, updated_at, roles) + SELECT benchUserID(i), ($3::uuid[])[(i %% $2) + 1], now(), now(), + CASE WHEN i %% 2 = 0 THEN ARRAY['%s']::text[] ELSE '{}'::text[] END + FROM generate_series(1, $1) AS i + ON CONFLICT DO NOTHING; + `, rbac.RoleOrgWorkspaceAccess()), s.users, s.orgs, pqStringArray(orgIDText)) + require.NoError(b, err) + _, err = sqlDB.ExecContext(ctx, ` + INSERT INTO organization_members (user_id, organization_id, created_at, updated_at, roles) + SELECT benchUserID(i), ($3::uuid[])[((i / $2) % $2) + 1], now(), now(), '{}'::text[] + FROM generate_series(1, $1) AS i + ON CONFLICT DO NOTHING; + `, s.users, s.orgs, pqStringArray(orgIDText)) + require.NoError(b, err) + case s.customRolesPerOrg > 0: + // One workspace-create custom role per (org, slot), granted + // round-robin: users cycle through orgs, and within an org + // through its roles. + _, err = sqlDB.ExecContext(ctx, ` + INSERT INTO custom_roles (name, display_name, organization_id, org_permissions) + SELECT + 'bench-role-' || slot, + 'Bench Role ' || slot, + ($2::uuid[])[(slot % $3) + 1], + '[{"negate":false,"resource_type":"workspace","action":"create"}]'::jsonb + FROM generate_series(0, $1 - 1) AS slot; + `, s.orgs*s.customRolesPerOrg, pqStringArray(orgIDText), s.orgs) + require.NoError(b, err) + _, err = sqlDB.ExecContext(ctx, ` + INSERT INTO organization_members (user_id, organization_id, created_at, updated_at, roles) + SELECT + benchUserID(i), + ($2::uuid[])[(i % $3) + 1], + now(), now(), + ARRAY['bench-role-' || ((i % ($3 * $4) / $3) * $3 + (i % $3))]::text[] + FROM generate_series(1, $1) AS i; + `, s.users, pqStringArray(orgIDText), s.orgs, s.customRolesPerOrg) + require.NoError(b, err) + default: + // Round-robin plain membership; every even user additionally + // holds the workspace-access grant so capability varies. + _, err = sqlDB.ExecContext(ctx, fmt.Sprintf(` + INSERT INTO organization_members (user_id, organization_id, created_at, updated_at, roles) + SELECT + benchUserID(i), + ($2::uuid[])[(i %% $3) + 1], + now(), now(), + CASE WHEN i %% 2 = 0 THEN ARRAY['%s']::text[] ELSE '{}'::text[] END + FROM generate_series(1, $1) AS i; + `, rbac.RoleOrgWorkspaceAccess()), s.users, pqStringArray(orgIDText), s.orgs) + require.NoError(b, err) + } + + // Bulk inserts leave planner statistics claiming near-empty tables, + // which makes the roles query fall into a nested-loop plan that + // re-executes its aggregation per user row. Refresh them the way + // autovacuum would have in a live deployment. + _, err = sqlDB.ExecContext(ctx, `ANALYZE users; ANALYZE organization_members; ANALYZE organizations; ANALYZE custom_roles;`) + require.NoError(b, err) +} + +func emptyOrgDefaultRoles(ctx context.Context, b *testing.B, db database.Store, org database.Organization) { + b.Helper() + _, err := db.UpdateOrganization(ctx, database.UpdateOrganizationParams{ + ID: org.ID, + UpdatedAt: org.UpdatedAt, + Name: org.Name, + DisplayName: org.DisplayName, + Description: org.Description, + Icon: org.Icon, + DefaultOrgMemberRoles: []string{}, + }) + require.NoError(b, err) +} + +func pqStringArray(elems []string) string { + out := "{" + for i, e := range elems { + if i > 0 { + out += "," + } + out += e + } + return out + "}" +} diff --git a/enterprise/coderd/license/usercount_test.go b/enterprise/coderd/license/usercount_test.go new file mode 100644 index 00000000000..52751096451 --- /dev/null +++ b/enterprise/coderd/license/usercount_test.go @@ -0,0 +1,684 @@ +package license_test + +import ( + "context" + "testing" + "time" + + "github.com/google/uuid" + "github.com/prometheus/client_golang/prometheus" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbgen" + "github.com/coder/coder/v2/coderd/database/dbmock" + "github.com/coder/coder/v2/coderd/database/dbtestutil" + "github.com/coder/coder/v2/coderd/database/dbtime" + "github.com/coder/coder/v2/coderd/rbac" + "github.com/coder/coder/v2/coderd/rbac/policy" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" + "github.com/coder/coder/v2/enterprise/coderd/license" + "github.com/coder/coder/v2/testutil" +) + +// TestCountWorkspaceCapableUsers verifies workspace-capable license seat +// counting: only users the RBAC engine authorizes to create workspaces +// consume seats, so members without workspace-create ("gateway accounts") +// are excluded. +// +// The subtests toggle the global builtin roles via ReloadBuiltinRoles, so +// they must run serially. +// +//nolint:tparallel,paralleltest +func TestCountWorkspaceCapableUsers(t *testing.T) { + ctx := context.Background() + authorizer := rbac.NewCachingAuthorizer(prometheus.NewRegistry()) + + activeUser := func(t *testing.T, db database.Store, seed database.User) database.User { + seed.Status = database.UserStatusActive + return dbgen.User(t, db, seed) + } + member := func(t *testing.T, db database.Store, orgID uuid.UUID, user database.User, roles ...string) { + dbgen.OrganizationMember(t, db, database.OrganizationMember{ + OrganizationID: orgID, + UserID: user.ID, + Roles: roles, + }) + } + emptyDefaultRoles := func(t *testing.T, db database.Store, org database.Organization) { + _, err := db.UpdateOrganization(ctx, database.UpdateOrganizationParams{ + ID: org.ID, + UpdatedAt: dbtime.Now(), + Name: org.Name, + DisplayName: org.DisplayName, + Description: org.Description, + Icon: org.Icon, + DefaultOrgMemberRoles: []string{}, + }) + require.NoError(t, err) + } + + t.Run("ElevationBundledParity", func(t *testing.T) { + // MinimumImplicitMember off (default): organization-member bundles + // the workspace-ops elevation, so every active org member counts + // and the workspace-capable count matches the legacy count except + // for zero-org plain members. + rbac.ReloadBuiltinRoles(nil) + t.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + db, _ := dbtestutil.NewDB(t) + org := dbgen.Organization(t, db, database.Organization{}) + + plainMember := activeUser(t, db, database.User{}) + member(t, db, org.ID, plainMember) + + orgAdmin := activeUser(t, db, database.User{}) + member(t, db, org.ID, orgAdmin, rbac.RoleOrgAdmin()) + + owner := activeUser(t, db, database.User{RBACRoles: []string{rbac.RoleOwner().Name}}) + member(t, db, org.ID, owner) + + // Counts under legacy, not under workspace-capable counting: no org, no + // workspace-create anywhere. + activeUser(t, db, database.User{}) + + // Counts under both: the owner site role grants workspace-create + // in any organization, independent of membership. + activeUser(t, db, database.User{RBACRoles: []string{rbac.RoleOwner().Name}}) + + // Never counted: not active. + suspended := dbgen.User(t, db, database.User{Status: database.UserStatusSuspended}) + member(t, db, org.ID, suspended) + dormant := dbgen.User(t, db, database.User{Status: database.UserStatusDormant}) + member(t, db, org.ID, dormant) + + // Never counted: service accounts are excluded from seat counts. + sa := activeUser(t, db, database.User{IsServiceAccount: true}) + member(t, db, org.ID, sa) + + legacy, err := db.GetActiveUserCount(ctx, false) + require.NoError(t, err) + require.Equal(t, int64(5), legacy) + + count, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), db, authorizer) + require.NoError(t, err) + require.Equal(t, int64(4), count, "zero-org plain member must not count") + }) + + t.Run("MinimumImplicitMember", func(t *testing.T) { + // MinimumImplicitMember on: organization-member carries only the + // floor. Workspace-create flows exclusively through the + // organization-workspace-access role, granted explicitly or via + // default_org_member_roles. + rbac.ReloadBuiltinRoles(&rbac.RoleOptions{MinimumImplicitMember: true}) + t.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + db, _ := dbtestutil.NewDB(t) + org := dbgen.Organization(t, db, database.Organization{}) + emptyDefaultRoles(t, db, org) + + // Gateway account: floor only, no workspace-create. Not counted. + gateway := activeUser(t, db, database.User{}) + member(t, db, org.ID, gateway) + + // Explicit organization-workspace-access grant. Counted. + wsUser := activeUser(t, db, database.User{}) + member(t, db, org.ID, wsUser, rbac.RoleOrgWorkspaceAccess()) + + // The creation ban negates workspace-create even when the + // workspace-access role is present. Not counted. + banned := activeUser(t, db, database.User{}) + member(t, db, org.ID, banned, rbac.RoleOrgWorkspaceAccess(), rbac.RoleOrgWorkspaceCreationBan()) + + // Org admins retain workspace-create. Counted. + orgAdmin := activeUser(t, db, database.User{}) + member(t, db, org.ID, orgAdmin, rbac.RoleOrgAdmin()) + + // Owners retain workspace-create. Counted. + owner := activeUser(t, db, database.User{RBACRoles: []string{rbac.RoleOwner().Name}}) + member(t, db, org.ID, owner) + + // Members of an org that keeps organization-workspace-access in + // default_org_member_roles inherit workspace-create. Counted. + defaultOrg := dbgen.Organization(t, db, database.Organization{}) + defaultMember := activeUser(t, db, database.User{}) + member(t, db, defaultOrg.ID, defaultMember) + + count, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), db, authorizer) + require.NoError(t, err) + require.Equal(t, int64(4), count) + }) + + t.Run("MultiOrgSplitCapability", func(t *testing.T) { + // Users whose capability differs between their organizations: + // workspace-create in any one org is sufficient to be counted. + rbac.ReloadBuiltinRoles(&rbac.RoleOptions{MinimumImplicitMember: true}) + t.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + db, _ := dbtestutil.NewDB(t) + orgA := dbgen.Organization(t, db, database.Organization{}) + orgB := dbgen.Organization(t, db, database.Organization{}) + emptyDefaultRoles(t, db, orgA) + emptyDefaultRoles(t, db, orgB) + + // Gateway in org A, workspace-create in org B. Counted. + split := activeUser(t, db, database.User{}) + member(t, db, orgA.ID, split) + member(t, db, orgB.ID, split, rbac.RoleOrgWorkspaceAccess()) + + // The creation ban is scoped to org A and must not negate the + // org B grant. Counted. + bannedSplit := activeUser(t, db, database.User{}) + member(t, db, orgA.ID, bannedSplit, rbac.RoleOrgWorkspaceAccess(), rbac.RoleOrgWorkspaceCreationBan()) + member(t, db, orgB.ID, bannedSplit, rbac.RoleOrgWorkspaceAccess()) + + // Gateway in both orgs. Not counted. + gateway := activeUser(t, db, database.User{}) + member(t, db, orgA.ID, gateway) + member(t, db, orgB.ID, gateway) + + count, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), db, authorizer) + require.NoError(t, err) + require.Equal(t, int64(2), count) + }) + + t.Run("CustomOrgRole", func(t *testing.T) { + rbac.ReloadBuiltinRoles(&rbac.RoleOptions{MinimumImplicitMember: true}) + t.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + db, _ := dbtestutil.NewDB(t) + org := dbgen.Organization(t, db, database.Organization{}) + emptyDefaultRoles(t, db, org) + + creatorRole, err := db.InsertCustomRole(ctx, database.InsertCustomRoleParams{ + Name: "workspace-creator", + DisplayName: "Workspace Creator", + OrganizationID: uuid.NullUUID{UUID: org.ID, Valid: true}, + OrgPermissions: []database.CustomRolePermission{{ + ResourceType: rbac.ResourceWorkspace.Type, + Action: policy.ActionCreate, + }}, + }) + require.NoError(t, err) + + auditRole, err := db.InsertCustomRole(ctx, database.InsertCustomRoleParams{ + Name: "org-reader", + DisplayName: "Org Reader", + OrganizationID: uuid.NullUUID{UUID: org.ID, Valid: true}, + OrgPermissions: []database.CustomRolePermission{{ + ResourceType: rbac.ResourceOrganization.Type, + Action: policy.ActionRead, + }}, + }) + require.NoError(t, err) + + // Custom org role with workspace-create. Counted. + creator := activeUser(t, db, database.User{}) + member(t, db, org.ID, creator, creatorRole.Name) + + // Custom org role without workspace-create. Not counted. + reader := activeUser(t, db, database.User{}) + member(t, db, org.ID, reader, auditRole.Name) + + count, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), db, authorizer) + require.NoError(t, err) + require.Equal(t, int64(1), count) + }) + + t.Run("MalformedRoleNotCounted", func(t *testing.T) { + rbac.ReloadBuiltinRoles(nil) + t.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + db, _ := dbtestutil.NewDB(t) + org := dbgen.Organization(t, db, database.Organization{}) + + // Authorization fails closed on an unparsable stored role, so + // this user is not workspace-capable even though their org + // membership would otherwise qualify them. + corrupt := activeUser(t, db, database.User{RBACRoles: []string{"bad:role:extra"}}) + member(t, db, org.ID, corrupt) + + // The bad row must not fail the count for everyone else. + capable := activeUser(t, db, database.User{}) + member(t, db, org.ID, capable) + + count, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), db, authorizer) + require.NoError(t, err) + require.Equal(t, int64(1), count) + }) + + t.Run("EntitlementsAddonGate", func(t *testing.T) { + // Permission-based counting is gated on both the experiment and a + // valid license carrying the AI Governance addon. Without either, + // the legacy active user count applies. + rbac.ReloadBuiltinRoles(&rbac.RoleOptions{MinimumImplicitMember: true}) + t.Cleanup(func() { rbac.ReloadBuiltinRoles(nil) }) + + db, _ := dbtestutil.NewDB(t) + org := dbgen.Organization(t, db, database.Organization{}) + emptyDefaultRoles(t, db, org) + + gateway := activeUser(t, db, database.User{}) + member(t, db, org.ID, gateway) + wsUser := activeUser(t, db, database.User{}) + member(t, db, org.ID, wsUser, rbac.RoleOrgWorkspaceAccess()) + + enablements := map[codersdk.FeatureName]bool{} + experimentOn := codersdk.Experiments{codersdk.ExperimentWorkspaceCapableLicensing} + + // No license: legacy count, even with the experiment on. + entitlements, err := license.Entitlements(ctx, testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, authorizer, experimentOn) + require.NoError(t, err) + require.Equal(t, int64(2), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + + // License without the AI Governance addon: still the legacy count. + _, err = db.InsertLicense(ctx, database.InsertLicenseParams{ + JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }), + Exp: dbtime.Now().Add(time.Hour), + }) + require.NoError(t, err) + entitlements, err = license.Entitlements(ctx, testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, authorizer, experimentOn) + require.NoError(t, err) + require.Equal(t, int64(2), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + + // License with the AI Governance addon: only the workspace-capable + // user counts. + _, err = db.InsertLicense(ctx, database.InsertLicenseParams{ + JWT: coderdenttest.GenerateLicense(t, *(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).AIGovernanceAddon(10)), + Exp: dbtime.Now().Add(time.Hour), + }) + require.NoError(t, err) + entitlements, err = license.Entitlements(ctx, testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, authorizer, experimentOn) + require.NoError(t, err) + require.Empty(t, entitlements.Errors) + require.Equal(t, int64(1), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + + // Addon present but experiment off: legacy count. + entitlements, err = license.Entitlements(ctx, testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, authorizer, nil) + require.NoError(t, err) + require.Equal(t, int64(2), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + + // Addon present, experiment on, but no authorizer: fall back to the + // legacy count instead of failing. + entitlements, err = license.Entitlements(ctx, testutil.Logger(t), db, 1, 1, coderdenttest.Keys, enablements, nil, experimentOn) + require.NoError(t, err) + require.Equal(t, int64(2), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + }) + + t.Run("LicensesEntitlementsCountFn", func(t *testing.T) { + // Exercises LicensesEntitlements directly: the count function is + // only invoked when a valid license carries the addon, grace + // period licenses still gate the count, and count errors fall + // back to the legacy count with a recorded error. + now := time.Now() + enablements := map[codersdk.FeatureName]bool{} + + dbLicense := func(opts coderdenttest.LicenseOptions) database.License { + return database.License{ + UUID: uuid.New(), + JWT: coderdenttest.GenerateLicense(t, opts), + Exp: now.Add(time.Hour * 24 * 60), + } + } + addonLicense := func() database.License { + return dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).Valid(now).AIGovernanceAddon(10)) + } + + t.Run("NoAddonFnNotCalled", func(t *testing.T) { + licenses := []database.License{dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).Valid(now))} + entitlements, err := license.LicensesEntitlements(ctx, now, licenses, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + t.Fatal("count fn must not be called without the addon") + return 0, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(7), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + }) + + t.Run("AddonMissingDependenciesIgnored", func(t *testing.T) { + // A license carrying the addon without its required features + // records a validation error and the addon is skipped, so + // workspace-capable counting must not activate. + opts := (&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).Valid(now) + opts.Addons = append(opts.Addons, codersdk.AddonAIGovernance) + entitlements, err := license.LicensesEntitlements(ctx, now, []database.License{dbLicense(*opts)}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + t.Fatal("count fn must not be called when addon dependencies are unmet") + return 0, nil + }, + }) + require.NoError(t, err) + require.NotEmpty(t, entitlements.Errors) + require.Equal(t, int64(7), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + }) + + t.Run("ActiveModeIgnoresFn", func(t *testing.T) { + // UserCountingMode is authoritative: with the mode left at its + // active-users zero value, the counting function must not be + // called even though it is set and the addon is present. + entitlements, err := license.LicensesEntitlements(ctx, now, []database.License{addonLicense()}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + t.Fatal("count fn must not be called in active counting mode") + return 0, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(7), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(100), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + }) + + t.Run("AddonUsesFn", func(t *testing.T) { + entitlements, err := license.LicensesEntitlements(ctx, now, []database.License{addonLicense()}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + ActiveAISeatCount: 5, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 3, nil + }, + }) + require.NoError(t, err) + require.Empty(t, entitlements.Errors) + require.Equal(t, int64(3), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + // Permission-based counting applies to workspace seats only: + // AI Governance seats keep their own count and limit. + aiSeats := entitlements.Features[codersdk.FeatureAIGovernanceUserLimit] + require.Equal(t, int64(5), *aiSeats.Actual) + require.Equal(t, int64(10), *aiSeats.Limit) + // Under the limit: no user-limit warning, even though the + // legacy active user count would also have been under it. + for _, warning := range entitlements.Warnings { + require.NotContains(t, warning, "users but") + } + // A fully valid addon license must not warn about the + // counting-mode revert. + for _, warning := range entitlements.Warnings { + require.NotContains(t, warning, "fully expires") + } + }) + + t.Run("BestPairSelection", func(t *testing.T) { + // A deployment holding both an addon license and a non-addon + // license has two user_limit candidates, each evaluated with + // its own counting mode. Limits and modes never mix. + licenses := []database.License{ + dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 200}, + }).Valid(now)), + addonLicense(), // user_limit 100, AI Governance addon. + } + run := func(t *testing.T, activeUsers, capableUsers int64) codersdk.Entitlements { + entitlements, err := license.LicensesEntitlements(ctx, now, licenses, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: activeUsers, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return capableUsers, nil + }, + }) + require.NoError(t, err) + return entitlements + } + + t.Run("LegacyPairCompliant", func(t *testing.T) { + // 180 active <= 200 wins over 150 capable > 100: the + // non-addon license keeps the deployment compliant. + entitlements := run(t, 180, 150) + require.Equal(t, int64(180), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(200), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + for _, warning := range entitlements.Warnings { + require.NotContains(t, warning, "users but") + } + }) + + t.Run("AddonPairCompliant", func(t *testing.T) { + // 90 capable <= 100 wins over 250 active > 200: the addon + // license keeps the deployment compliant. + entitlements := run(t, 250, 90) + require.Equal(t, int64(90), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(100), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + for _, warning := range entitlements.Warnings { + require.NotContains(t, warning, "users but") + } + }) + + t.Run("NeitherPairCompliant", func(t *testing.T) { + // Both pairs over: the higher limit is reported, with the + // counting mode of its own license. + entitlements := run(t, 250, 150) + require.Equal(t, int64(250), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(200), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + require.Contains(t, entitlements.Warnings, + "Your deployment has 250 active users but is only licensed for 200.") + }) + + t.Run("GraceAddonCompliantBeatsEntitledOver", func(t *testing.T) { + // A grace-period addon pair that fits its count wins over an + // entitled non-addon pair that does not, carrying its grace + // entitlement and the revert warning with it. + licenses := []database.License{ + dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 200}, + }).Valid(now)), + dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).GracePeriod(now).AIGovernanceAddon(10)), + } + entitlements, err := license.LicensesEntitlements(ctx, now, licenses, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 250, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 90, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(90), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(100), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + require.Equal(t, codersdk.EntitlementGracePeriod, entitlements.Features[codersdk.FeatureUserLimit].Entitlement) + require.Contains(t, entitlements.Warnings, + "Your deployment has 90 workspace-capable users but the license with the limit 100 is expired.") + require.Contains(t, entitlements.Warnings, + "Your license with the AI Governance addon is expired. When it fully expires, all 250 active users will count toward the user limit instead of the 90 workspace-capable users.") + }) + + t.Run("EqualLimitsPreferAddon", func(t *testing.T) { + // Identical limit and entitlement on an addon and a + // non-addon license: the addon pair wins the tie, so the + // workspace-capable count is displayed. + licenses := []database.License{ + dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).Valid(now)), + addonLicense(), + } + entitlements, err := license.LicensesEntitlements(ctx, now, licenses, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 80, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 30, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(30), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(100), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + }) + + t.Run("TwoAddonCandidates", func(t *testing.T) { + // Two addon licenses: the entitled higher-limit pair fits + // the capable count and wins over the grace pair, and its + // presence suppresses the revert warning. + licenses := []database.License{ + dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).GracePeriod(now).AIGovernanceAddon(10)), + dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 300}, + }).Valid(now).AIGovernanceAddon(10)), + } + entitlements, err := license.LicensesEntitlements(ctx, now, licenses, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 500, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 150, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(150), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Equal(t, int64(300), *entitlements.Features[codersdk.FeatureUserLimit].Limit) + require.Equal(t, codersdk.EntitlementEntitled, entitlements.Features[codersdk.FeatureUserLimit].Entitlement) + for _, warning := range entitlements.Warnings { + require.NotContains(t, warning, "fully expires") + require.NotContains(t, warning, "users but") + } + }) + }) + + t.Run("ModeWithoutFnIsDevError", func(t *testing.T) { + _, err := license.LicensesEntitlements(ctx, now, []database.License{addonLicense()}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + }) + require.ErrorContains(t, err, "dev error") + }) + + t.Run("OverLimitWarnsWithCapableCount", func(t *testing.T) { + // The over-limit warning must report the workspace-capable + // count it was compared against, and say so, rather than + // claiming that many "active users" exist. + entitlements, err := license.LicensesEntitlements(ctx, now, []database.License{addonLicense()}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 150, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(150), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Contains(t, entitlements.Warnings, + "Your deployment has 150 workspace-capable users but is only licensed for 100.") + }) + + t.Run("GracePeriodAddonUsesFn", func(t *testing.T) { + // A license in its grace period still includes the addon, so + // counting must not revert until the license hard-expires. + licenses := []database.License{dbLicense(*(&coderdenttest.LicenseOptions{ + Features: license.Features{codersdk.FeatureUserLimit: 100}, + }).GracePeriod(now).AIGovernanceAddon(10))} + entitlements, err := license.LicensesEntitlements(ctx, now, licenses, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 3, nil + }, + }) + require.NoError(t, err) + require.Equal(t, int64(3), *entitlements.Features[codersdk.FeatureUserLimit].Actual) + require.Contains(t, entitlements.Warnings, + "Your deployment has 3 workspace-capable users but the license with the limit 100 is expired.") + // The revert warning gives admins the legacy count they will be + // measured by once the grace period ends. + require.Contains(t, entitlements.Warnings, + "Your license with the AI Governance addon is expired. When it fully expires, all 7 active users will count toward the user limit instead of the 3 workspace-capable users.") + }) + + t.Run("FnErrorPropagates", func(t *testing.T) { + // A failed capable count aborts the computation, matching the + // legacy active-user-count error semantics; the caller keeps + // the previous entitlements rather than seeing a silently + // different count. + _, err := license.LicensesEntitlements(ctx, now, []database.License{addonLicense()}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 0, xerrors.New("boom") + }, + }) + require.ErrorContains(t, err, "count workspace capable users") + require.ErrorContains(t, err, "boom") + }) + + t.Run("ContextCanceledBails", func(t *testing.T) { + _, err := license.LicensesEntitlements(ctx, now, []database.License{addonLicense()}, enablements, coderdenttest.Keys, license.FeatureArguments{ + ActiveUserCount: 7, + UserCountingMode: license.UserCountingModeWorkspaceCapable, + WorkspaceCapableUserCountFn: func(context.Context) (int64, error) { + return 0, context.Canceled + }, + }) + require.ErrorIs(t, err, context.Canceled) + }) + }) +} + +// TestCountWorkspaceCapableUsersErrors covers the count's database +// failure paths, which abort the count rather than skewing it. +// +// Reads the builtin role registry that sibling tests reload, so it must +// run serially. +// +//nolint:paralleltest +func TestCountWorkspaceCapableUsersErrors(t *testing.T) { + ctx := context.Background() + authorizer := rbac.NewCachingAuthorizer(prometheus.NewRegistry()) + + prefetchParams := database.CustomRolesParams{IncludeSystemRoles: true} + + t.Run("NilAuthorizer", func(t *testing.T) { + mDB := dbmock.NewMockStore(gomock.NewController(t)) + _, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), mDB, nil) + require.ErrorContains(t, err, "dev error") + }) + + t.Run("PrefetchError", func(t *testing.T) { + mDB := dbmock.NewMockStore(gomock.NewController(t)) + mDB.EXPECT().CustomRoles(gomock.Any(), prefetchParams).Return(nil, xerrors.New("boom")) + + _, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), mDB, authorizer) + require.ErrorContains(t, err, "prefetch custom roles") + }) + + t.Run("RolesQueryError", func(t *testing.T) { + mDB := dbmock.NewMockStore(gomock.NewController(t)) + mDB.EXPECT().CustomRoles(gomock.Any(), prefetchParams).Return([]database.CustomRole{}, nil) + mDB.EXPECT().GetActiveUsersAuthorizationRoles(gomock.Any()).Return(nil, xerrors.New("boom")) + + _, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), mDB, authorizer) + require.ErrorContains(t, err, "get active users authorization roles") + }) + + t.Run("ExpandLookupError", func(t *testing.T) { + // A custom role that was not prefetched (deleted, or created + // mid-count) is looked up individually; a database failure there + // aborts the count. + mDB := dbmock.NewMockStore(gomock.NewController(t)) + userID := uuid.New() + orgID := uuid.New() + mDB.EXPECT().CustomRoles(gomock.Any(), prefetchParams).Return([]database.CustomRole{}, nil) + mDB.EXPECT().GetActiveUsersAuthorizationRoles(gomock.Any()).Return([]database.GetActiveUsersAuthorizationRolesRow{{ + ID: userID, + Roles: []string{"member", "dangling-role:" + orgID.String()}, + }}, nil) + mDB.EXPECT().CustomRoles(gomock.Any(), gomock.Not(prefetchParams)).Return(nil, xerrors.New("boom")) + + _, err := license.CountWorkspaceCapableUsers(ctx, testutil.Logger(t), mDB, authorizer) + require.ErrorContains(t, err, "evaluate workspace-create for user "+userID.String()) + require.ErrorContains(t, err, "expand roles") + }) +} diff --git a/enterprise/coderd/license/userlimit_internal_test.go b/enterprise/coderd/license/userlimit_internal_test.go new file mode 100644 index 00000000000..fe0ca14628f --- /dev/null +++ b/enterprise/coderd/license/userlimit_internal_test.go @@ -0,0 +1,82 @@ +package license + +import ( + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/codersdk" +) + +func TestBetterUserLimit(t *testing.T) { + t.Parallel() + + cand := func(limit, count int64, entitlement codersdk.Entitlement, addon bool) resolvedCandidate { + return resolvedCandidate{ + userLimitCandidate: userLimitCandidate{limit: limit, entitlement: entitlement, aiGovernanceAddon: addon}, + count: count, + } + } + entitled := codersdk.EntitlementEntitled + grace := codersdk.EntitlementGracePeriod + + cases := []struct { + name string + a, b resolvedCandidate + want bool + }{ + { + name: "ComplianceBeatsEntitlement", + a: cand(200, 150, grace, false), + b: cand(100, 150, entitled, false), + want: true, + }, + { + name: "ComplianceBeatsHigherLimit", + a: cand(100, 90, entitled, true), + b: cand(200, 250, entitled, false), + want: true, + }, + { + name: "EntitlementBeatsLimitWhenBothCompliant", + a: cand(100, 50, entitled, false), + b: cand(200, 50, grace, false), + want: true, + }, + { + name: "HigherLimitWinsWhenBothCompliantAndEqualEntitlement", + a: cand(200, 50, entitled, false), + b: cand(100, 50, entitled, false), + want: true, + }, + { + name: "HigherLimitWinsWhenBothOver", + a: cand(200, 250, entitled, false), + b: cand(100, 150, entitled, true), + want: true, + }, + { + name: "AddonBreaksExactTies", + a: cand(100, 50, entitled, true), + b: cand(100, 80, entitled, false), + want: true, + }, + { + name: "EqualCandidatesAreNotBetter", + a: cand(100, 50, entitled, false), + b: cand(100, 50, entitled, false), + want: false, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tc.want, betterUserLimit(tc.a, tc.b)) + if tc.want { + require.False(t, betterUserLimit(tc.b, tc.a), + "strict ordering must not hold both ways") + } + }) + } +} diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 3d89b954e2c..df238f98d8c 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -4842,6 +4842,7 @@ export type Experiment = | "notifications" | "oauth2" | "workspace-build-updates" + | "workspace-capable-licensing" | "workspace-usage"; export const Experiments: Experiment[] = [ @@ -4856,6 +4857,7 @@ export const Experiments: Experiment[] = [ "notifications", "oauth2", "workspace-build-updates", + "workspace-capable-licensing", "workspace-usage", ];