Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 0 additions & 39 deletions coderd/database/dbauthz/dbauthz.go
Original file line number Diff line number Diff line change
Expand Up @@ -3878,14 +3878,6 @@ func (q *querier) GetDeploymentID(ctx context.Context) (string, error) {
return q.db.GetDeploymentID(ctx)
}

func (q *querier) GetDeploymentWorkspaceAgentStats(ctx context.Context, createdAfter time.Time) (database.GetDeploymentWorkspaceAgentStatsRow, error) {
return q.db.GetDeploymentWorkspaceAgentStats(ctx, createdAfter)
}

func (q *querier) GetDeploymentWorkspaceAgentUsageStats(ctx context.Context, createdAt time.Time) (database.GetDeploymentWorkspaceAgentUsageStatsRow, error) {
return q.db.GetDeploymentWorkspaceAgentUsageStats(ctx, createdAt)
}

func (q *querier) GetDeploymentWorkspaceStats(ctx context.Context) (database.GetDeploymentWorkspaceStatsRow, error) {
return q.db.GetDeploymentWorkspaceStats(ctx)
}
Expand Down Expand Up @@ -4877,14 +4869,6 @@ func (q *querier) GetTemplateInsightsByInterval(ctx context.Context, arg databas
return q.db.GetTemplateInsightsByInterval(ctx, arg)
}

func (q *querier) GetTemplateInsightsByTemplate(ctx context.Context, arg database.GetTemplateInsightsByTemplateParams) ([]database.GetTemplateInsightsByTemplateRow, error) {
// Only used by prometheus metrics collector. No need to check update template perms.
Comment thread
EhabY marked this conversation as resolved.
if err := q.authorizeContext(ctx, policy.ActionViewInsights, rbac.ResourceTemplate); err != nil {
return nil, err
}
return q.db.GetTemplateInsightsByTemplate(ctx, arg)
}

