chore: wire up usage tracking for managed agents (#19096)

Wires up the usage collector and publisher to coderd.

Relates to coder/internal#814
This commit is contained in:
Dean Sheather
2025-08-20 23:38:09 +10:00
committed by GitHub
parent dd867bd743
commit 6eb02d1c2a
43 changed files with 540 additions and 346 deletions
@@ -29,6 +29,7 @@ import (
"cdr.dev/slog"
"github.com/coder/coder/v2/coderd/usage"
"github.com/coder/coder/v2/coderd/util/slice"
"github.com/coder/coder/v2/codersdk/drpcsdk"
@@ -121,6 +122,7 @@ type server struct {
DeploymentValues *codersdk.DeploymentValues
NotificationsEnqueuer notifications.Enqueuer
PrebuildsOrchestrator *atomic.Pointer[prebuilds.ReconciliationOrchestrator]
UsageInserter *atomic.Pointer[usage.Inserter]
OIDCConfig promoauth.OAuth2Config
@@ -174,6 +176,7 @@ func NewServer(
auditor *atomic.Pointer[audit.Auditor],
templateScheduleStore *atomic.Pointer[schedule.TemplateScheduleStore],
userQuietHoursScheduleStore *atomic.Pointer[schedule.UserQuietHoursScheduleStore],
usageInserter *atomic.Pointer[usage.Inserter],
deploymentValues *codersdk.DeploymentValues,
options Options,
enqueuer notifications.Enqueuer,
@@ -195,6 +198,9 @@ func NewServer(
if userQuietHoursScheduleStore == nil {
return nil, xerrors.New("userQuietHoursScheduleStore is nil")
}
if usageInserter == nil {
return nil, xerrors.New("usageCollector is nil")
}
if deploymentValues == nil {
return nil, xerrors.New("deploymentValues is nil")
}
@@ -244,6 +250,7 @@ func NewServer(
heartbeatInterval: options.HeartbeatInterval,
heartbeatFn: options.HeartbeatFn,
PrebuildsOrchestrator: prebuildsOrchestrator,
UsageInserter: usageInserter,
}
if s.heartbeatFn == nil {
@@ -2030,6 +2037,20 @@ func (s *server) completeWorkspaceBuildJob(ctx context.Context, job database.Pro
sidebarAppID = uuid.NullUUID{}
}
if hasAITask && workspaceBuild.Transition == database.WorkspaceTransitionStart {
// Insert usage event for managed agents.
usageInserter := s.UsageInserter.Load()
if usageInserter != nil {
event := usage.DCManagedAgentsV1{
Count: 1,
}
err = (*usageInserter).InsertDiscreteUsageEvent(ctx, db, event)
if err != nil {
return xerrors.Errorf("insert %q event: %w", event.EventType(), err)
}
}
}
hasExternalAgent := false
for _, resource := range jobType.WorkspaceBuild.Resources {
if resource.Type == "coder_external_agent" {
@@ -16,6 +16,7 @@ import (
"time"
"github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/trace"
@@ -30,7 +31,9 @@ import (
"github.com/coder/coder/v2/buildinfo"
"github.com/coder/coder/v2/coderd/audit"
"github.com/coder/coder/v2/coderd/coderdtest"
"github.com/coder/coder/v2/coderd/database"
"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"
@@ -44,6 +47,7 @@ import (
"github.com/coder/coder/v2/coderd/schedule"
"github.com/coder/coder/v2/coderd/schedule/cron"
"github.com/coder/coder/v2/coderd/telemetry"
"github.com/coder/coder/v2/coderd/usage"
"github.com/coder/coder/v2/coderd/wspubsub"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/codersdk/agentsdk"
@@ -67,6 +71,13 @@ func testUserQuietHoursScheduleStore() *atomic.Pointer[schedule.UserQuietHoursSc
return ptr
}
func testUsageInserter() *atomic.Pointer[usage.Inserter] {
ptr := &atomic.Pointer[usage.Inserter]{}
inserter := usage.NewAGPLInserter()
ptr.Store(&inserter)
return ptr
}
func TestAcquireJob_LongPoll(t *testing.T) {
t.Parallel()
//nolint:dogsled
@@ -681,12 +692,20 @@ func TestUpdateJob(t *testing.T) {
t.Run("NotRunning", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, nil)
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
Input: json.RawMessage("{}"),
ID: version.JobID,
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: version.ID,
})),
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
@@ -700,12 +719,20 @@ func TestUpdateJob(t *testing.T) {
t.Run("NotOwner", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, nil)
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
Input: json.RawMessage("{}"),
ID: version.JobID,
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: version.ID,
})),
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
@@ -730,38 +757,57 @@ func TestUpdateJob(t *testing.T) {
require.ErrorContains(t, err, "you don't own this job")
})
setupJob := func(t *testing.T, db database.Store, srvID, orgID uuid.UUID, tags database.StringMap) uuid.UUID {
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
OrganizationID: orgID,
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionImport,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: json.RawMessage("{}"),
Tags: tags,
})
setupJob := func(t *testing.T, db database.Store, srvID, orgID uuid.UUID, tags database.StringMap) (templateVersionID, jobID uuid.UUID) {
templateVersionID = uuid.New()
jobID = uuid.New()
err := db.InTx(func(db database.Store) error {
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
ID: templateVersionID,
CreatedBy: user.ID,
OrganizationID: orgID,
JobID: jobID,
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: version.JobID,
OrganizationID: orgID,
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionImport,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: version.ID,
})),
Tags: tags,
})
if err != nil {
return xerrors.Errorf("insert provisioner job: %w", err)
}
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
WorkerID: uuid.NullUUID{
UUID: srvID,
Valid: true,
},
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
StartedAt: sql.NullTime{
Time: dbtime.Now(),
Valid: true,
},
OrganizationID: orgID,
ProvisionerTags: must(json.Marshal(job.Tags)),
})
if err != nil {
return xerrors.Errorf("acquire provisioner job: %w", err)
}
return nil
}, nil)
require.NoError(t, err)
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
WorkerID: uuid.NullUUID{
UUID: srvID,
Valid: true,
},
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
StartedAt: sql.NullTime{
Time: dbtime.Now(),
Valid: true,
},
OrganizationID: orgID,
ProvisionerTags: must(json.Marshal(job.Tags)),
})
require.NoError(t, err)
return job.ID
return templateVersionID, jobID
}
t.Run("Success", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, err := srv.UpdateJob(ctx, &proto.UpdateJobRequest{
JobId: job.String(),
})
@@ -771,7 +817,7 @@ func TestUpdateJob(t *testing.T) {
t.Run("Logs", func(t *testing.T) {
t.Parallel()
srv, db, ps, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
published := make(chan struct{})
@@ -796,23 +842,14 @@ func TestUpdateJob(t *testing.T) {
t.Run("Readme", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
versionID := uuid.New()
user := dbgen.User(t, db, database.User{})
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
ID: versionID,
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: job,
})
require.NoError(t, err)
_, err = srv.UpdateJob(ctx, &proto.UpdateJobRequest{
templateVersionID, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, err := srv.UpdateJob(ctx, &proto.UpdateJobRequest{
JobId: job.String(),
Readme: []byte("# hello world"),
})
require.NoError(t, err)
version, err := db.GetTemplateVersionByID(ctx, versionID)
version, err := db.GetTemplateVersionByID(ctx, templateVersionID)
require.NoError(t, err)
require.Equal(t, "# hello world", version.Readme)
})
@@ -825,16 +862,7 @@ func TestUpdateJob(t *testing.T) {
defer cancel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
versionID := uuid.New()
user := dbgen.User(t, db, database.User{})
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
ID: versionID,
CreatedBy: user.ID,
JobID: job,
OrganizationID: pd.OrganizationID,
})
require.NoError(t, err)
templateVersionID, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
firstTemplateVariable := &sdkproto.TemplateVariable{
Name: "first",
Type: "string",
@@ -863,7 +891,7 @@ func TestUpdateJob(t *testing.T) {
require.NoError(t, err)
require.Len(t, response.VariableValues, 2)
templateVariables, err := db.GetTemplateVersionVariables(ctx, versionID)
templateVariables, err := db.GetTemplateVersionVariables(ctx, templateVersionID)
require.NoError(t, err)
require.Len(t, templateVariables, 2)
require.Equal(t, templateVariables[0].Value, firstTemplateVariable.DefaultValue)
@@ -875,16 +903,7 @@ func TestUpdateJob(t *testing.T) {
defer cancel()
srv, db, _, pd := setup(t, false, &overrides{})
user := dbgen.User(t, db, database.User{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
versionID := uuid.New()
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
CreatedBy: user.ID,
ID: versionID,
JobID: job,
OrganizationID: pd.OrganizationID,
})
require.NoError(t, err)
templateVersionID, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
firstTemplateVariable := &sdkproto.TemplateVariable{
Name: "first",
Type: "string",
@@ -909,7 +928,7 @@ func TestUpdateJob(t *testing.T) {
// Even though there is an error returned, variables are stored in the database
// to show the schema in the site UI.
templateVariables, err := db.GetTemplateVersionVariables(ctx, versionID)
templateVariables, err := db.GetTemplateVersionVariables(ctx, templateVersionID)
require.NoError(t, err)
require.Len(t, templateVariables, 2)
require.Equal(t, templateVariables[0].Value, firstTemplateVariable.DefaultValue)
@@ -923,18 +942,9 @@ func TestUpdateJob(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
versionID := uuid.New()
user := dbgen.User(t, db, database.User{})
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
ID: versionID,
CreatedBy: user.ID,
JobID: job,
OrganizationID: pd.OrganizationID,
})
require.NoError(t, err)
_, err = srv.UpdateJob(ctx, &proto.UpdateJobRequest{
srv, db, _, pd := setup(t, false, nil)
templateVersionID, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, err := srv.UpdateJob(ctx, &proto.UpdateJobRequest{
JobId: job.String(),
WorkspaceTags: map[string]string{
"bird": "tweety",
@@ -943,7 +953,7 @@ func TestUpdateJob(t *testing.T) {
})
require.NoError(t, err)
workspaceTags, err := db.GetTemplateVersionWorkspaceTags(ctx, versionID)
workspaceTags, err := db.GetTemplateVersionWorkspaceTags(ctx, templateVersionID)
require.NoError(t, err)
require.Len(t, workspaceTags, 2)
require.Equal(t, workspaceTags[0].Key, "bird")
@@ -955,7 +965,7 @@ func TestUpdateJob(t *testing.T) {
t.Run("LogSizeLimit", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
// Create a log message that exceeds the 1MB limit
largeOutput := strings.Repeat("a", 1048577) // 1MB + 1 byte
@@ -979,7 +989,7 @@ func TestUpdateJob(t *testing.T) {
t.Run("IncrementalLogSizeOverflow", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
// Send logs that together exceed the limit
mediumOutput := strings.Repeat("b", 524289) // Half a MB + 1 byte
@@ -1020,7 +1030,7 @@ func TestUpdateJob(t *testing.T) {
t.Run("LogSizeTracking", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
logOutput := "test log message"
expectedSize := int32(len(logOutput)) // #nosec G115 - Log length is 16.
@@ -1045,7 +1055,7 @@ func TestUpdateJob(t *testing.T) {
t.Run("LogOverflowStopsProcessing", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
_, job := setupJob(t, db, pd.ID, pd.OrganizationID, pd.Tags)
// First: trigger overflow
largeOutput := strings.Repeat("a", 1048577) // 1MB + 1 byte
@@ -1108,12 +1118,20 @@ func TestFailJob(t *testing.T) {
t.Run("NotOwner", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, nil)
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionImport,
Input: json.RawMessage("{}"),
ID: version.JobID,
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionImport,
Input: must(json.Marshal(provisionerdserver.TemplateVersionImportJob{
TemplateVersionID: version.ID,
})),
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
@@ -1139,13 +1157,21 @@ func TestFailJob(t *testing.T) {
})
t.Run("AlreadyCompleted", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
srv, db, _, pd := setup(t, false, nil)
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionImport,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: json.RawMessage("{}"),
ID: version.JobID,
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionImport,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: must(json.Marshal(provisionerdserver.TemplateVersionImportJob{
TemplateVersionID: version.ID,
})),
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
@@ -1310,14 +1336,22 @@ func TestCompleteJob(t *testing.T) {
t.Run("NotOwner", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, nil)
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
ID: version.JobID,
Provisioner: database.ProvisionerTypeEcho,
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeWorkspaceBuild,
Type: database.ProvisionerJobTypeTemplateVersionImport,
OrganizationID: pd.OrganizationID,
Input: json.RawMessage("{}"),
Tags: pd.Tags,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: version.ID,
})),
Tags: pd.Tags,
})
require.NoError(t, err)
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
@@ -1361,10 +1395,12 @@ func TestCompleteJob(t *testing.T) {
OrganizationID: pd.OrganizationID,
ID: jobID,
Provisioner: database.ProvisionerTypeEcho,
Input: []byte(`{"template_version_id": "` + versionID.String() + `"}`),
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionImport,
Tags: pd.Tags,
Input: must(json.Marshal(provisionerdserver.TemplateVersionImportJob{
TemplateVersionID: versionID,
})),
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionImport,
Tags: pd.Tags,
})
require.NoError(t, err)
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
@@ -1410,14 +1446,22 @@ func TestCompleteJob(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
org := dbgen.Organization(t, db, database.Organization{})
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: org.ID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
OrganizationID: org.ID,
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: json.RawMessage("{}"),
Tags: pd.Tags,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: version.ID,
})),
Tags: pd.Tags,
})
require.NoError(t, err)
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
@@ -1628,25 +1672,49 @@ func TestCompleteJob(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
jobID := uuid.New()
versionID := uuid.New()
user := dbgen.User(t, db, database.User{})
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
tv := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
ID: versionID,
JobID: jobID,
OrganizationID: pd.OrganizationID,
JobID: jobID,
})
template := dbgen.Template(t, db, database.Template{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
ActiveVersionID: tv.ID,
})
err := db.UpdateTemplateVersionByID(ctx, database.UpdateTemplateVersionByIDParams{
ID: tv.ID,
TemplateID: uuid.NullUUID{
UUID: template.ID,
Valid: true,
},
UpdatedAt: dbtime.Now(),
Name: tv.Name,
Message: tv.Message,
})
require.NoError(t, err)
workspace := dbgen.Workspace(t, db, database.WorkspaceTable{
OwnerID: user.ID,
OrganizationID: pd.OrganizationID,
TemplateID: template.ID,
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: jobID,
Provisioner: database.ProvisionerTypeEcho,
Input: []byte(`{"template_version_id": "` + versionID.String() + `"}`),
Input: json.RawMessage("{}"),
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeWorkspaceBuild,
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
require.NoError(t, err)
_ = dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
WorkspaceID: workspace.ID,
TemplateVersionID: tv.ID,
InitiatorID: user.ID,
JobID: jobID,
})
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
OrganizationID: pd.OrganizationID,
WorkerID: uuid.NullUUID{
@@ -1697,11 +1765,13 @@ func TestCompleteJob(t *testing.T) {
})
require.NoError(t, err)
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: jobID,
Provisioner: database.ProvisionerTypeEcho,
Input: []byte(`{"template_version_id": "` + versionID.String() + `"}`),
ID: jobID,
Provisioner: database.ProvisionerTypeEcho,
Input: must(json.Marshal(provisionerdserver.TemplateVersionImportJob{
TemplateVersionID: versionID,
})),
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeWorkspaceBuild,
Type: database.ProvisionerJobTypeTemplateVersionImport,
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
@@ -1766,10 +1836,12 @@ func TestCompleteJob(t *testing.T) {
OrganizationID: pd.OrganizationID,
ID: jobID,
Provisioner: database.ProvisionerTypeEcho,
Input: []byte(`{"template_version_id": "` + versionID.String() + `"}`),
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeWorkspaceBuild,
Tags: pd.Tags,
Input: must(json.Marshal(provisionerdserver.TemplateVersionImportJob{
TemplateVersionID: versionID,
})),
StorageMethod: database.ProvisionerStorageMethodFile,
Type: database.ProvisionerJobTypeTemplateVersionImport,
Tags: pd.Tags,
})
require.NoError(t, err)
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
@@ -2091,12 +2163,20 @@ func TestCompleteJob(t *testing.T) {
t.Run("TemplateDryRun", func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
user := dbgen.User(t, db, database.User{})
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
JobID: uuid.New(),
})
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
ID: uuid.New(),
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: json.RawMessage("{}"),
ID: version.JobID,
Provisioner: database.ProvisionerTypeEcho,
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
StorageMethod: database.ProvisionerStorageMethodFile,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: version.ID,
})),
OrganizationID: pd.OrganizationID,
Tags: pd.Tags,
})
@@ -2191,8 +2271,10 @@ func TestCompleteJob(t *testing.T) {
Transition: database.WorkspaceTransitionStart,
}},
provisionerJobParams: database.InsertProvisionerJobParams{
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
Input: json.RawMessage("{}"),
Type: database.ProvisionerJobTypeTemplateVersionDryRun,
Input: must(json.Marshal(provisionerdserver.TemplateVersionDryRunJob{
TemplateVersionID: templateVersionID,
})),
},
},
{
@@ -2349,22 +2431,26 @@ func TestCompleteJob(t *testing.T) {
OrganizationID: pd.OrganizationID,
})
tv := dbgen.TemplateVersion(t, db, database.TemplateVersion{
ID: templateVersionID,
CreatedBy: user.ID,
OrganizationID: pd.OrganizationID,
TemplateID: uuid.NullUUID{UUID: tpl.ID, Valid: true},
JobID: job.ID,
})
workspace := dbgen.Workspace(t, db, database.WorkspaceTable{
TemplateID: tpl.ID,
OrganizationID: pd.OrganizationID,
OwnerID: user.ID,
})
_ = dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
ID: workspaceBuildID,
JobID: job.ID,
WorkspaceID: workspace.ID,
TemplateVersionID: tv.ID,
})
if jobParams.Type == database.ProvisionerJobTypeWorkspaceBuild {
workspace := dbgen.Workspace(t, db, database.WorkspaceTable{
TemplateID: tpl.ID,
OrganizationID: pd.OrganizationID,
OwnerID: user.ID,
})
_ = dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
ID: workspaceBuildID,
JobID: job.ID,
WorkspaceID: workspace.ID,
TemplateVersionID: tv.ID,
})
}
require.NoError(t, err)
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
@@ -2672,7 +2758,10 @@ func TestCompleteJob(t *testing.T) {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
fakeUsageInserter, usageInserterPtr := newFakeUsageInserter()
srv, db, _, pd := setup(t, false, &overrides{
usageInserter: usageInserterPtr,
})
importJobID := uuid.New()
tvID := uuid.New()
@@ -2741,6 +2830,10 @@ func TestCompleteJob(t *testing.T) {
require.NoError(t, err)
require.True(t, version.HasAITask.Valid) // We ALWAYS expect a value to be set, therefore not nil, i.e. valid = true.
require.Equal(t, tc.expected, version.HasAITask.Bool)
// We never expect a usage event to be collected for
// template imports.
require.Empty(t, fakeUsageInserter.collectedEvents)
})
}
})
@@ -2750,22 +2843,27 @@ func TestCompleteJob(t *testing.T) {
// will be set as well in that case.
t.Run("WorkspaceBuild", func(t *testing.T) {
type testcase struct {
name string
input *proto.CompletedJob_WorkspaceBuild
expected bool
name string
transition database.WorkspaceTransition
input *proto.CompletedJob_WorkspaceBuild
expectHasAiTask bool
expectUsageEvent bool
}
sidebarAppID := uuid.NewString()
for _, tc := range []testcase{
{
name: "has_ai_task is false by default",
input: &proto.CompletedJob_WorkspaceBuild{
name: "has_ai_task is false by default",
transition: database.WorkspaceTransitionStart,
input: &proto.CompletedJob_WorkspaceBuild{
// No AiTasks defined.
},
expected: false,
expectHasAiTask: false,
expectUsageEvent: false,
},
{
name: "has_ai_task is set to true",
name: "has_ai_task is set to true",
transition: database.WorkspaceTransitionStart,
input: &proto.CompletedJob_WorkspaceBuild{
AiTasks: []*sdkproto.AITask{
{
@@ -2792,11 +2890,13 @@ func TestCompleteJob(t *testing.T) {
},
},
},
expected: true,
expectHasAiTask: true,
expectUsageEvent: true,
},
// Checks regression for https://github.com/coder/coder/issues/18776
{
name: "non-existing app",
name: "non-existing app",
transition: database.WorkspaceTransitionStart,
input: &proto.CompletedJob_WorkspaceBuild{
AiTasks: []*sdkproto.AITask{
{
@@ -2808,13 +2908,49 @@ func TestCompleteJob(t *testing.T) {
},
},
},
expected: false,
expectHasAiTask: false,
expectUsageEvent: false,
},
{
name: "has_ai_task is set to true, but transition is not start",
transition: database.WorkspaceTransitionStop,
input: &proto.CompletedJob_WorkspaceBuild{
AiTasks: []*sdkproto.AITask{
{
Id: uuid.NewString(),
SidebarApp: &sdkproto.AITaskSidebarApp{
Id: sidebarAppID,
},
},
},
Resources: []*sdkproto.Resource{
{
Agents: []*sdkproto.Agent{
{
Id: uuid.NewString(),
Name: "a",
Apps: []*sdkproto.App{
{
Id: sidebarAppID,
Slug: "test-app",
},
},
},
},
},
},
},
expectHasAiTask: true,
expectUsageEvent: false,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
srv, db, _, pd := setup(t, false, &overrides{})
fakeUsageInserter, usageInserterPtr := newFakeUsageInserter()
srv, db, _, pd := setup(t, false, &overrides{
usageInserter: usageInserterPtr,
})
importJobID := uuid.New()
tvID := uuid.New()
@@ -2868,7 +3004,7 @@ func TestCompleteJob(t *testing.T) {
WorkspaceID: workspaceTable.ID,
TemplateVersionID: version.ID,
InitiatorID: user.ID,
Transition: database.WorkspaceTransitionStart,
Transition: tc.transition,
})
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
@@ -2899,11 +3035,22 @@ func TestCompleteJob(t *testing.T) {
build, err = db.GetWorkspaceBuildByID(ctx, build.ID)
require.NoError(t, err)
require.True(t, build.HasAITask.Valid) // We ALWAYS expect a value to be set, therefore not nil, i.e. valid = true.
require.Equal(t, tc.expected, build.HasAITask.Bool)
require.Equal(t, tc.expectHasAiTask, build.HasAITask.Bool)
if tc.expected {
if tc.expectHasAiTask {
require.Equal(t, sidebarAppID, build.AITaskSidebarAppID.UUID.String())
}
if tc.expectUsageEvent {
// Check that a usage event was collected.
require.Len(t, fakeUsageInserter.collectedEvents, 1)
require.Equal(t, usage.DCManagedAgentsV1{
Count: 1,
}, fakeUsageInserter.collectedEvents[0])
} else {
// Check that no usage event was collected.
require.Empty(t, fakeUsageInserter.collectedEvents)
}
})
}
})
@@ -3835,6 +3982,7 @@ type overrides struct {
externalAuthConfigs []*externalauth.Config
templateScheduleStore *atomic.Pointer[schedule.TemplateScheduleStore]
userQuietHoursScheduleStore *atomic.Pointer[schedule.UserQuietHoursScheduleStore]
usageInserter *atomic.Pointer[usage.Inserter]
clock *quartz.Mock
acquireJobLongPollDuration time.Duration
heartbeatFn func(ctx context.Context) error
@@ -3855,13 +4003,14 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi
var externalAuthConfigs []*externalauth.Config
tss := testTemplateScheduleStore()
uqhss := testUserQuietHoursScheduleStore()
usageInserter := testUsageInserter()
clock := quartz.NewReal()
pollDur := time.Duration(0)
if ov == nil {
ov = &overrides{}
}
if ov.ctx == nil {
ctx, cancel := context.WithCancel(context.Background())
ctx, cancel := context.WithCancel(dbauthz.AsProvisionerd(context.Background()))
t.Cleanup(cancel)
ov.ctx = ctx
}
@@ -3892,6 +4041,15 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi
require.True(t, swapped)
}
}
if ov.usageInserter != nil {
tUsageInserter := usageInserter.Load()
// keep the initial test value if the override hasn't set the atomic pointer.
usageInserter = ov.usageInserter
if usageInserter.Load() == nil {
swapped := usageInserter.CompareAndSwap(nil, tUsageInserter)
require.True(t, swapped)
}
}
if ov.clock != nil {
clock = ov.clock
}
@@ -3929,6 +4087,10 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi
var op atomic.Pointer[agplprebuilds.ReconciliationOrchestrator]
op.Store(&prebuildsOrchestrator)
// Use an authz wrapped database for the server to ensure permission checks
// work.
authorizer := rbac.NewStrictCachingAuthorizer(prometheus.NewRegistry())
serverDB := dbauthz.New(db, authorizer, logger, coderdtest.AccessControlStorePointer())
srv, err := provisionerdserver.NewServer(
ov.ctx,
proto.CurrentVersion.String(),
@@ -3938,7 +4100,7 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi
slogtest.Make(t, &slogtest.Options{IgnoreErrors: ignoreLogErrors}),
[]database.ProvisionerType{database.ProvisionerTypeEcho},
provisionerdserver.Tags(daemon.Tags),
db,
serverDB,
ps,
provisionerdserver.NewAcquirer(ov.ctx, logger.Named("acquirer"), db, ps),
telemetry.NewNoop(),
@@ -3947,6 +4109,7 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi
auditPtr,
tss,
uqhss,
usageInserter,
deploymentValues,
provisionerdserver.Options{
ExternalAuthConfigs: externalAuthConfigs,
@@ -4061,3 +4224,22 @@ func (s *fakeStream) cancel() {
s.canceled = true
s.c.Broadcast()
}
type fakeUsageInserter struct {
collectedEvents []usage.Event
}
var _ usage.Inserter = &fakeUsageInserter{}
func newFakeUsageInserter() (*fakeUsageInserter, *atomic.Pointer[usage.Inserter]) {
ptr := &atomic.Pointer[usage.Inserter]{}
fake := &fakeUsageInserter{}
var inserter usage.Inserter = fake
ptr.Store(&inserter)
return fake, ptr
}
func (f *fakeUsageInserter) InsertDiscreteUsageEvent(_ context.Context, _ database.Store, event usage.DiscreteEvent) error {
f.collectedEvents = append(f.collectedEvents, event)
return nil
}