mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(coderd): insert provisioner daemons (#11207)
* Adds UpdateProvisionerDaemonLastSeenAt * Adds heartbeat to provisioner daemons * Inserts provisioner daemons to database upon start * Ensures TagOwner is an empty string and not nil * Adds COALESCE() in idx_provisioner_daemons_name_owner_key
This commit is contained in:
@@ -44,9 +44,15 @@ import (
|
||||
sdkproto "github.com/coder/coder/v2/provisionersdk/proto"
|
||||
)
|
||||
|
||||
// DefaultAcquireJobLongPollDur is the time the (deprecated) AcquireJob rpc waits to try to obtain a job before
|
||||
// canceling and returning an empty job.
|
||||
const DefaultAcquireJobLongPollDur = time.Second * 5
|
||||
const (
|
||||
// DefaultAcquireJobLongPollDur is the time the (deprecated) AcquireJob rpc waits to try to obtain a job before
|
||||
// canceling and returning an empty job.
|
||||
DefaultAcquireJobLongPollDur = time.Second * 5
|
||||
|
||||
// DefaultHeartbeatInterval is the interval at which the provisioner daemon
|
||||
// will update its last seen at timestamp in the database.
|
||||
DefaultHeartbeatInterval = time.Minute
|
||||
)
|
||||
|
||||
type Options struct {
|
||||
OIDCConfig httpmw.OAuth2Config
|
||||
@@ -56,6 +62,16 @@ type Options struct {
|
||||
|
||||
// AcquireJobLongPollDur is used in tests
|
||||
AcquireJobLongPollDur time.Duration
|
||||
|
||||
// HeartbeatInterval is the interval at which the provisioner daemon
|
||||
// will update its last seen at timestamp in the database.
|
||||
HeartbeatInterval time.Duration
|
||||
|
||||
// HeartbeatFn is the function that will be called at the interval
|
||||
// specified by HeartbeatInterval.
|
||||
// The default function just calls UpdateProvisionerDaemonLastSeenAt.
|
||||
// This is mainly used for testing.
|
||||
HeartbeatFn func(context.Context) error
|
||||
}
|
||||
|
||||
type server struct {
|
||||
@@ -85,6 +101,9 @@ type server struct {
|
||||
TimeNowFn func() time.Time
|
||||
|
||||
acquireJobLongPollDur time.Duration
|
||||
|
||||
heartbeatInterval time.Duration
|
||||
heartbeatFn func(ctx context.Context) error
|
||||
}
|
||||
|
||||
// We use the null byte (0x00) in generating a canonical map key for tags, so
|
||||
@@ -161,7 +180,21 @@ func NewServer(
|
||||
if options.AcquireJobLongPollDur == 0 {
|
||||
options.AcquireJobLongPollDur = DefaultAcquireJobLongPollDur
|
||||
}
|
||||
return &server{
|
||||
if options.HeartbeatInterval == 0 {
|
||||
options.HeartbeatInterval = DefaultHeartbeatInterval
|
||||
}
|
||||
// Avoid a nil check in s.heartbeat.
|
||||
if options.HeartbeatFn == nil {
|
||||
options.HeartbeatFn = func(hbCtx context.Context) error {
|
||||
//nolint:gocritic // This is specifically for updating the last seen at timestamp.
|
||||
return db.UpdateProvisionerDaemonLastSeenAt(dbauthz.AsSystemRestricted(hbCtx), database.UpdateProvisionerDaemonLastSeenAtParams{
|
||||
ID: id,
|
||||
LastSeenAt: sql.NullTime{Time: time.Now(), Valid: true},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
s := &server{
|
||||
lifecycleCtx: lifecycleCtx,
|
||||
AccessURL: accessURL,
|
||||
ID: id,
|
||||
@@ -182,7 +215,12 @@ func NewServer(
|
||||
OIDCConfig: options.OIDCConfig,
|
||||
TimeNowFn: options.TimeNowFn,
|
||||
acquireJobLongPollDur: options.AcquireJobLongPollDur,
|
||||
}, nil
|
||||
heartbeatInterval: options.HeartbeatInterval,
|
||||
heartbeatFn: options.HeartbeatFn,
|
||||
}
|
||||
|
||||
go s.heartbeatLoop()
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// timeNow should be used when trying to get the current time for math
|
||||
@@ -194,6 +232,50 @@ func (s *server) timeNow() time.Time {
|
||||
return dbtime.Now()
|
||||
}
|
||||
|
||||
// heartbeatLoop runs heartbeatOnce at the interval specified by HeartbeatInterval
|
||||
// until the lifecycle context is canceled.
|
||||
func (s *server) heartbeatLoop() {
|
||||
tick := time.NewTicker(time.Nanosecond)
|
||||
defer tick.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-s.lifecycleCtx.Done():
|
||||
s.Logger.Debug(s.lifecycleCtx, "heartbeat loop canceled")
|
||||
return
|
||||
case <-tick.C:
|
||||
if s.lifecycleCtx.Err() != nil {
|
||||
return
|
||||
}
|
||||
start := s.timeNow()
|
||||
hbCtx, hbCancel := context.WithTimeout(s.lifecycleCtx, s.heartbeatInterval)
|
||||
if err := s.heartbeat(hbCtx); err != nil {
|
||||
if !xerrors.Is(err, context.DeadlineExceeded) && !xerrors.Is(err, context.Canceled) {
|
||||
s.Logger.Error(hbCtx, "heartbeat failed", slog.Error(err))
|
||||
}
|
||||
}
|
||||
hbCancel()
|
||||
elapsed := s.timeNow().Sub(start)
|
||||
nextBeat := s.heartbeatInterval - elapsed
|
||||
// avoid negative interval
|
||||
if nextBeat <= 0 {
|
||||
nextBeat = time.Nanosecond
|
||||
}
|
||||
tick.Reset(nextBeat)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// heartbeat updates the last seen at timestamp in the database.
|
||||
// If HeartbeatFn is set, it will be called instead.
|
||||
func (s *server) heartbeat(ctx context.Context) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil
|
||||
default:
|
||||
return s.heartbeatFn(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
// AcquireJob queries the database to lock a job.
|
||||
//
|
||||
// Deprecated: This method is only available for back-level provisioner daemons.
|
||||
|
||||
@@ -66,7 +66,8 @@ func testUserQuietHoursScheduleStore() *atomic.Pointer[schedule.UserQuietHoursSc
|
||||
|
||||
func TestAcquireJob_LongPoll(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, _, _ := setup(t, false, &overrides{acquireJobLongPollDuration: time.Microsecond})
|
||||
//nolint:dogsled // ૮・ᴥ・ა
|
||||
srv, _, _, _ := setup(t, false, &overrides{acquireJobLongPollDuration: time.Microsecond})
|
||||
job, err := srv.AcquireJob(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, &proto.AcquiredJob{}, job)
|
||||
@@ -74,7 +75,8 @@ func TestAcquireJob_LongPoll(t *testing.T) {
|
||||
|
||||
func TestAcquireJobWithCancel_Cancel(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, _, _ := setup(t, false, nil)
|
||||
//nolint:dogsled // ૮ ˶′ﻌ ‵˶ ა
|
||||
srv, _, _, _ := setup(t, false, nil)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
fs := newFakeStream(ctx)
|
||||
@@ -95,6 +97,46 @@ func TestAcquireJobWithCancel_Cancel(t *testing.T) {
|
||||
require.Equal(t, "", job.JobId)
|
||||
}
|
||||
|
||||
func TestHeartbeat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
heartbeatChan := make(chan struct{})
|
||||
heartbeatFn := func(hbCtx context.Context) error {
|
||||
t.Logf("heartbeat")
|
||||
select {
|
||||
case <-hbCtx.Done():
|
||||
return hbCtx.Err()
|
||||
default:
|
||||
heartbeatChan <- struct{}{}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
//nolint:dogsled // 。:゚૮ ˶ˆ ﻌ ˆ˶ ა ゚:。
|
||||
_, _, _, _ = setup(t, false, &overrides{
|
||||
ctx: ctx,
|
||||
heartbeatFn: heartbeatFn,
|
||||
heartbeatInterval: testutil.IntervalFast,
|
||||
})
|
||||
|
||||
_, ok := <-heartbeatChan
|
||||
require.True(t, ok, "first heartbeat not received")
|
||||
_, ok = <-heartbeatChan
|
||||
require.True(t, ok, "second heartbeat not received")
|
||||
cancel()
|
||||
// Close the channel to ensure we don't receive any more heartbeats.
|
||||
// The test will fail if we do.
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
t.Fatalf("heartbeat received after cancel: %v", r)
|
||||
}
|
||||
}()
|
||||
|
||||
close(heartbeatChan)
|
||||
<-time.After(testutil.IntervalMedium)
|
||||
}
|
||||
|
||||
func TestAcquireJob(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -120,7 +162,7 @@ func TestAcquireJob(t *testing.T) {
|
||||
tc := tc
|
||||
t.Run(tc.name+"_InitiatorNotFound", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, _ := setup(t, false, nil)
|
||||
srv, db, _, _ := setup(t, false, nil)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
_, err := db.InsertProvisionerJob(context.Background(), database.InsertProvisionerJobParams{
|
||||
@@ -141,7 +183,7 @@ func TestAcquireJob(t *testing.T) {
|
||||
// deployment config.
|
||||
dv := &codersdk.DeploymentValues{MaxTokenLifetime: clibase.Duration(time.Hour)}
|
||||
gitAuthProvider := "github"
|
||||
srv, db, ps := setup(t, false, &overrides{
|
||||
srv, db, ps, _ := setup(t, false, &overrides{
|
||||
deploymentValues: dv,
|
||||
externalAuthConfigs: []*externalauth.Config{{
|
||||
ID: gitAuthProvider,
|
||||
@@ -359,7 +401,7 @@ func TestAcquireJob(t *testing.T) {
|
||||
|
||||
t.Run(tc.name+"_TemplateVersionDryRun", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, ps := setup(t, false, nil)
|
||||
srv, db, ps, _ := setup(t, false, nil)
|
||||
ctx := context.Background()
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
@@ -396,7 +438,7 @@ func TestAcquireJob(t *testing.T) {
|
||||
})
|
||||
t.Run(tc.name+"_TemplateVersionImport", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, ps := setup(t, false, nil)
|
||||
srv, db, ps, _ := setup(t, false, nil)
|
||||
ctx := context.Background()
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
@@ -427,7 +469,7 @@ func TestAcquireJob(t *testing.T) {
|
||||
})
|
||||
t.Run(tc.name+"_TemplateVersionImportWithUserVariable", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, ps := setup(t, false, nil)
|
||||
srv, db, ps, _ := setup(t, false, nil)
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{})
|
||||
@@ -476,7 +518,7 @@ func TestUpdateJob(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
t.Run("NotFound", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, _, _ := setup(t, false, nil)
|
||||
srv, _, _, _ := setup(t, false, nil)
|
||||
_, err := srv.UpdateJob(ctx, &proto.UpdateJobRequest{
|
||||
JobId: "hello",
|
||||
})
|
||||
@@ -489,7 +531,7 @@ func TestUpdateJob(t *testing.T) {
|
||||
})
|
||||
t.Run("NotRunning", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, _ := setup(t, false, nil)
|
||||
srv, db, _, _ := setup(t, false, nil)
|
||||
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
|
||||
ID: uuid.New(),
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
@@ -505,7 +547,7 @@ func TestUpdateJob(t *testing.T) {
|
||||
// This test prevents runners from updating jobs they don't own!
|
||||
t.Run("NotOwner", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, _ := setup(t, false, nil)
|
||||
srv, db, _, _ := setup(t, false, nil)
|
||||
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
|
||||
ID: uuid.New(),
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
@@ -548,9 +590,8 @@ func TestUpdateJob(t *testing.T) {
|
||||
|
||||
t.Run("Success", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
job := setupJob(t, db, srvID)
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
job := setupJob(t, db, pd.ID)
|
||||
_, err := srv.UpdateJob(ctx, &proto.UpdateJobRequest{
|
||||
JobId: job.String(),
|
||||
})
|
||||
@@ -559,9 +600,8 @@ func TestUpdateJob(t *testing.T) {
|
||||
|
||||
t.Run("Logs", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, ps := setup(t, false, &overrides{id: &srvID})
|
||||
job := setupJob(t, db, srvID)
|
||||
srv, db, ps, pd := setup(t, false, &overrides{})
|
||||
job := setupJob(t, db, pd.ID)
|
||||
|
||||
published := make(chan struct{})
|
||||
|
||||
@@ -585,9 +625,8 @@ func TestUpdateJob(t *testing.T) {
|
||||
})
|
||||
t.Run("Readme", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
job := setupJob(t, db, srvID)
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
job := setupJob(t, db, pd.ID)
|
||||
versionID := uuid.New()
|
||||
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
|
||||
ID: versionID,
|
||||
@@ -612,9 +651,8 @@ func TestUpdateJob(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
|
||||
defer cancel()
|
||||
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
job := setupJob(t, db, srvID)
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
job := setupJob(t, db, pd.ID)
|
||||
versionID := uuid.New()
|
||||
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
|
||||
ID: versionID,
|
||||
@@ -660,9 +698,8 @@ func TestUpdateJob(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
|
||||
defer cancel()
|
||||
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
job := setupJob(t, db, srvID)
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
job := setupJob(t, db, pd.ID)
|
||||
versionID := uuid.New()
|
||||
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
|
||||
ID: versionID,
|
||||
@@ -707,7 +744,7 @@ func TestFailJob(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
t.Run("NotFound", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, _, _ := setup(t, false, nil)
|
||||
srv, _, _, _ := setup(t, false, nil)
|
||||
_, err := srv.FailJob(ctx, &proto.FailedJob{
|
||||
JobId: "hello",
|
||||
})
|
||||
@@ -721,7 +758,7 @@ func TestFailJob(t *testing.T) {
|
||||
// This test prevents runners from updating jobs they don't own!
|
||||
t.Run("NotOwner", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, _ := setup(t, false, nil)
|
||||
srv, db, _, _ := setup(t, false, nil)
|
||||
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
|
||||
ID: uuid.New(),
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
@@ -744,8 +781,7 @@ func TestFailJob(t *testing.T) {
|
||||
})
|
||||
t.Run("AlreadyCompleted", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
|
||||
ID: uuid.New(),
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
@@ -755,7 +791,7 @@ func TestFailJob(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -780,8 +816,7 @@ func TestFailJob(t *testing.T) {
|
||||
//
|
||||
// (*Server).FailJob audit log - get build {"error": "sql: no rows in result set"}
|
||||
ignoreLogErrors := true
|
||||
srvID := uuid.New()
|
||||
srv, db, ps := setup(t, ignoreLogErrors, &overrides{id: &srvID})
|
||||
srv, db, ps, pd := setup(t, ignoreLogErrors, &overrides{})
|
||||
workspace, err := db.InsertWorkspace(ctx, database.InsertWorkspaceParams{
|
||||
ID: uuid.New(),
|
||||
AutomaticUpdates: database.AutomaticUpdatesNever,
|
||||
@@ -810,7 +845,7 @@ func TestFailJob(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -852,7 +887,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
t.Run("NotFound", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, _, _ := setup(t, false, nil)
|
||||
srv, _, _, _ := setup(t, false, nil)
|
||||
_, err := srv.CompleteJob(ctx, &proto.CompletedJob{
|
||||
JobId: "hello",
|
||||
})
|
||||
@@ -866,7 +901,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
// This test prevents runners from updating jobs they don't own!
|
||||
t.Run("NotOwner", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srv, db, _ := setup(t, false, nil)
|
||||
srv, db, _, _ := setup(t, false, nil)
|
||||
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
|
||||
ID: uuid.New(),
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
@@ -890,8 +925,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
|
||||
t.Run("TemplateImport_MissingGitAuth", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
jobID := uuid.New()
|
||||
versionID := uuid.New()
|
||||
err := db.InsertTemplateVersion(ctx, database.InsertTemplateVersionParams{
|
||||
@@ -909,7 +943,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -939,9 +973,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
|
||||
t.Run("TemplateImport_WithGitAuth", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{
|
||||
id: &srvID,
|
||||
srv, db, _, pd := setup(t, false, &overrides{
|
||||
externalAuthConfigs: []*externalauth.Config{{
|
||||
ID: "github",
|
||||
}},
|
||||
@@ -963,7 +995,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -1106,9 +1138,8 @@ func TestCompleteJob(t *testing.T) {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srvID := uuid.New()
|
||||
tss := &atomic.Pointer[schedule.TemplateScheduleStore]{}
|
||||
srv, db, ps := setup(t, false, &overrides{id: &srvID, templateScheduleStore: tss})
|
||||
srv, db, ps, pd := setup(t, false, &overrides{templateScheduleStore: tss})
|
||||
|
||||
var store schedule.TemplateScheduleStore = schedule.MockTemplateScheduleStore{
|
||||
GetFn: func(_ context.Context, _ database.Store, _ uuid.UUID) (schedule.TemplateScheduleOptions, error) {
|
||||
@@ -1123,10 +1154,19 @@ func TestCompleteJob(t *testing.T) {
|
||||
}
|
||||
tss.Store(&store)
|
||||
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
template := dbgen.Template(t, db, database.Template{
|
||||
Name: "template",
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
Name: "template",
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
|
||||
TemplateID: uuid.NullUUID{
|
||||
UUID: template.ID,
|
||||
Valid: true,
|
||||
},
|
||||
JobID: uuid.New(),
|
||||
})
|
||||
err := db.UpdateTemplateScheduleByID(ctx, database.UpdateTemplateScheduleByIDParams{
|
||||
ID: template.ID,
|
||||
@@ -1148,13 +1188,6 @@ func TestCompleteJob(t *testing.T) {
|
||||
TemplateID: template.ID,
|
||||
Ttl: workspaceTTL,
|
||||
})
|
||||
version := dbgen.TemplateVersion(t, db, database.TemplateVersion{
|
||||
TemplateID: uuid.NullUUID{
|
||||
UUID: template.ID,
|
||||
Valid: true,
|
||||
},
|
||||
JobID: uuid.New(),
|
||||
})
|
||||
build := dbgen.WorkspaceBuild(t, db, database.WorkspaceBuild{
|
||||
WorkspaceID: workspace.ID,
|
||||
TemplateVersionID: version.ID,
|
||||
@@ -1170,7 +1203,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
})
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -1315,19 +1348,17 @@ func TestCompleteJob(t *testing.T) {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
srvID := uuid.New()
|
||||
// Simulate the given time starting from now.
|
||||
require.False(t, c.now.IsZero())
|
||||
start := time.Now()
|
||||
tss := &atomic.Pointer[schedule.TemplateScheduleStore]{}
|
||||
uqhss := &atomic.Pointer[schedule.UserQuietHoursScheduleStore]{}
|
||||
srv, db, ps := setup(t, false, &overrides{
|
||||
srv, db, ps, pd := setup(t, false, &overrides{
|
||||
timeNowFn: func() time.Time {
|
||||
return c.now.Add(time.Since(start))
|
||||
},
|
||||
templateScheduleStore: tss,
|
||||
userQuietHoursScheduleStore: uqhss,
|
||||
id: &srvID,
|
||||
})
|
||||
|
||||
var templateScheduleStore schedule.TemplateScheduleStore = schedule.MockTemplateScheduleStore{
|
||||
@@ -1418,7 +1449,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
})
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -1484,8 +1515,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
})
|
||||
t.Run("TemplateDryRun", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
srvID := uuid.New()
|
||||
srv, db, _ := setup(t, false, &overrides{id: &srvID})
|
||||
srv, db, _, pd := setup(t, false, &overrides{})
|
||||
job, err := db.InsertProvisionerJob(ctx, database.InsertProvisionerJobParams{
|
||||
ID: uuid.New(),
|
||||
Provisioner: database.ProvisionerTypeEcho,
|
||||
@@ -1495,7 +1525,7 @@ func TestCompleteJob(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
_, err = db.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{
|
||||
WorkerID: uuid.NullUUID{
|
||||
UUID: srvID,
|
||||
UUID: pd.ID,
|
||||
Valid: true,
|
||||
},
|
||||
Types: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
@@ -1686,73 +1716,89 @@ func TestInsertWorkspaceResource(t *testing.T) {
|
||||
}
|
||||
|
||||
type overrides struct {
|
||||
ctx context.Context
|
||||
deploymentValues *codersdk.DeploymentValues
|
||||
externalAuthConfigs []*externalauth.Config
|
||||
id *uuid.UUID
|
||||
templateScheduleStore *atomic.Pointer[schedule.TemplateScheduleStore]
|
||||
userQuietHoursScheduleStore *atomic.Pointer[schedule.UserQuietHoursScheduleStore]
|
||||
timeNowFn func() time.Time
|
||||
acquireJobLongPollDuration time.Duration
|
||||
heartbeatFn func(ctx context.Context) error
|
||||
heartbeatInterval time.Duration
|
||||
}
|
||||
|
||||
func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisionerDaemonServer, database.Store, pubsub.Pubsub) {
|
||||
func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisionerDaemonServer, database.Store, pubsub.Pubsub, database.ProvisionerDaemon) {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
db := dbmem.New()
|
||||
ps := pubsub.NewInMemory()
|
||||
deploymentValues := &codersdk.DeploymentValues{}
|
||||
var externalAuthConfigs []*externalauth.Config
|
||||
srvID := uuid.New()
|
||||
tss := testTemplateScheduleStore()
|
||||
uqhss := testUserQuietHoursScheduleStore()
|
||||
var timeNowFn func() time.Time
|
||||
pollDur := time.Duration(0)
|
||||
if ov != nil {
|
||||
if ov.deploymentValues != nil {
|
||||
deploymentValues = ov.deploymentValues
|
||||
}
|
||||
if ov.externalAuthConfigs != nil {
|
||||
externalAuthConfigs = ov.externalAuthConfigs
|
||||
}
|
||||
if ov.id != nil {
|
||||
srvID = *ov.id
|
||||
}
|
||||
if ov.templateScheduleStore != nil {
|
||||
ttss := tss.Load()
|
||||
// keep the initial test value if the override hasn't set the atomic pointer.
|
||||
tss = ov.templateScheduleStore
|
||||
if tss.Load() == nil {
|
||||
swapped := tss.CompareAndSwap(nil, ttss)
|
||||
require.True(t, swapped)
|
||||
}
|
||||
}
|
||||
if ov.userQuietHoursScheduleStore != nil {
|
||||
tuqhss := uqhss.Load()
|
||||
// keep the initial test value if the override hasn't set the atomic pointer.
|
||||
uqhss = ov.userQuietHoursScheduleStore
|
||||
if uqhss.Load() == nil {
|
||||
swapped := uqhss.CompareAndSwap(nil, tuqhss)
|
||||
require.True(t, swapped)
|
||||
}
|
||||
}
|
||||
if ov.timeNowFn != nil {
|
||||
timeNowFn = ov.timeNowFn
|
||||
}
|
||||
pollDur = ov.acquireJobLongPollDuration
|
||||
if ov == nil {
|
||||
ov = &overrides{}
|
||||
}
|
||||
if ov.ctx == nil {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
ov.ctx = ctx
|
||||
}
|
||||
if ov.heartbeatInterval == 0 {
|
||||
ov.heartbeatInterval = testutil.IntervalMedium
|
||||
}
|
||||
if ov.deploymentValues != nil {
|
||||
deploymentValues = ov.deploymentValues
|
||||
}
|
||||
if ov.externalAuthConfigs != nil {
|
||||
externalAuthConfigs = ov.externalAuthConfigs
|
||||
}
|
||||
if ov.templateScheduleStore != nil {
|
||||
ttss := tss.Load()
|
||||
// keep the initial test value if the override hasn't set the atomic pointer.
|
||||
tss = ov.templateScheduleStore
|
||||
if tss.Load() == nil {
|
||||
swapped := tss.CompareAndSwap(nil, ttss)
|
||||
require.True(t, swapped)
|
||||
}
|
||||
}
|
||||
if ov.userQuietHoursScheduleStore != nil {
|
||||
tuqhss := uqhss.Load()
|
||||
// keep the initial test value if the override hasn't set the atomic pointer.
|
||||
uqhss = ov.userQuietHoursScheduleStore
|
||||
if uqhss.Load() == nil {
|
||||
swapped := uqhss.CompareAndSwap(nil, tuqhss)
|
||||
require.True(t, swapped)
|
||||
}
|
||||
}
|
||||
if ov.timeNowFn != nil {
|
||||
timeNowFn = ov.timeNowFn
|
||||
}
|
||||
pollDur = ov.acquireJobLongPollDuration
|
||||
|
||||
daemon, err := db.UpsertProvisionerDaemon(ov.ctx, database.UpsertProvisionerDaemonParams{
|
||||
Name: "test",
|
||||
CreatedAt: dbtime.Now(),
|
||||
Provisioners: []database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
Tags: database.StringMap{},
|
||||
LastSeenAt: sql.NullTime{},
|
||||
Version: "",
|
||||
APIVersion: "1.0",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
srv, err := provisionerdserver.NewServer(
|
||||
ctx,
|
||||
ov.ctx,
|
||||
&url.URL{},
|
||||
srvID,
|
||||
daemon.ID,
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: ignoreLogErrors}),
|
||||
[]database.ProvisionerType{database.ProvisionerTypeEcho},
|
||||
provisionerdserver.Tags{},
|
||||
provisionerdserver.Tags(daemon.Tags),
|
||||
db,
|
||||
ps,
|
||||
provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), db, ps),
|
||||
provisionerdserver.NewAcquirer(ov.ctx, logger.Named("acquirer"), db, ps),
|
||||
telemetry.NewNoop(),
|
||||
trace.NewNoopTracerProvider().Tracer("noop"),
|
||||
&atomic.Pointer[proto.QuotaCommitter]{},
|
||||
@@ -1765,10 +1811,12 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi
|
||||
TimeNowFn: timeNowFn,
|
||||
OIDCConfig: &oauth2.Config{},
|
||||
AcquireJobLongPollDur: pollDur,
|
||||
HeartbeatInterval: ov.heartbeatInterval,
|
||||
HeartbeatFn: ov.heartbeatFn,
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
return srv, db, ps
|
||||
return srv, db, ps, daemon
|
||||
}
|
||||
|
||||
func must[T any](value T, err error) T {
|
||||
|
||||
Reference in New Issue
Block a user