func (q *querier) GetTemplateParameterInsights(ctx context.Context, arg database.GetTemplateParameterInsightsParams) ([]database.GetTemplateParameterInsightsRow, error) {
if err := q.authorizeTemplateInsights(ctx, arg.TemplateIDs); err != nil {
return nil, err
Expand Down Expand Up @@ -5627,22 +5611,6 @@ func (q *querier) GetWorkspaceAgentScriptsByAgentIDs(ctx context.Context, ids []
return q.db.GetWorkspaceAgentScriptsByAgentIDs(ctx, ids)
}

func (q *querier) GetWorkspaceAgentStats(ctx context.Context, createdAfter time.Time) ([]database.GetWorkspaceAgentStatsRow, error) {
return q.db.GetWorkspaceAgentStats(ctx, createdAfter)
}

func (q *querier) GetWorkspaceAgentStatsAndLabels(ctx context.Context, createdAfter time.Time) ([]database.GetWorkspaceAgentStatsAndLabelsRow, error) {
return q.db.GetWorkspaceAgentStatsAndLabels(ctx, createdAfter)
}

func (q *querier) GetWorkspaceAgentUsageStats(ctx context.Context, createdAt time.Time) ([]database.GetWorkspaceAgentUsageStatsRow, error) {
return q.db.GetWorkspaceAgentUsageStats(ctx, createdAt)
}

func (q *querier) GetWorkspaceAgentUsageStatsAndLabels(ctx context.Context, createdAt time.Time) ([]database.GetWorkspaceAgentUsageStatsAndLabelsRow, error) {
return q.db.GetWorkspaceAgentUsageStatsAndLabels(ctx, createdAt)
}

func (q *querier) GetWorkspaceAgentsByInstanceID(ctx context.Context, authInstanceID string) ([]database.WorkspaceAgent, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err == nil {
return q.db.GetWorkspaceAgentsByInstanceID(ctx, authInstanceID)
Expand Down Expand Up @@ -9379,13 +9347,6 @@ func (q *querier) UpsertTelemetryItem(ctx context.Context, arg database.UpsertTe
return q.db.UpsertTelemetryItem(ctx, arg)
}

func (q *querier) UpsertTemplateUsageStats(ctx context.Context) error {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil {
return err
}
return q.db.UpsertTemplateUsageStats(ctx)
}

func (q *querier) UpsertUserAIBudgetOverride(ctx context.Context, arg database.UpsertUserAIBudgetOverrideParams) (database.UserAIBudgetOverride, error) {
// Setting a user's AI budget override affects both the user (their
// per-user spend cap) and the group (spend attribution).
Expand Down
159 changes: 138 additions & 21 deletions coderd/database/dbauthz/dbauthz_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"database/sql"
"encoding/json"
"fmt"
"maps"
"net"
"reflect"
"strconv"
Expand Down Expand Up @@ -3009,7 +3010,7 @@ func (s *MethodTestSuite) TestTemplate() {
check.Args(arg).Asserts(rbac.ResourceTemplate, policy.ActionViewInsights)
}))
s.Run("GetTemplateInsightsByTemplate", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
arg := database.GetTemplateInsightsByTemplateParams{}
arg := database.GetTemplateInsightsByTemplateParams{AppFamilies: codersdk.SessionCountAppFamiliesJSON()}
dbm.EXPECT().GetTemplateInsightsByTemplate(gomock.Any(), arg).Return([]database.GetTemplateInsightsByTemplateRow{}, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceTemplate, policy.ActionViewInsights)
}))
Expand All @@ -3034,8 +3035,9 @@ func (s *MethodTestSuite) TestTemplate() {
check.Args(arg).Asserts(rbac.ResourceTemplate, policy.ActionViewInsights).Returns([]database.TemplateUsageStat{})
}))
s.Run("UpsertTemplateUsageStats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
dbm.EXPECT().UpsertTemplateUsageStats(gomock.Any()).Return(nil).AnyTimes()
check.Asserts(rbac.ResourceSystem, policy.ActionUpdate)
arg := codersdk.SessionCountAppFamiliesJSON()
dbm.EXPECT().UpsertTemplateUsageStats(gomock.Any(), arg).Return(nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceSystem, policy.ActionUpdate)
}))
s.Run("UpdatePresetsLastInvalidatedAt", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
t1 := testutil.Fake(s.T(), faker, database.Template{})
Expand Down Expand Up @@ -5534,14 +5536,20 @@ func (s *MethodTestSuite) TestSystemFunctions() {
check.Args("foo").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
}))
s.Run("GetDeploymentWorkspaceAgentStats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
t := time.Time{}
dbm.EXPECT().GetDeploymentWorkspaceAgentStats(gomock.Any(), t).Return(database.GetDeploymentWorkspaceAgentStatsRow{}, nil).AnyTimes()
check.Args(t).Asserts()
arg := database.GetDeploymentWorkspaceAgentStatsParams{
CreatedAt: time.Time{},
AppFamilies: codersdk.SessionCountAppFamiliesJSON(),
}
dbm.EXPECT().GetDeploymentWorkspaceAgentStats(gomock.Any(), arg).Return(database.GetDeploymentWorkspaceAgentStatsRow{}, nil).AnyTimes()
check.Args(arg).Asserts()
}))
s.Run("GetDeploymentWorkspaceAgentUsageStats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
t := time.Time{}
dbm.EXPECT().GetDeploymentWorkspaceAgentUsageStats(gomock.Any(), t).Return(database.GetDeploymentWorkspaceAgentUsageStatsRow{}, nil).AnyTimes()
check.Args(t).Asserts()
arg := database.GetDeploymentWorkspaceAgentUsageStatsParams{
CreatedAt: time.Time{},
AppFamilies: codersdk.SessionCountAppFamiliesJSON(),
}
dbm.EXPECT().GetDeploymentWorkspaceAgentUsageStats(gomock.Any(), arg).Return(database.GetDeploymentWorkspaceAgentUsageStatsRow{}, nil).AnyTimes()
check.Args(arg).Asserts()
}))
s.Run("GetDeploymentWorkspaceStats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
dbm.EXPECT().GetDeploymentWorkspaceStats(gomock.Any()).Return(database.GetDeploymentWorkspaceStatsRow{}, nil).AnyTimes()
Expand Down Expand Up @@ -5569,24 +5577,36 @@ func (s *MethodTestSuite) TestSystemFunctions() {
check.Args(arg).Asserts(rbac.ResourceSystem, policy.ActionUpdate)
}))
s.Run("GetWorkspaceAgentStatsAndLabels", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
t := time.Time{}
dbm.EXPECT().GetWorkspaceAgentStatsAndLabels(gomock.Any(), t).Return([]database.GetWorkspaceAgentStatsAndLabelsRow{}, nil).AnyTimes()
check.Args(t).Asserts()
arg := database.GetWorkspaceAgentStatsAndLabelsParams{
CreatedAt: time.Time{},
AppFamilies: codersdk.SessionCountAppFamiliesJSON(),
}
dbm.EXPECT().GetWorkspaceAgentStatsAndLabels(gomock.Any(), arg).Return([]database.GetWorkspaceAgentStatsAndLabelsRow{}, nil).AnyTimes()
check.Args(arg).Asserts()
}))
s.Run("GetWorkspaceAgentUsageStatsAndLabels", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
t := time.Time{}
dbm.EXPECT().GetWorkspaceAgentUsageStatsAndLabels(gomock.Any(), t).Return([]database.GetWorkspaceAgentUsageStatsAndLabelsRow{}, nil).AnyTimes()
check.Args(t).Asserts()
arg := database.GetWorkspaceAgentUsageStatsAndLabelsParams{
CreatedAt: time.Time{},
AppFamilies: codersdk.SessionCountAppFamiliesJSON(),
}
dbm.EXPECT().GetWorkspaceAgentUsageStatsAndLabels(gomock.Any(), arg).Return([]database.GetWorkspaceAgentUsageStatsAndLabelsRow{}, nil).AnyTimes()
check.Args(arg).Asserts()
}))
s.Run("GetWorkspaceAgentStats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
t := time.Time{}
dbm.EXPECT().GetWorkspaceAgentStats(gomock.Any(), t).Return([]database.GetWorkspaceAgentStatsRow{}, nil).AnyTimes()
check.Args(t).Asserts()
arg := database.GetWorkspaceAgentStatsParams{
CreatedAt: time.Time{},
AppFamilies: codersdk.SessionCountAppFamiliesJSON(),
}
dbm.EXPECT().GetWorkspaceAgentStats(gomock.Any(), arg).Return([]database.GetWorkspaceAgentStatsRow{}, nil).AnyTimes()
check.Args(arg).Asserts()
}))
s.Run("GetWorkspaceAgentUsageStats", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
t := time.Time{}
dbm.EXPECT().GetWorkspaceAgentUsageStats(gomock.Any(), t).Return([]database.GetWorkspaceAgentUsageStatsRow{}, nil).AnyTimes()
check.Args(t).Asserts()
arg := database.GetWorkspaceAgentUsageStatsParams{
CreatedAt: time.Time{},
AppFamilies: codersdk.SessionCountAppFamiliesJSON(),
}
dbm.EXPECT().GetWorkspaceAgentUsageStats(gomock.Any(), arg).Return([]database.GetWorkspaceAgentUsageStatsRow{}, nil).AnyTimes()
check.Args(arg).Asserts()
}))
s.Run("GetWorkspaceProxyByHostname", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
p := testutil.Fake(s.T(), faker, database.WorkspaceProxy{WildcardHostname: "*.example.com"})
Expand Down Expand Up @@ -7947,3 +7967,100 @@ func TestAsExternalAuthChecker(t *testing.T) {
}
})
}

