mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add managed agent license limit checks (#18937)
- Adds a query for counting managed agent workspace builds between two timestamps - The "Actual" field in the feature entitlement for managed agents is now populated with the value read from the database - The wsbuilder package now validates AI agent usage against the limit when a license is installed Closes coder/internal#777
This commit is contained in:
@@ -2193,6 +2193,14 @@ func (q *querier) GetLogoURL(ctx context.Context) (string, error) {
|
||||
return q.db.GetLogoURL(ctx)
|
||||
}
|
||||
|
||||
func (q *querier) GetManagedAgentCount(ctx context.Context, arg database.GetManagedAgentCountParams) (int64, error) {
|
||||
// Must be able to read all workspaces to check usage.
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceWorkspace); err != nil {
|
||||
return 0, xerrors.Errorf("authorize read all workspaces: %w", err)
|
||||
}
|
||||
return q.db.GetManagedAgentCount(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetNotificationMessagesByStatus(ctx context.Context, arg database.GetNotificationMessagesByStatusParams) ([]database.NotificationMessage, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceNotificationMessage); err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -17,20 +17,18 @@ import (
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/notifications"
|
||||
"github.com/coder/coder/v2/coderd/rbac/policy"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/coderd/notifications"
|
||||
"github.com/coder/coder/v2/coderd/rbac"
|
||||
"github.com/coder/coder/v2/coderd/rbac/policy"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/provisionersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
@@ -903,6 +901,14 @@ func (s *MethodTestSuite) TestLicense() {
|
||||
require.NoError(s.T(), err)
|
||||
check.Args().Asserts().Returns("value")
|
||||
}))
|
||||
s.Run("GetManagedAgentCount", s.Subtest(func(db database.Store, check *expects) {
|
||||
start := dbtime.Now()
|
||||
end := start.Add(time.Hour)
|
||||
check.Args(database.GetManagedAgentCountParams{
|
||||
StartTime: start,
|
||||
EndTime: end,
|
||||
}).Asserts(rbac.ResourceWorkspace, policy.ActionRead).Returns(int64(0))
|
||||
}))
|
||||
}
|
||||
|
||||
func (s *MethodTestSuite) TestOrganization() {
|
||||
|
||||
@@ -964,6 +964,13 @@ func (m queryMetricsStore) GetLogoURL(ctx context.Context) (string, error) {
|
||||
return url, err
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetManagedAgentCount(ctx context.Context, arg database.GetManagedAgentCountParams) (int64, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetManagedAgentCount(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetManagedAgentCount").Observe(time.Since(start).Seconds())
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetNotificationMessagesByStatus(ctx context.Context, arg database.GetNotificationMessagesByStatusParams) ([]database.NotificationMessage, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetNotificationMessagesByStatus(ctx, arg)
|
||||
|
||||
@@ -2012,6 +2012,21 @@ func (mr *MockStoreMockRecorder) GetLogoURL(ctx any) *gomock.Call {
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetLogoURL", reflect.TypeOf((*MockStore)(nil).GetLogoURL), ctx)
|
||||
}
|
||||
|
||||
// GetManagedAgentCount mocks base method.
|
||||
func (m *MockStore) GetManagedAgentCount(ctx context.Context, arg database.GetManagedAgentCountParams) (int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetManagedAgentCount", ctx, arg)
|
||||
ret0, _ := ret[0].(int64)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetManagedAgentCount indicates an expected call of GetManagedAgentCount.
|
||||
func (mr *MockStoreMockRecorder) GetManagedAgentCount(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetManagedAgentCount", reflect.TypeOf((*MockStore)(nil).GetManagedAgentCount), ctx, arg)
|
||||
}
|
||||
|
||||
// GetNotificationMessagesByStatus mocks base method.
|
||||
func (m *MockStore) GetNotificationMessagesByStatus(ctx context.Context, arg database.GetNotificationMessagesByStatusParams) ([]database.NotificationMessage, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -216,6 +216,8 @@ type sqlcQuerier interface {
|
||||
GetLicenseByID(ctx context.Context, id int32) (License, error)
|
||||
GetLicenses(ctx context.Context) ([]License, error)
|
||||
GetLogoURL(ctx context.Context) (string, error)
|
||||
// This isn't strictly a license query, but it's related to license enforcement.
|
||||
GetManagedAgentCount(ctx context.Context, arg GetManagedAgentCountParams) (int64, error)
|
||||
GetNotificationMessagesByStatus(ctx context.Context, arg GetNotificationMessagesByStatusParams) ([]NotificationMessage, error)
|
||||
// Fetch the notification report generator log indicating recent activity.
|
||||
GetNotificationReportGeneratorLogByTemplate(ctx context.Context, templateID uuid.UUID) (NotificationReportGeneratorLog, error)
|
||||
|
||||
@@ -4286,6 +4286,44 @@ func (q *sqlQuerier) GetLicenses(ctx context.Context) ([]License, error) {
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const getManagedAgentCount = `-- name: GetManagedAgentCount :one
|
||||
SELECT
|
||||
COUNT(DISTINCT wb.id) AS count
|
||||
FROM
|
||||
workspace_builds AS wb
|
||||
JOIN
|
||||
provisioner_jobs AS pj
|
||||
ON
|
||||
wb.job_id = pj.id
|
||||
WHERE
|
||||
wb.transition = 'start'::workspace_transition
|
||||
AND wb.has_ai_task = true
|
||||
-- Only count jobs that are pending, running or succeeded. Other statuses
|
||||
-- like cancel(ed|ing), failed or unknown are not considered as managed
|
||||
-- agent usage. These workspace builds are typically unusable anyway.
|
||||
AND pj.job_status IN (
|
||||
'pending'::provisioner_job_status,
|
||||
'running'::provisioner_job_status,
|
||||
'succeeded'::provisioner_job_status
|
||||
)
|
||||
-- Jobs are counted at the time they are created, not when they are
|
||||
-- completed, as pending jobs haven't completed yet.
|
||||
AND wb.created_at BETWEEN $1::timestamptz AND $2::timestamptz
|
||||
`
|
||||
|
||||
type GetManagedAgentCountParams struct {
|
||||
StartTime time.Time `db:"start_time" json:"start_time"`
|
||||
EndTime time.Time `db:"end_time" json:"end_time"`
|
||||
}
|
||||
|
||||
// This isn't strictly a license query, but it's related to license enforcement.
|
||||
func (q *sqlQuerier) GetManagedAgentCount(ctx context.Context, arg GetManagedAgentCountParams) (int64, error) {
|
||||
row := q.db.QueryRowContext(ctx, getManagedAgentCount, arg.StartTime, arg.EndTime)
|
||||
var count int64
|
||||
err := row.Scan(&count)
|
||||
return count, err
|
||||
}
|
||||
|
||||
const getUnexpiredLicenses = `-- name: GetUnexpiredLicenses :many
|
||||
SELECT id, uploaded_at, jwt, exp, uuid
|
||||
FROM licenses
|
||||
|
||||
@@ -35,3 +35,28 @@ DELETE
|
||||
FROM licenses
|
||||
WHERE id = $1
|
||||
RETURNING id;
|
||||
|
||||
-- name: GetManagedAgentCount :one
|
||||
-- This isn't strictly a license query, but it's related to license enforcement.
|
||||
SELECT
|
||||
COUNT(DISTINCT wb.id) AS count
|
||||
FROM
|
||||
workspace_builds AS wb
|
||||
JOIN
|
||||
provisioner_jobs AS pj
|
||||
ON
|
||||
wb.job_id = pj.id
|
||||
WHERE
|
||||
wb.transition = 'start'::workspace_transition
|
||||
AND wb.has_ai_task = true
|
||||
-- Only count jobs that are pending, running or succeeded. Other statuses
|
||||
-- like cancel(ed|ing), failed or unknown are not considered as managed
|
||||
-- agent usage. These workspace builds are typically unusable anyway.
|
||||
AND pj.job_status IN (
|
||||
'pending'::provisioner_job_status,
|
||||
'running'::provisioner_job_status,
|
||||
'succeeded'::provisioner_job_status
|
||||
)
|
||||
-- Jobs are counted at the time they are created, not when they are
|
||||
-- completed, as pending jobs haven't completed yet.
|
||||
AND wb.created_at BETWEEN @start_time::timestamptz AND @end_time::timestamptz;
|
||||
|
||||
Reference in New Issue
Block a user