diff --git a/coderd/coderd.go b/coderd/coderd.go index aabd12188c..bcc95f182e 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -319,6 +319,9 @@ type Options struct { // rotator is the sole creator of nats_ca rows, so this cache is read-only. NATSCACache cryptokeys.SigningKeycache Clock quartz.Clock + // Acquirer acquires provisioner jobs. Defaults to provisionerdserver.Acquirer + // backed by Database and Pubsub. + Acquirer *provisionerdserver.Acquirer // WebPushDispatcher is a way to send notifications over Web Push. WebPushDispatcher webpush.Dispatcher @@ -671,6 +674,15 @@ func New(options *Options) *API { var buildUsageChecker atomic.Pointer[wsbuilder.UsageChecker] var noopUsageChecker wsbuilder.UsageChecker = wsbuilder.NoopUsageChecker{} buildUsageChecker.Store(&noopUsageChecker) + acquirer := options.Acquirer + if acquirer == nil { + acquirer = provisionerdserver.NewAcquirer( + ctx, + options.Logger.Named("acquirer"), + options.Database, + options.Pubsub, + ) + } api := &API{ ctx: ctx, cancel: cancel, @@ -696,15 +708,10 @@ func New(options *Options) *API { Experiments: experiments, WebpushDispatcher: options.WebPushDispatcher, healthCheckGroup: &singleflight.Group[string, *healthsdk.HealthcheckReport]{}, - Acquirer: provisionerdserver.NewAcquirer( - ctx, - options.Logger.Named("acquirer"), - options.Database, - options.Pubsub, - ), - dbRolluper: options.DatabaseRolluper, - ProfileCollector: defaultProfileCollector{}, - AISeatTracker: aiseats.Noop{}, + Acquirer: acquirer, + dbRolluper: options.DatabaseRolluper, + ProfileCollector: defaultProfileCollector{}, + AISeatTracker: aiseats.Noop{}, } api.WorkspaceAppsProvider = workspaceapps.NewDBTokenProvider( diff --git a/coderd/coderdtest/coderdtest.go b/coderd/coderdtest/coderdtest.go index 8ff4334828..c790db406c 100644 --- a/coderd/coderdtest/coderdtest.go +++ b/coderd/coderdtest/coderdtest.go @@ -198,6 +198,7 @@ type Options struct { APIKeyEncryptionCache cryptokeys.EncryptionKeycache OIDCConvertKeyCache cryptokeys.SigningKeycache Clock quartz.Clock + Acquirer *provisionerdserver.Acquirer TelemetryReporter telemetry.Reporter ProvisionerdServerMetrics *provisionerdserver.Metrics @@ -655,6 +656,7 @@ func NewOptions(t testing.TB, options *Options) (func(http.Handler), context.Can NotificationsEnqueuer: options.NotificationsEnqueuer, OneTimePasscodeValidityPeriod: options.OneTimePasscodeValidityPeriod, Clock: options.Clock, + Acquirer: options.Acquirer, AppEncryptionKeyCache: options.APIKeyEncryptionCache, OIDCConvertKeyCache: options.OIDCConvertKeyCache, ProvisionerdServerMetrics: options.ProvisionerdServerMetrics, diff --git a/coderd/provisionerdserver/acquirer.go b/coderd/provisionerdserver/acquirer.go index adb508de10..e082a9651e 100644 --- a/coderd/provisionerdserver/acquirer.go +++ b/coderd/provisionerdserver/acquirer.go @@ -18,6 +18,7 @@ import ( "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/coder/v2/coderd/database/provisionerjobs" "github.com/coder/coder/v2/coderd/database/pubsub" + "github.com/coder/quartz" ) const ( @@ -49,15 +50,14 @@ type Acquirer struct { mu sync.Mutex q map[dKey]domain - // testing only - backupPollDuration time.Duration + clock quartz.Clock } type AcquirerOption func(*Acquirer) -func TestingBackupPollDuration(dur time.Duration) AcquirerOption { +func WithClock(clock quartz.Clock) AcquirerOption { return func(a *Acquirer) { - a.backupPollDuration = dur + a.clock = clock } } @@ -70,12 +70,12 @@ func NewAcquirer(ctx context.Context, logger slog.Logger, store AcquirerStore, p opts ...AcquirerOption, ) *Acquirer { a := &Acquirer{ - ctx: ctx, - logger: logger, - store: store, - ps: ps, - q: make(map[dKey]domain), - backupPollDuration: backupPollDuration, + ctx: ctx, + logger: logger, + store: store, + ps: ps, + q: make(map[dKey]domain), + clock: quartz.NewReal(), } for _, opt := range opts { opt(a) @@ -173,7 +173,7 @@ func (a *Acquirer) want(organization uuid.UUID, pt []database.ProvisionerType, t acquirees: make(map[chan<- struct{}]*acquiree), } a.q[dk] = d - go d.poll(a.backupPollDuration) + go d.poll(backupPollDuration) // this is a new request for this dKey, so is cleared. cleared = true } @@ -483,7 +483,7 @@ func (d domain) contains(p provisionerjobs.JobPosting) bool { } func (d domain) poll(dur time.Duration) { - tkr := time.NewTicker(dur) + tkr := d.a.clock.NewTicker(dur, "acquirer", "backup_poll") defer tkr.Stop() for { select { diff --git a/coderd/provisionerdserver/acquirer_test.go b/coderd/provisionerdserver/acquirer_test.go index 0f724ad173..3198fad25c 100644 --- a/coderd/provisionerdserver/acquirer_test.go +++ b/coderd/provisionerdserver/acquirer_test.go @@ -9,7 +9,6 @@ import ( "strings" "sync" "testing" - "time" "github.com/google/uuid" "github.com/sqlc-dev/pqtype" @@ -25,6 +24,7 @@ import ( "github.com/coder/coder/v2/coderd/provisionerdserver" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/testutil" + "github.com/coder/quartz" ) func TestMain(m *testing.M) { @@ -38,7 +38,9 @@ func TestAcquirer_Store(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) - _ = provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), db, ps) + _ = provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), db, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) } func TestAcquirer_Single(t *testing.T) { @@ -48,7 +50,9 @@ func TestAcquirer_Single(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) - uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) orgID := uuid.New() workerID := uuid.New() @@ -75,7 +79,9 @@ func TestAcquirer_MultipleSameDomain(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) - uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) acquirees := make([]*testAcquiree, 0, 10) jobIDs := make(map[uuid.UUID]bool) @@ -121,7 +127,9 @@ func TestAcquirer_WaitsOnNoJobs(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) - uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) orgID := uuid.New() workerID := uuid.New() @@ -173,7 +181,9 @@ func TestAcquirer_RetriesPending(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) - uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) orgID := uuid.New() workerID := uuid.New() @@ -193,11 +203,8 @@ func TestAcquirer_RetriesPending(t *testing.T) { // First call to DB is in progress. Send in posting postJob(t, ps, database.ProvisionerTypeEcho, provisionerdserver.Tags{}) - // there is a race between the posting being processed and the DB call - // returning. In either case we should retry, but we're trying to hit the - // case where the posting is processed first, so sleep a little bit to give - // it a chance. - time.Sleep(testutil.IntervalMedium) + // MemoryPubsub.Publish waits for the listener to finish, so the pending + // notification has been processed before the first database call returns. // Now, when first DB call returns ErrNoRows we retry. err := fs.sendCtx(ctx, database.ProvisionerJob{}, sql.ErrNoRows) @@ -217,6 +224,9 @@ func TestAcquirer_DifferentDomains(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) orgID := uuid.New() pt := []database.ProvisionerType{database.ProvisionerTypeEcho} @@ -235,8 +245,6 @@ func TestAcquirer_DifferentDomains(t *testing.T) { {ID: jobID, Provisioner: database.ProvisionerTypeEcho, Tags: database.StringMap{"worker": "1"}}, } - uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) - ctx0, cancel0 := context.WithCancel(ctx) defer cancel0() acquiree0.startAcquire(ctx0, uut) @@ -264,9 +272,11 @@ func TestAcquirer_BackupPoll(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) defer cancel() logger := testutil.Logger(t) + clock := quartz.NewMock(t) + tickerTrap := clock.Trap().NewTicker("acquirer", "backup_poll") uut := provisionerdserver.NewAcquirer( ctx, logger.Named("acquirer"), fs, ps, - provisionerdserver.TestingBackupPollDuration(testutil.IntervalMedium), + provisionerdserver.WithClock(clock), ) workerID := uuid.New() @@ -282,6 +292,15 @@ func TestAcquirer_BackupPoll(t *testing.T) { err = fs.sendCtx(ctx, database.ProvisionerJob{ID: jobID}, nil) require.NoError(t, err) acquiree.startAcquire(ctx, uut) + select { + case <-fs.callStarted: + case <-ctx.Done(): + t.Fatal("timed out waiting for initial database call") + } + tickerCall := tickerTrap.MustWait(ctx) + tickerCall.MustRelease(ctx) + _, waiter := clock.AdvanceNext() + waiter.MustWait(ctx) job := acquiree.success(ctx) require.Equal(t, jobID, job.ID) } @@ -307,7 +326,9 @@ func TestAcquirer_UnblockOnCancel(t *testing.T) { acquiree1 := newTestAcquiree(t, orgID, worker1, pt, tags) jobID := uuid.New() - uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) // queue up 2 responses --- we may not need both, since acquiree0 will // usually cancel before calling, but cancel is async, so it might call. @@ -498,20 +519,43 @@ func TestAcquirer_MatchTags(t *testing.T) { }) require.NoError(t, err) ptypes := []database.ProvisionerType{database.ProvisionerTypeEcho} - acq := provisionerdserver.NewAcquirer(ctx, log, db, ps) - acquireOrgID := org.ID if tt.unmatchedOrg { acquireOrgID = uuid.New() } - aj, err := acq.AcquireJob(ctx, acquireOrgID, uuid.New(), ptypes, tt.acquireJobTags) + if tt.expectAcquire { + acq := provisionerdserver.NewAcquirer(ctx, log, db, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) + aj, err := acq.AcquireJob(ctx, acquireOrgID, uuid.New(), ptypes, tt.acquireJobTags) assert.NoError(t, err) assert.Equal(t, pj.ID, aj.ID) - } else { - assert.Empty(t, aj, "should not have acquired job") - assert.ErrorIs(t, err, context.DeadlineExceeded, "should have timed out") + return } + + store := &acquirerStoreSpy{ + Store: db, + callCompleted: make(chan struct{}, 1), + } + acq := provisionerdserver.NewAcquirer(ctx, log, store, ps, + provisionerdserver.WithClock(quartz.NewMock(t)), + ) + acquireCtx, acquireCancel := context.WithCancel(ctx) + acquiree := newTestAcquiree(t, acquireOrgID, uuid.New(), ptypes, tt.acquireJobTags) + acquiree.startAcquire(acquireCtx, acq) + select { + case <-store.callCompleted: + case <-ctx.Done(): + t.Fatal("timed out waiting for initial database call") + } + acquireCancel() + acquiree.requireCanceled(ctx) + + job, err := db.GetProvisionerJobByID(ctx, pj.ID) + require.NoError(t, err) + require.False(t, job.StartedAt.Valid) + require.False(t, job.WorkerID.Valid) }) } @@ -562,11 +606,28 @@ func postJob(t *testing.T, ps pubsub.Pubsub, pt database.ProvisionerType, tags p require.NoError(t, err) } +type acquirerStoreSpy struct { + database.Store + callCompleted chan struct{} +} + +func (s *acquirerStoreSpy) AcquireProvisionerJob( + ctx context.Context, params database.AcquireProvisionerJobParams, +) (database.ProvisionerJob, error) { + job, err := s.Store.AcquireProvisionerJob(ctx, params) + select { + case s.callCompleted <- struct{}{}: + default: + } + return job, err +} + // fakeOrderedStore is a fake store that lets tests send AcquireProvisionerJob // results in order over a channel, and tests for overlapped calls. type fakeOrderedStore struct { - jobs chan database.ProvisionerJob - errors chan error + jobs chan database.ProvisionerJob + errors chan error + callStarted chan struct{} mu sync.Mutex params []database.AcquireProvisionerJobParams @@ -581,9 +642,10 @@ func newFakeOrderedStore() *fakeOrderedStore { return &fakeOrderedStore{ // buffer the channels so that we can queue up lots of responses to // occur nearly simultaneously - jobs: make(chan database.ProvisionerJob, 100), - errors: make(chan error, 100), - inflight: make(map[uuid.UUID]bool), + jobs: make(chan database.ProvisionerJob, 100), + errors: make(chan error, 100), + callStarted: make(chan struct{}, 100), + inflight: make(map[uuid.UUID]bool), } } @@ -599,6 +661,10 @@ func (s *fakeOrderedStore) AcquireProvisionerJob( } s.inflight[params.WorkerID.UUID] = true s.mu.Unlock() + select { + case s.callStarted <- struct{}{}: + default: + } job := <-s.jobs err := <-s.errors diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index 80ad75a493..4713dbe399 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -5294,7 +5294,13 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi provisionerdserver.Tags(daemon.Tags), serverDB, ps, - provisionerdserver.NewAcquirer(ov.ctx, logger.Named("acquirer"), db, ps), + provisionerdserver.NewAcquirer( + ov.ctx, + logger.Named("acquirer"), + db, + ps, + provisionerdserver.WithClock(clock), + ), telemetry.NewNoop(), trace.NewNoopTracerProvider().Tracer("noop"), &atomic.Pointer[proto.QuotaCommitter]{}, diff --git a/enterprise/coderd/prebuilds/claim_test.go b/enterprise/coderd/prebuilds/claim_test.go index e58913ed40..bbfd381779 100644 --- a/enterprise/coderd/prebuilds/claim_test.go +++ b/enterprise/coderd/prebuilds/claim_test.go @@ -23,6 +23,7 @@ import ( "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/coder/v2/coderd/files" agplprebuilds "github.com/coder/coder/v2/coderd/prebuilds" + "github.com/coder/coder/v2/coderd/provisionerdserver" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" @@ -131,19 +132,28 @@ func TestClaimPrebuild(t *testing.T) { // Setup clock := quartz.NewMock(t) + acquirerClock := quartz.NewMock(t) clock.Set(dbtime.Now()) ctx := testutil.Context(t, testutil.WaitSuperLong) db, pubsub := dbtestutil.NewDB(t) + logger := testutil.Logger(t) + acquirer := provisionerdserver.NewAcquirer( + ctx, + logger.Named("acquirer"), + db, + pubsub, + provisionerdserver.WithClock(acquirerClock), + ) spy := newStoreSpy(db, tc.claimingErr) expectedPrebuildsCount := desiredInstances * presetCount - logger := testutil.Logger(t) client, _, api, owner := coderdenttest.NewWithAPI(t, &coderdenttest.Options{ Options: &coderdtest.Options{ Database: spy, Pubsub: pubsub, Clock: clock, + Acquirer: acquirer, }, LicenseOptions: &coderdenttest.LicenseOptions{ Features: license.Features{ @@ -160,10 +170,18 @@ func TestClaimPrebuild(t *testing.T) { orgID = secondOrg.ID } + acquirerTickerTrap := acquirerClock.Trap().NewTicker("acquirer", "backup_poll") + defer acquirerTickerTrap.Close() provisionerCloser := coderdenttest.NewExternalProvisionerDaemon(t, client, orgID, map[string]string{ provisionersdk.TagScope: provisionersdk.ScopeOrganization, }) defer provisionerCloser.Close() + acquirerTickerTrap.MustWait(ctx).MustRelease(ctx) + secondAcquirerTickerReady := make(chan struct{}) + go func() { + acquirerTickerTrap.MustWait(ctx).MustRelease(ctx) + close(secondAcquirerTickerReady) + }() cache := files.New(prometheus.NewRegistry(), &coderdtest.FakeAuthorizer{}) reconciler := prebuilds.NewStoreReconciler( @@ -201,12 +219,12 @@ func TestClaimPrebuild(t *testing.T) { actions, err := reconciler.CalculateActions(ctx, *ps) require.NoError(t, err) require.NotNil(t, actions) - require.NoError(t, reconciler.ReconcilePreset(ctx, *ps)) } // Given: a set of running, eligible prebuilds eventually starts up. runningPrebuilds := make(map[uuid.UUID]database.GetRunningPrebuiltWorkspacesRow, desiredInstances*presetCount) + advancedAcquirerClock := false require.Eventually(t, func() bool { rows, err := spy.GetRunningPrebuiltWorkspaces(ctx) if err != nil { @@ -239,6 +257,15 @@ func TestClaimPrebuild(t *testing.T) { } } + if !advancedAcquirerClock { + select { + case <-secondAcquirerTickerReady: + acquirerClock.Advance(30 * time.Second).MustWait(ctx) + advancedAcquirerClock = true + default: + } + } + t.Logf("found %d running prebuilds so far, want %d", len(runningPrebuilds), expectedPrebuildsCount) return len(runningPrebuilds) == expectedPrebuildsCount diff --git a/enterprise/coderd/workspaces_test.go b/enterprise/coderd/workspaces_test.go index b4913a7d2c..cfc99a84db 100644 --- a/enterprise/coderd/workspaces_test.go +++ b/enterprise/coderd/workspaces_test.go @@ -2559,12 +2559,21 @@ func TestPrebuildsAutobuild(t *testing.T) { // Set the clock to Monday, January 1st, 2024 at 8:00 AM UTC to keep the test deterministic clock := quartz.NewMock(t) + acquirerClock := quartz.NewMock(t) clock.Set(time.Date(2024, 1, 1, 8, 0, 0, 0, time.UTC)) + acquirerTickerTrap := acquirerClock.Trap().NewTicker("acquirer", "backup_poll") // Setup ctx := testutil.Context(t, testutil.WaitSuperLong) db, pb := dbtestutil.NewDB(t, dbtestutil.WithDumpOnFailure()) logger := testutil.Logger(t) + acquirer := provisionerdserver.NewAcquirer( + ctx, + logger.Named("acquirer"), + db, + pb, + provisionerdserver.WithClock(acquirerClock), + ) tickCh := make(chan time.Time) statsCh := make(chan autobuild.Stats) notificationsNoop := notifications.NewNoopEnqueuer() @@ -2576,6 +2585,7 @@ func TestPrebuildsAutobuild(t *testing.T) { IncludeProvisionerDaemon: true, AutobuildStats: statsCh, Clock: clock, + Acquirer: acquirer, TemplateScheduleStore: schedule.NewEnterpriseTemplateScheduleStore( agplUserQuietHoursScheduleStore(), notificationsNoop, @@ -2589,6 +2599,10 @@ func TestPrebuildsAutobuild(t *testing.T) { }, }, }) + // The Acquirer creates a fresh backup-poll ticker for the initial idle + // wait and again after completing the template import job. Release both + // so the second ticker exists before the clock advances below. + acquirerTickerTrap.MustWait(ctx).MustRelease(ctx) // Setup Prebuild reconciler cache := files.New(prometheus.NewRegistry(), &coderdtest.FakeAuthorizer{}) @@ -2620,8 +2634,12 @@ func TestPrebuildsAutobuild(t *testing.T) { require.NoError(t, err) require.Len(t, presets, 1) + acquirerTickerTrap.MustWait(ctx).MustRelease(ctx) + acquirerTickerTrap.Close() + // Given: reconciliation loop runs and starts prebuilt workspace in failed state runReconciliationLoop(t, ctx, db, reconciler, presets) + acquirerClock.Advance(30 * time.Second).MustWait(ctx) var failedWorkspaceBuilds []database.GetFailedWorkspaceBuildsByTemplateIDRow require.Eventually(t, func() bool { rows, err := db.GetFailedWorkspaceBuildsByTemplateID(ctx, database.GetFailedWorkspaceBuildsByTemplateIDParams{