// TestSessionCountAppFamiliesRequired ensures the session count queries fail
// loudly when the app family registry is empty, so a forgotten parameter
// surfaces as an error instead of silently dropping every family's sessions
// from usage reporting.
func TestSessionCountAppFamiliesRequired(t *testing.T) {
t.Parallel()

ctrl := gomock.NewController(t)
defer ctrl.Finish()
dbm := dbmock.NewMockStore(ctrl)
dbm.EXPECT().Wrappers().Return([]string{}).AnyTimes()
q := dbauthz.New(dbm, &coderdtest.RecordingAuthorizer{Wrapped: &coderdtest.FakeAuthorizer{}}, slog.Make(), coderdtest.AccessControlStorePointer())
ctx := dbauthz.As(context.Background(), coderdtest.RandomRBACSubject())

_, err := q.GetDeploymentWorkspaceAgentStats(ctx, database.GetDeploymentWorkspaceAgentStatsParams{})
require.ErrorContains(t, err, "developer error")
_, err = q.GetDeploymentWorkspaceAgentUsageStats(ctx, database.GetDeploymentWorkspaceAgentUsageStatsParams{})
require.ErrorContains(t, err, "developer error")
_, err = q.GetWorkspaceAgentStats(ctx, database.GetWorkspaceAgentStatsParams{})
require.ErrorContains(t, err, "developer error")
_, err = q.GetWorkspaceAgentStatsAndLabels(ctx, database.GetWorkspaceAgentStatsAndLabelsParams{})
require.ErrorContains(t, err, "developer error")
_, err = q.GetWorkspaceAgentUsageStats(ctx, database.GetWorkspaceAgentUsageStatsParams{})
require.ErrorContains(t, err, "developer error")
_, err = q.GetWorkspaceAgentUsageStatsAndLabels(ctx, database.GetWorkspaceAgentUsageStatsAndLabelsParams{})
require.ErrorContains(t, err, "developer error")
_, err = q.GetTemplateInsightsByTemplate(ctx, database.GetTemplateInsightsByTemplateParams{})
require.ErrorContains(t, err, "developer error")
err = q.UpsertTemplateUsageStats(ctx, nil)
require.ErrorContains(t, err, "developer error")
}

// TestSessionCountAppFamiliesMustMatchQueries covers registries that are
// present but wrong. Each query hardcodes one probe per family, so a registry
// whose keys drifted from codersdk.AttributedAppFamilies would report zero
// for the affected family instead of failing.
func TestSessionCountAppFamiliesMustMatchQueries(t *testing.T) {
t.Parallel()

valid := map[codersdk.AppFamilyName][]string{}
for _, family := range codersdk.AttributedAppFamilies() {
valid[family] = []string{string(family)}
}
without := func(drop codersdk.AppFamilyName) json.RawMessage {
families := maps.Clone(valid)
delete(families, drop)
return mustMarshalAppFamilies(t, families)
}

for _, tc := range []struct {
name string
appFamilies json.RawMessage
errContains string
}{
{"EmptyObject", json.RawMessage(`{}`), `missing family "vscode"`},
{"JSONNull", json.RawMessage(`null`), `missing family "vscode"`},
{"NotAnObject", json.RawMessage(`["vscode"]`), "must be a JSON object"},
{"MissingFamily", without(codersdk.AppFamilySSH), `missing family "ssh"`},
{"EmptyAppNames", mustMarshalAppFamilies(t, map[codersdk.AppFamilyName][]string{
codersdk.AppFamilyVSCode: {"vscode"},
codersdk.AppFamilyJetBrains: {"jetbrains"},
codersdk.AppFamilySSH: {},
codersdk.AppFamilyReconnectingPTY: {"reconnecting_pty"},
}), `no app names for family "ssh"`},
{"UnknownFamily", mustMarshalAppFamilies(t, map[codersdk.AppFamilyName][]string{
codersdk.AppFamilyVSCode: {"vscode"},
codersdk.AppFamilyJetBrains: {"jetbrains"},
codersdk.AppFamilySSH: {"ssh"},
codersdk.AppFamilyReconnectingPTY: {"reconnecting_pty"},
"emacs": {"emacs"},
}), `has family "emacs"`},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

ctrl := gomock.NewController(t)
defer ctrl.Finish()
dbm := dbmock.NewMockStore(ctrl)
dbm.EXPECT().Wrappers().Return([]string{}).AnyTimes()
q := dbauthz.New(dbm, &coderdtest.RecordingAuthorizer{Wrapped: &coderdtest.FakeAuthorizer{}}, slog.Make(), coderdtest.AccessControlStorePointer())
ctx := dbauthz.As(context.Background(), coderdtest.RandomRBACSubject())

_, err := q.GetDeploymentWorkspaceAgentStats(ctx, database.GetDeploymentWorkspaceAgentStatsParams{AppFamilies: tc.appFamilies})
require.ErrorContains(t, err, tc.errContains)
err = q.UpsertTemplateUsageStats(ctx, tc.appFamilies)
require.ErrorContains(t, err, tc.errContains)
})
}
}

func mustMarshalAppFamilies(t *testing.T, families map[codersdk.AppFamilyName][]string) json.RawMessage {
t.Helper()
raw, err := json.Marshal(families)
require.NoError(t, err)
return raw
}
118 changes: 118 additions & 0 deletions coderd/database/dbauthz/sessioncountparams.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
package dbauthz

import (
"context"
"encoding/json"
"slices"

"golang.org/x/xerrors"

"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/rbac/policy"
"github.com/coder/coder/v2/codersdk"
)

// The session count read queries take the app family attribution registry as
// a single jsonb parameter. Every query hardcodes one probe per family, so a
// registry that is empty, malformed, or keyed differently from
// codersdk.AttributedAppFamilies would still run and return zero counts for
// the affected families, silently dropping sessions from deployment stats,
// insights, Prometheus, and telemetry. dbauthz wraps every production store,
// including the transaction stores used by the rollup, so validating here
// makes a wrong registry fail the call loudly instead of failing the data
// quietly. These methods override the generated ones in dbauthz.go;
// scripts/dbgen preserves methods defined outside that file.

// validateSessionCountAppFamilies checks that the registry has exactly the
// families the queries probe, each with at least one app name.
func validateSessionCountAppFamilies(appFamilies json.RawMessage) error {
if len(appFamilies) == 0 {
return xerrors.New("developer error: session count app families must not be empty, populate them with codersdk.SessionCountAppFamiliesJSON()")
}

var families map[codersdk.AppFamilyName][]string
if err := json.Unmarshal(appFamilies, &families); err != nil {
return xerrors.Errorf("developer error: session count app families must be a JSON object of family to app names, populate them with codersdk.SessionCountAppFamiliesJSON(): %w", err)
}

required := codersdk.AttributedAppFamilies()
for _, family := range required {
appNames, ok := families[family]
if !ok {
return xerrors.Errorf("developer error: session count app families is missing family %q, which the queries probe; populate them with codersdk.SessionCountAppFamiliesJSON()", family)
}
if len(appNames) == 0 {
return xerrors.Errorf("developer error: session count app families has no app names for family %q, so its sessions would go uncounted", family)
}
}
for family := range families {
if !slices.Contains(required, family) {
return xerrors.Errorf("developer error: session count app families has family %q, which no query probes; add a probe per query or drop it from codersdk.AttributedAppFamilies", family)
}
}
return nil
}

func (q *querier) GetDeploymentWorkspaceAgentStats(ctx context.Context, arg database.GetDeploymentWorkspaceAgentStatsParams) (database.GetDeploymentWorkspaceAgentStatsRow, error) {
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return database.GetDeploymentWorkspaceAgentStatsRow{}, err
}
return q.db.GetDeploymentWorkspaceAgentStats(ctx, arg)
}

func (q *querier) GetDeploymentWorkspaceAgentUsageStats(ctx context.Context, arg database.GetDeploymentWorkspaceAgentUsageStatsParams) (database.GetDeploymentWorkspaceAgentUsageStatsRow, error) {
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return database.GetDeploymentWorkspaceAgentUsageStatsRow{}, err
}
return q.db.GetDeploymentWorkspaceAgentUsageStats(ctx, arg)
}

func (q *querier) GetWorkspaceAgentStats(ctx context.Context, arg database.GetWorkspaceAgentStatsParams) ([]database.GetWorkspaceAgentStatsRow, error) {
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return nil, err
}
return q.db.GetWorkspaceAgentStats(ctx, arg)
}

func (q *querier) GetWorkspaceAgentStatsAndLabels(ctx context.Context, arg database.GetWorkspaceAgentStatsAndLabelsParams) ([]database.GetWorkspaceAgentStatsAndLabelsRow, error) {
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return nil, err
}
return q.db.GetWorkspaceAgentStatsAndLabels(ctx, arg)
}

func (q *querier) GetWorkspaceAgentUsageStats(ctx context.Context, arg database.GetWorkspaceAgentUsageStatsParams) ([]database.GetWorkspaceAgentUsageStatsRow, error) {
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return nil, err
}
return q.db.GetWorkspaceAgentUsageStats(ctx, arg)
}

func (q *querier) GetWorkspaceAgentUsageStatsAndLabels(ctx context.Context, arg database.GetWorkspaceAgentUsageStatsAndLabelsParams) ([]database.GetWorkspaceAgentUsageStatsAndLabelsRow, error) {
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return nil, err
}
return q.db.GetWorkspaceAgentUsageStatsAndLabels(ctx, arg)
}

func (q *querier) GetTemplateInsightsByTemplate(ctx context.Context, arg database.GetTemplateInsightsByTemplateParams) ([]database.GetTemplateInsightsByTemplateRow, error) {
// Only used by prometheus metrics collector. No need to check update template perms.
if err := q.authorizeContext(ctx, policy.ActionViewInsights, rbac.ResourceTemplate); err != nil {
return nil, err
}
if err := validateSessionCountAppFamilies(arg.AppFamilies); err != nil {
return nil, err
}
return q.db.GetTemplateInsightsByTemplate(ctx, arg)
}

func (q *querier) UpsertTemplateUsageStats(ctx context.Context, appFamilies json.RawMessage) error {
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil {
return err
}
if err := validateSessionCountAppFamilies(appFamilies); err != nil {
return err
}
return q.db.UpsertTemplateUsageStats(ctx, appFamilies)
}
Loading
Loading