diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 5fb892171f..ca4e9a82e9 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -6921,6 +6921,19 @@ func (q *querier) LockChatByID(ctx context.Context, id uuid.UUID) (uuid.UUID, er return q.db.LockChatByID(ctx, id) } +func (q *querier) LockProvisionerKeyByIDForShare(ctx context.Context, id uuid.UUID) (uuid.UUID, error) { + // The lock query returns only the key ID, so fetch the key to authorize + // the read against its RBAC object. + key, err := q.db.GetProvisionerKeyByID(ctx, id) + if err != nil { + return uuid.Nil, err + } + if err := q.authorizeContext(ctx, policy.ActionRead, key); err != nil { + return uuid.Nil, err + } + return q.db.LockProvisionerKeyByIDForShare(ctx, id) +} + func (q *querier) MarkAllInboxNotificationsAsRead(ctx context.Context, arg database.MarkAllInboxNotificationsAsReadParams) error { resource := rbac.ResourceInboxNotification.WithOwner(arg.UserID.String()) diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 41c072c429..22dbc2618d 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -4488,6 +4488,13 @@ func (s *MethodTestSuite) TestProvisionerKeys() { dbm.EXPECT().GetProvisionerKeyByID(gomock.Any(), pk.ID).Return(pk, nil).AnyTimes() check.Args(pk.ID).Asserts(pk, policy.ActionRead).Returns(pk) })) + s.Run("LockProvisionerKeyByIDForShare", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + org := testutil.Fake(s.T(), faker, database.Organization{}) + pk := testutil.Fake(s.T(), faker, database.ProvisionerKey{OrganizationID: org.ID}) + dbm.EXPECT().GetProvisionerKeyByID(gomock.Any(), pk.ID).Return(pk, nil).AnyTimes() + dbm.EXPECT().LockProvisionerKeyByIDForShare(gomock.Any(), pk.ID).Return(pk.ID, nil).AnyTimes() + check.Args(pk.ID).Asserts(pk, policy.ActionRead).Returns(pk.ID) + })) s.Run("GetProvisionerKeyByHashedSecret", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { org := testutil.Fake(s.T(), faker, database.Organization{}) pk := testutil.Fake(s.T(), faker, database.ProvisionerKey{OrganizationID: org.ID, HashedSecret: []byte("foo")}) diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 871b5a7a9d..026c6ae069 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -4857,6 +4857,14 @@ func (m queryMetricsStore) LockChatByID(ctx context.Context, id uuid.UUID) (uuid return r0, r1 } +func (m queryMetricsStore) LockProvisionerKeyByIDForShare(ctx context.Context, id uuid.UUID) (uuid.UUID, error) { + start := time.Now() + r0, r1 := m.s.LockProvisionerKeyByIDForShare(ctx, id) + m.queryLatencies.WithLabelValues("LockProvisionerKeyByIDForShare").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "LockProvisionerKeyByIDForShare").Inc() + return r0, r1 +} + func (m queryMetricsStore) MarkAllInboxNotificationsAsRead(ctx context.Context, arg database.MarkAllInboxNotificationsAsReadParams) error { start := time.Now() r0 := m.s.MarkAllInboxNotificationsAsRead(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index ef4755b318..4c459c231f 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -9178,6 +9178,21 @@ func (mr *MockStoreMockRecorder) LockChatByID(ctx, id any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LockChatByID", reflect.TypeOf((*MockStore)(nil).LockChatByID), ctx, id) } +// LockProvisionerKeyByIDForShare mocks base method. +func (m *MockStore) LockProvisionerKeyByIDForShare(ctx context.Context, id uuid.UUID) (uuid.UUID, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "LockProvisionerKeyByIDForShare", ctx, id) + ret0, _ := ret[0].(uuid.UUID) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// LockProvisionerKeyByIDForShare indicates an expected call of LockProvisionerKeyByIDForShare. +func (mr *MockStoreMockRecorder) LockProvisionerKeyByIDForShare(ctx, id any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LockProvisionerKeyByIDForShare", reflect.TypeOf((*MockStore)(nil).LockProvisionerKeyByIDForShare), ctx, id) +} + // MarkAllInboxNotificationsAsRead mocks base method. func (m *MockStore) MarkAllInboxNotificationsAsRead(ctx context.Context, arg database.MarkAllInboxNotificationsAsReadParams) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 0ebccc223a..7e052926c4 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -1280,6 +1280,11 @@ type sqlcQuerier interface { // allocate a new snapshot version in one round trip. LockChatAndBumpSnapshotVersion(ctx context.Context, id uuid.UUID) (Chat, error) LockChatByID(ctx context.Context, id uuid.UUID) (uuid.UUID, error) + // Locks the provisioner key row with FOR KEY SHARE for the remainder of the + // current transaction. FOR KEY SHARE conflicts with DELETE, so while the lock + // is held the key cannot be deleted, and a committed deletion is observed as + // no rows by later calls. + LockProvisionerKeyByIDForShare(ctx context.Context, id uuid.UUID) (uuid.UUID, error) MarkAllInboxNotificationsAsRead(ctx context.Context, arg MarkAllInboxNotificationsAsReadParams) error // Flips active, already-hydrated chats for an agent to dirty when the // agent's latest snapshot hash differs from the chat's pinned hash. The diff --git a/coderd/database/querier_test.go b/coderd/database/querier_test.go index b6ac7c3c58..d416a3f688 100644 --- a/coderd/database/querier_test.go +++ b/coderd/database/querier_test.go @@ -2613,6 +2613,27 @@ func TestAcquireProvisionerJob(t *testing.T) { }) require.ErrorIs(t, err, sql.ErrNoRows) }) + + t.Run("ProvisionerKeyLock", func(t *testing.T) { + t.Parallel() + var ( + db, _ = dbtestutil.NewDB(t) + ctx = testutil.Context(t, testutil.WaitMedium) + org = dbgen.Organization(t, db, database.Organization{}) + key = dbgen.ProvisionerKey(t, db, database.ProvisionerKey{OrganizationID: org.ID}) + ) + + // While the key exists, the lock returns its ID. + id, err := db.LockProvisionerKeyByIDForShare(ctx, key.ID) + require.NoError(t, err) + require.Equal(t, key.ID, id) + + // Once the key is deleted, the lock reports no rows. + err = db.DeleteProvisionerKey(ctx, key.ID) + require.NoError(t, err) + _, err = db.LockProvisionerKeyByIDForShare(ctx, key.ID) + require.ErrorIs(t, err, sql.ErrNoRows) + }) } func TestUserLastSeenFilter(t *testing.T) { diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 9434475b87..4697e3499a 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -23804,6 +23804,27 @@ func (q *sqlQuerier) ListProvisionerKeysByOrganizationExcludeReserved(ctx contex return items, nil } +const lockProvisionerKeyByIDForShare = `-- name: LockProvisionerKeyByIDForShare :one +SELECT + id +FROM + provisioner_keys +WHERE + id = $1 +FOR KEY SHARE +` + +// Locks the provisioner key row with FOR KEY SHARE for the remainder of the +// current transaction. FOR KEY SHARE conflicts with DELETE, so while the lock +// is held the key cannot be deleted, and a committed deletion is observed as +// no rows by later calls. +func (q *sqlQuerier) LockProvisionerKeyByIDForShare(ctx context.Context, id uuid.UUID) (uuid.UUID, error) { + row := q.db.QueryRowContext(ctx, lockProvisionerKeyByIDForShare, id) + var id_2 uuid.UUID + err := row.Scan(&id_2) + return id_2, err +} + const getWorkspaceProxies = `-- name: GetWorkspaceProxies :many SELECT id, name, display_name, icon, url, wildcard_hostname, created_at, updated_at, deleted, token_hashed_secret, region_id, derp_enabled, derp_only, version diff --git a/coderd/database/queries/provisionerkeys.sql b/coderd/database/queries/provisionerkeys.sql index 0bf95069dd..6e89104e2b 100644 --- a/coderd/database/queries/provisionerkeys.sql +++ b/coderd/database/queries/provisionerkeys.sql @@ -19,6 +19,19 @@ FROM WHERE id = $1; +-- name: LockProvisionerKeyByIDForShare :one +-- Locks the provisioner key row with FOR KEY SHARE for the remainder of the +-- current transaction. FOR KEY SHARE conflicts with DELETE, so while the lock +-- is held the key cannot be deleted, and a committed deletion is observed as +-- no rows by later calls. +SELECT + id +FROM + provisioner_keys +WHERE + id = $1 +FOR KEY SHARE; + -- name: GetProvisionerKeyByHashedSecret :one SELECT * diff --git a/coderd/provisionerdserver/acquirer.go b/coderd/provisionerdserver/acquirer.go index e082a9651e..e0779a177f 100644 --- a/coderd/provisionerdserver/acquirer.go +++ b/coderd/provisionerdserver/acquirer.go @@ -15,9 +15,11 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" "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/coder/v2/codersdk" "github.com/coder/quartz" ) @@ -61,9 +63,12 @@ func WithClock(clock quartz.Clock) AcquirerOption { } } -// AcquirerStore is the subset of database.Store that the Acquirer needs +// AcquirerStore is the subset of database.Store that the Acquirer needs. Job +// acquisition runs in a transaction that locks the worker's deletable +// provisioner key (LockProvisionerKeyByIDForShare) before claiming a job +// (AcquireProvisionerJob), so a claim cannot commit after the key's deletion. type AcquirerStore interface { - AcquireProvisionerJob(context.Context, database.AcquireProvisionerJobParams) (database.ProvisionerJob, error) + InTx(func(database.Store) error, *database.TxOptions) error } func NewAcquirer(ctx context.Context, logger slog.Logger, store AcquirerStore, ps pubsub.Pubsub, @@ -88,11 +93,15 @@ func NewAcquirer(ctx context.Context, logger slog.Logger, store AcquirerStore, p // tags from the database. The call blocks until a job is acquired, the context is // done, or the database returns an error _other_ than that no jobs are available. // If no jobs are available, this method handles retrying as appropriate. +// When keyID is a deletable provisioner key, the claim only succeeds while +// that key row still exists. Reserved keys and the zero value are not +// checked, as they have no row to delete. func (a *Acquirer) AcquireJob( - ctx context.Context, organization uuid.UUID, worker uuid.UUID, pt []database.ProvisionerType, tags Tags, + ctx context.Context, organization uuid.UUID, worker uuid.UUID, pt []database.ProvisionerType, tags Tags, keyID uuid.UUID, ) ( retJob database.ProvisionerJob, retErr error, ) { + deletableKey := codersdk.IsDeletableProvisionerKey(keyID) logger := a.logger.With( slog.F("organization_id", organization), slog.F("worker_id", worker), @@ -120,19 +129,52 @@ func (a *Acquirer) AcquireJob( return database.ProvisionerJob{}, err case <-clearance: logger.Debug(ctx, "got clearance to call database") - job, err := a.store.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{ - OrganizationID: organization, - StartedAt: sql.NullTime{ - Time: dbtime.Now(), - Valid: true, - }, - WorkerID: uuid.NullUUID{ - UUID: worker, - Valid: true, - }, - Types: pt, - ProvisionerTags: dbTags, - }) + var job database.ProvisionerJob + err := a.store.InTx(func(tx database.Store) error { + if deletableKey { + // Lock the key for the rest of the transaction so the claim + // below cannot commit after the key's deletion. A missing row + // means the key was deleted. + _, err := tx.LockProvisionerKeyByIDForShare( + //nolint:gocritic // The acquire context has no actor that can + // read provisioner keys, so scope the read to this narrow subject. + dbauthz.AsSystemReadProvisionerDaemons(ctx), keyID) + if xerrors.Is(err, sql.ErrNoRows) { + return ErrProvisionerKeyDeleted + } + if err != nil { + return xerrors.Errorf("lock provisioner key: %w", err) + } + } + acquired, err := tx.AcquireProvisionerJob(ctx, database.AcquireProvisionerJobParams{ + OrganizationID: organization, + StartedAt: sql.NullTime{ + Time: dbtime.Now(), + Valid: true, + }, + WorkerID: uuid.NullUUID{ + UUID: worker, + Valid: true, + }, + Types: pt, + ProvisionerTags: dbTags, + }) + if err != nil { + return err + } + job = acquired + return nil + }, nil) + if xerrors.Is(err, ErrProvisionerKeyDeleted) { + logger.Debug(ctx, "provisioner key deleted, exiting acquire") + // cancel (not done) hands an in-progress clearance to another + // acquiree in the domain, re-dispatching the wakeup this + // acquiree consumed. + if internalError := a.cancel(dk, clearance); internalError != nil { + return database.ProvisionerJob{}, internalError + } + return database.ProvisionerJob{}, ErrProvisionerKeyDeleted + } if xerrors.Is(err, sql.ErrNoRows) { logger.Debug(ctx, "no job available") continue diff --git a/coderd/provisionerdserver/acquirer_test.go b/coderd/provisionerdserver/acquirer_test.go index 3198fad25c..03fce24730 100644 --- a/coderd/provisionerdserver/acquirer_test.go +++ b/coderd/provisionerdserver/acquirer_test.go @@ -15,6 +15,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/goleak" + "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbtestutil" @@ -118,6 +119,123 @@ func TestAcquirer_MultipleSameDomain(t *testing.T) { require.Equal(t, workerIDs, gotWorkerCalls) } +// TestAcquirer_ProvisionerKeyDeleted verifies that an acquiree whose deletable +// key no longer exists exits with ErrProvisionerKeyDeleted and hands its +// clearance to another acquiree in the same domain. +func TestAcquirer_ProvisionerKeyDeleted(t *testing.T) { + t.Parallel() + fs := newFakeOrderedStore() + ps := pubsub.NewInMemory() + ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) + defer cancel() + logger := testutil.Logger(t) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + + orgID := uuid.New() + pt := []database.ProvisionerType{database.ProvisionerTypeEcho} + tags := provisionerdserver.Tags{"environment": "on-prem"} + + // The keyed acquiree starts first; as the domain's first member it gets + // immediate clearance and blocks in the key lock. + keyed := newTestAcquiree(t, orgID, uuid.New(), pt, tags) + keyed.startAcquireWithKey(ctx, uut, uuid.New()) + require.Eventually(t, func() bool { return fs.lockCallCount() == 1 }, testutil.WaitShort, testutil.IntervalFast) + + // The unkeyed acquiree joins the same domain and parks without clearance. + unkeyed := newTestAcquiree(t, orgID, uuid.New(), pt, tags) + unkeyed.startAcquire(ctx, uut) + + // Release the lock with no rows: the key is deleted. + err := fs.sendLock(ctx, sql.ErrNoRows) + require.NoError(t, err) + + // The keyed acquiree exits terminally rather than re-parking, without + // ever attempting a claim. Count only its own calls: its clearance passes + // to the unkeyed acquiree, whose call can land before this assertion. + select { + case <-ctx.Done(): + t.Fatal("timeout waiting for keyed acquiree to exit") + case err := <-keyed.ec: + require.ErrorIs(t, err, provisionerdserver.ErrProvisionerKeyDeleted) + } + <-keyed.jc + require.Equal(t, 0, fs.callCountForWorker(keyed.workerID)) + + // Its clearance is handed to the unkeyed acquiree, which claims a job + // without a new posting or backup poll. + jobID := uuid.New() + err = fs.sendCtx(ctx, database.ProvisionerJob{ID: jobID}, nil) + require.NoError(t, err) + job := unkeyed.success(ctx) + require.Equal(t, jobID, job.ID) +} + +// TestAcquirer_ProvisionerKeyExists verifies that a no-rows acquire result +// with the key still present re-parks the acquiree; ErrProvisionerKeyDeleted +// requires the key lock to find no row. +func TestAcquirer_ProvisionerKeyExists(t *testing.T) { + t.Parallel() + fs := newFakeOrderedStore() + ps := pubsub.NewInMemory() + ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) + defer cancel() + logger := testutil.Logger(t) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + + orgID := uuid.New() + pt := []database.ProvisionerType{database.ProvisionerTypeEcho} + tags := provisionerdserver.Tags{"environment": "on-prem"} + + acquiree := newTestAcquiree(t, orgID, uuid.New(), pt, tags) + jobID := uuid.New() + // Two acquire rounds: the first locks the key and finds no job, the + // second locks the key and claims the job. + require.NoError(t, fs.sendLock(ctx, nil)) + require.NoError(t, fs.sendCtx(ctx, database.ProvisionerJob{}, sql.ErrNoRows)) + require.NoError(t, fs.sendLock(ctx, nil)) + require.NoError(t, fs.sendCtx(ctx, database.ProvisionerJob{ID: jobID}, nil)) + acquiree.startAcquireWithKey(ctx, uut, uuid.New()) + require.Eventually(t, func() bool { return fs.callCount() == 1 }, testutil.WaitShort, testutil.IntervalFast) + acquiree.requireBlocked() + + // A compatible posting wakes the parked acquiree and it claims the job. + postJob(t, ps, database.ProvisionerTypeEcho, provisionerdserver.Tags{}) + job := acquiree.success(ctx) + require.Equal(t, jobID, job.ID) +} + +// TestAcquirer_ProvisionerKeyCheckError verifies that a transient error from +// the key lock fails the acquire with a plain error, not the key-deleted +// sentinel. +func TestAcquirer_ProvisionerKeyCheckError(t *testing.T) { + t.Parallel() + fs := newFakeOrderedStore() + ps := pubsub.NewInMemory() + ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) + defer cancel() + logger := testutil.Logger(t) + uut := provisionerdserver.NewAcquirer(ctx, logger.Named("acquirer"), fs, ps) + + orgID := uuid.New() + pt := []database.ProvisionerType{database.ProvisionerTypeEcho} + tags := provisionerdserver.Tags{"environment": "on-prem"} + + acquiree := newTestAcquiree(t, orgID, uuid.New(), pt, tags) + require.NoError(t, fs.sendLock(ctx, xerrors.New("transient database error"))) + acquiree.startAcquireWithKey(ctx, uut, uuid.New()) + + select { + case <-ctx.Done(): + t.Fatal("timeout waiting for acquiree to exit") + case err := <-acquiree.ec: + require.Error(t, err) + require.NotErrorIs(t, err, provisionerdserver.ErrProvisionerKeyDeleted) + require.ErrorContains(t, err, "lock provisioner key") + } + <-acquiree.jc + require.Equal(t, 0, fs.callCount()) +} + // TestAcquirer_WaitsOnNoJobs tests that after a call that returns no jobs, Acquirer waits for a new // job posting before retrying func TestAcquirer_WaitsOnNoJobs(t *testing.T) { @@ -523,12 +641,11 @@ func TestAcquirer_MatchTags(t *testing.T) { if tt.unmatchedOrg { acquireOrgID = uuid.New() } - 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) + aj, err := acq.AcquireJob(ctx, acquireOrgID, uuid.New(), ptypes, tt.acquireJobTags, uuid.Nil) assert.NoError(t, err) assert.Equal(t, pj.ID, aj.ID) return @@ -622,15 +739,31 @@ func (s *acquirerStoreSpy) AcquireProvisionerJob( return job, err } +// InTx re-wraps the transaction handle in a spy so the acquirer's +// AcquireProvisionerJob call inside the transaction is still observed. +func (s *acquirerStoreSpy) InTx(fn func(database.Store) error, opts *database.TxOptions) error { + return s.Store.InTx(func(tx database.Store) error { + return fn(&acquirerStoreSpy{Store: tx, callCompleted: s.callCompleted}) + }, opts) +} + // fakeOrderedStore is a fake store that lets tests send AcquireProvisionerJob -// results in order over a channel, and tests for overlapped calls. +// results in order over a channel, and tests for overlapped calls. Keyed +// acquires also block in LockProvisionerKeyByIDForShare until a result is +// sent over lockResults. The embedded Store panics on any other method. type fakeOrderedStore struct { + database.Store jobs chan database.ProvisionerJob errors chan error callStarted chan struct{} + // lockResults releases LockProvisionerKeyByIDForShare calls: nil locks + // the key, sql.ErrNoRows reads as deleted, other errors are returned + // as-is. + lockResults chan error - mu sync.Mutex - params []database.AcquireProvisionerJobParams + mu sync.Mutex + params []database.AcquireProvisionerJobParams + lockCalls int // inflight and overlaps track whether any calls from workers overlap with // one another @@ -645,10 +778,17 @@ func newFakeOrderedStore() *fakeOrderedStore { jobs: make(chan database.ProvisionerJob, 100), errors: make(chan error, 100), callStarted: make(chan struct{}, 100), + lockResults: make(chan error, 100), inflight: make(map[uuid.UUID]bool), } } +// InTx runs fn against the fake itself; the fake does not implement +// transactional semantics. +func (s *fakeOrderedStore) InTx(fn func(database.Store) error, _ *database.TxOptions) error { + return fn(s) +} + func (s *fakeOrderedStore) AcquireProvisionerJob( _ context.Context, params database.AcquireProvisionerJobParams, ) ( @@ -676,6 +816,51 @@ func (s *fakeOrderedStore) AcquireProvisionerJob( return job, err } +func (s *fakeOrderedStore) LockProvisionerKeyByIDForShare(_ context.Context, id uuid.UUID) (uuid.UUID, error) { + s.mu.Lock() + s.lockCalls++ + s.mu.Unlock() + if err := <-s.lockResults; err != nil { + return uuid.Nil, err + } + return id, nil +} + +func (s *fakeOrderedStore) callCount() int { + s.mu.Lock() + defer s.mu.Unlock() + return len(s.params) +} + +// callCountForWorker returns the number of AcquireProvisionerJob calls made on +// behalf of workerID. +func (s *fakeOrderedStore) callCountForWorker(workerID uuid.UUID) int { + s.mu.Lock() + defer s.mu.Unlock() + var n int + for _, p := range s.params { + if p.WorkerID.UUID == workerID { + n++ + } + } + return n +} + +func (s *fakeOrderedStore) lockCallCount() int { + s.mu.Lock() + defer s.mu.Unlock() + return s.lockCalls +} + +func (s *fakeOrderedStore) sendLock(ctx context.Context, err error) error { + select { + case <-ctx.Done(): + return ctx.Err() + case s.lockResults <- err: + return nil + } +} + func (s *fakeOrderedStore) sendCtx(ctx context.Context, job database.ProvisionerJob, err error) error { select { case <-ctx.Done(): @@ -694,8 +879,10 @@ func (s *fakeOrderedStore) sendCtx(ctx context.Context, job database.Provisioner // fakeTaggedStore is a test store that allows tests to specify which jobs are // available, and returns them to callers with the appropriate provisioner type -// and tags. It doesn't care about the order. +// and tags. It doesn't care about the order. The embedded Store panics on any +// unstubbed method. type fakeTaggedStore struct { + database.Store t *testing.T mu sync.Mutex jobs []database.ProvisionerJob @@ -743,6 +930,12 @@ jobLoop: return database.ProvisionerJob{}, sql.ErrNoRows } +// InTx runs fn against the fake itself; the fake does not implement +// transactional semantics. +func (s *fakeTaggedStore) InTx(fn func(database.Store) error, _ *database.TxOptions) error { + return fn(s) +} + // testAcquiree is a helper type that handles asynchronously calling AcquireJob // and asserting whether or not it returns, blocks, or is canceled. type testAcquiree struct { @@ -768,8 +961,12 @@ func newTestAcquiree(t *testing.T, orgID uuid.UUID, workerID uuid.UUID, pt []dat } func (a *testAcquiree) startAcquire(ctx context.Context, uut *provisionerdserver.Acquirer) { + a.startAcquireWithKey(ctx, uut, uuid.Nil) +} + +func (a *testAcquiree) startAcquireWithKey(ctx context.Context, uut *provisionerdserver.Acquirer, keyID uuid.UUID) { go func() { - j, e := uut.AcquireJob(ctx, a.orgID, a.workerID, a.pt, a.tags) + j, e := uut.AcquireJob(ctx, a.orgID, a.workerID, a.pt, a.tags, keyID) a.ec <- e a.jc <- j }() diff --git a/coderd/provisionerdserver/provisionerdserver.go b/coderd/provisionerdserver/provisionerdserver.go index 34e0a891af..4bfe0b9acd 100644 --- a/coderd/provisionerdserver/provisionerdserver.go +++ b/coderd/provisionerdserver/provisionerdserver.go @@ -15,6 +15,7 @@ import ( "sort" "strconv" "strings" + "sync" "sync/atomic" "time" @@ -79,6 +80,14 @@ type Options struct { ExternalAuthConfigs []*externalauth.Config AISeatTracker aiseats.SeatTracker + // KeyID is the provisioner key the daemon authenticated with, or the + // zero value if it did not authenticate with a key. + KeyID uuid.UUID + + // SessionCancel terminates the daemon's session. Required when KeyID is a + // deletable provisioner key; optional otherwise. + SessionCancel context.CancelFunc + // Clock for testing Clock quartz.Clock @@ -104,6 +113,8 @@ type server struct { lifecycleCtx context.Context AccessURL *url.URL ID uuid.UUID + KeyID uuid.UUID + sessionCancel context.CancelFunc OrganizationID uuid.UUID Logger slog.Logger Provisioners []database.ProvisionerType @@ -134,6 +145,16 @@ type server struct { heartbeatInterval time.Duration heartbeatFn func(ctx context.Context) error + // jobMu guards activeJobs and terminationPending. + jobMu sync.Mutex + // activeJobs tracks jobs claimed by this session that have not yet been + // completed or failed. The in-tree provisioner daemon runs jobs serially, + // so at most one entry is expected; the protocol does not enforce this. + activeJobs map[uuid.UUID]struct{} + // terminationPending records a termination request that arrived while a + // job was active; it is performed when the last active job finishes. + terminationPending bool + metrics *Metrics } @@ -142,6 +163,10 @@ type server struct { var ErrTagsContainNullByte = xerrors.New("tags cannot contain the null byte (0x00)") +// ErrProvisionerKeyDeleted is returned from job acquisition when the +// provisioner key the daemon authenticated with no longer exists. +var ErrProvisionerKeyDeleted = xerrors.New("provisioner key was deleted") + type Tags map[string]string func (t Tags) ToJSON() (json.RawMessage, error) { @@ -161,6 +186,15 @@ func (t Tags) Valid() error { return nil } +// Server is the provisioner daemon DRPC server plus session-lifecycle hooks +// used by the serve handlers. +type Server interface { + proto.DRPCProvisionerDaemonServer + // TerminateSession cancels the session once no acquired job is + // active. Safe to call from any goroutine. + TerminateSession() +} + func NewServer( lifecycleCtx context.Context, apiVersion string, @@ -186,11 +220,16 @@ func NewServer( prebuildsOrchestrator *atomic.Pointer[prebuilds.ReconciliationOrchestrator], metrics *Metrics, experiments codersdk.Experiments, -) (proto.DRPCProvisionerDaemonServer, error) { +) (Server, error) { // Fail-fast if pointers are nil if lifecycleCtx == nil { return nil, xerrors.New("ctx is nil") } + // A deletable key's session must be cancelable, otherwise key deletion + // cannot terminate it. + if codersdk.IsDeletableProvisionerKey(options.KeyID) && options.SessionCancel == nil { + return nil, xerrors.New("SessionCancel is required when KeyID is a deletable provisioner key") + } if quotaCommitter == nil { return nil, xerrors.New("quotaCommitter is nil") } @@ -236,6 +275,8 @@ func NewServer( apiVersion: apiVersion, AccessURL: accessURL, ID: id, + KeyID: options.KeyID, + sessionCancel: options.SessionCancel, OrganizationID: organizationID, Logger: logger, Provisioners: provisioners, @@ -260,6 +301,7 @@ func NewServer( PrebuildsOrchestrator: prebuildsOrchestrator, UsageInserter: usageInserter, AISeatTracker: options.AISeatTracker, + activeJobs: map[uuid.UUID]struct{}{}, metrics: metrics, Experiments: experiments, } @@ -297,6 +339,16 @@ func (s *server) heartbeatLoop() { if err := s.heartbeat(hbCtx); err != nil && !database.IsQueryCanceledError(err) { s.Logger.Warn(hbCtx, "heartbeat failed", slog.Error(err)) } + // The key check rides the heartbeat tick so a session whose deletable + // key is gone terminates within one interval. Transient errors are + // logged and the session is left running. + if deleted, err := s.keyDeleted(hbCtx); err != nil && !database.IsQueryCanceledError(err) { + s.Logger.Warn(hbCtx, "check provisioner key on heartbeat", slog.Error(err)) + } else if deleted { + s.Logger.Warn(hbCtx, "provisioner key deleted, canceling session", + slog.F("provisioner_key_id", s.KeyID)) + s.TerminateSession() + } hbCancel() elapsed := s.timeNow().Sub(start) nextBeat := s.heartbeatInterval - elapsed @@ -328,27 +380,105 @@ func (s *server) defaultHeartbeat(ctx context.Context) error { }) } +// keyDeleted reports whether the provisioner key no longer exists. +func (s *server) keyDeleted(ctx context.Context) (bool, error) { + if !codersdk.IsDeletableProvisionerKey(s.KeyID) { + return false, nil + } + _, err := s.Database.GetProvisionerKeyByID( + //nolint:gocritic // Callers' contexts cannot read provisioner keys + // (provisionerd actor or no actor at all), so scope the read to this + // narrow subject. + dbauthz.AsSystemReadProvisionerDaemons(ctx), s.KeyID) + if errors.Is(err, sql.ErrNoRows) { + return true, nil + } + if err != nil { + return false, xerrors.Errorf("get provisioner key: %w", err) + } + return false, nil +} + +// TerminateSession cancels the session. Cancellation is deferred while a job +// claimed by this session is active so the daemon can report the job's +// result; the last active job's completion performs it. Requires a configured +// sessionCancel. +func (s *server) TerminateSession() { + s.jobMu.Lock() + if len(s.activeJobs) > 0 { + s.terminationPending = true + s.jobMu.Unlock() + s.Logger.Info(s.lifecycleCtx, "deferring session cancellation until active jobs finish", + slog.F("provisioner_key_id", s.KeyID)) + return + } + s.jobMu.Unlock() + s.sessionCancel() +} + +// jobStarted records a job claimed by this session as active. +func (s *server) jobStarted(id uuid.UUID) { + s.jobMu.Lock() + defer s.jobMu.Unlock() + s.activeJobs[id] = struct{}{} +} + +// jobFinished removes an active job and performs a termination deferred while +// jobs were active. The daemon may not receive the final RPC response when +// this cancels the session; the job's outcome is already persisted. +func (s *server) jobFinished(id uuid.UUID) { + s.jobMu.Lock() + delete(s.activeJobs, id) + terminate := s.terminationPending && len(s.activeJobs) == 0 + s.jobMu.Unlock() + if !terminate { + return + } + s.Logger.Warn(s.lifecycleCtx, "canceling session after job completion", + slog.F("provisioner_key_id", s.KeyID)) + s.sessionCancel() +} + // AcquireJob queries the database to lock a job. // // Deprecated: This method is only available for back-level provisioner daemons. func (s *server) AcquireJob(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { //nolint:gocritic // Provisionerd has specific authz rules. ctx = dbauthz.AsProvisionerd(ctx) + if deleted, err := s.keyDeleted(ctx); err != nil { + return nil, xerrors.Errorf("acquire job: check provisioner key: %w", err) + } else if deleted { + s.Logger.Warn(ctx, "provisioner key deleted, rejecting job acquisition", + slog.F("provisioner_key_id", s.KeyID)) + s.TerminateSession() + return nil, xerrors.Errorf("acquire job: %w", ErrProvisionerKeyDeleted) + } // Since AcquireJob blocks until a job is available, we set a long (5s by default) timeout. This allows back-level // provisioner daemons to gracefully shut down within a few seconds, but keeps them from rapidly polling the // database. acqCtx, acqCancel := context.WithTimeout(ctx, s.acquireJobLongPollDur) defer acqCancel() - job, err := s.Acquirer.AcquireJob(acqCtx, s.OrganizationID, s.ID, s.Provisioners, s.Tags) + job, err := s.Acquirer.AcquireJob(acqCtx, s.OrganizationID, s.ID, s.Provisioners, s.Tags, s.KeyID) if database.IsQueryCanceledError(err) { s.Logger.Debug(ctx, "successful cancel") return &proto.AcquiredJob{}, nil } + if errors.Is(err, ErrProvisionerKeyDeleted) { + s.Logger.Warn(ctx, "provisioner key deleted, rejecting job acquisition", + slog.F("provisioner_key_id", s.KeyID)) + s.TerminateSession() + } if err != nil { return nil, xerrors.Errorf("acquire job: %w", err) } s.Logger.Debug(ctx, "locked job from database", slog.F("job_id", job.ID)) - return s.acquireProtoJob(ctx, job) + s.jobStarted(job.ID) + pj, err := s.acquireProtoJob(ctx, job) + if err != nil { + s.jobFinished(job.ID) + return nil, err + } + return pj, nil } type jobAndErr struct { @@ -367,6 +497,14 @@ func (s *server) AcquireJobWithCancel(stream proto.DRPCProvisionerDaemon_Acquire retErr = closeErr } }() + if deleted, err := s.keyDeleted(streamCtx); err != nil { + return xerrors.Errorf("acquire job: check provisioner key: %w", err) + } else if deleted { + s.Logger.Warn(streamCtx, "provisioner key deleted, rejecting job acquisition", + slog.F("provisioner_key_id", s.KeyID)) + s.TerminateSession() + return xerrors.Errorf("acquire job: %w", ErrProvisionerKeyDeleted) + } acqCtx, acqCancel := context.WithCancel(streamCtx) defer acqCancel() recvCh := make(chan error, 1) @@ -376,7 +514,7 @@ func (s *server) AcquireJobWithCancel(stream proto.DRPCProvisionerDaemon_Acquire }() jec := make(chan jobAndErr, 1) go func() { - job, err := s.Acquirer.AcquireJob(acqCtx, s.OrganizationID, s.ID, s.Provisioners, s.Tags) + job, err := s.Acquirer.AcquireJob(acqCtx, s.OrganizationID, s.ID, s.Provisioners, s.Tags, s.KeyID) jec <- jobAndErr{job: job, err: err} }() var recvErr error @@ -397,67 +535,93 @@ func (s *server) AcquireJobWithCancel(stream proto.DRPCProvisionerDaemon_Acquire } return nil } + if errors.Is(je.err, ErrProvisionerKeyDeleted) { + s.Logger.Warn(streamCtx, "provisioner key deleted, rejecting job acquisition", + slog.F("provisioner_key_id", s.KeyID)) + s.TerminateSession() + } if je.err != nil { return xerrors.Errorf("acquire job: %w", je.err) } logger := s.Logger.With(slog.F("job_id", je.job.ID)) logger.Debug(streamCtx, "locked job from database") + s.jobStarted(je.job.ID) if recvErr != nil { logger.Error(streamCtx, "recv error and failed to cancel acquire job", slog.Error(recvErr)) // Well, this is awkward. We hit an error receiving from the stream, but didn't cancel before we locked a job // in the database. We need to mark this job as failed so the end user can retry if they want to. - now := s.timeNow() - err := s.Database.UpdateProvisionerJobWithCompleteByID( - //nolint:gocritic // Provisionerd has specific authz rules. - dbauthz.AsProvisionerd(context.Background()), - database.UpdateProvisionerJobWithCompleteByIDParams{ - ID: je.job.ID, - CompletedAt: sql.NullTime{ - Time: now, - Valid: true, - }, - UpdatedAt: now, - Error: sql.NullString{ - String: "connection to provisioner daemon broken", - Valid: true, - }, - ErrorCode: sql.NullString{}, - }) - if err != nil { - logger.Error(streamCtx, "error updating failed job", slog.Error(err)) - } + s.failAcquiredJob(logger, je.job.ID, "connection to provisioner daemon broken") + s.jobFinished(je.job.ID) return recvErr } pj, err := s.acquireProtoJob(streamCtx, je.job) if err != nil { + // acquireProtoJob marks the job failed itself. + s.jobFinished(je.job.ID) return err } err = stream.Send(pj) if err != nil { s.Logger.Error(streamCtx, "failed to send job", slog.Error(err)) + // The job was locked but never delivered, so mark it failed instead of + // leaving it assigned to a worker that does not have it. + s.failAcquiredJob(logger, je.job.ID, "connection to provisioner daemon broken") + s.jobFinished(je.job.ID) return err } return nil } -func (s *server) acquireProtoJob(ctx context.Context, job database.ProvisionerJob) (*proto.AcquiredJob, error) { - // Marks the acquired job as failed with the error message provided. - failJob := func(errorMessage string) error { - err := s.Database.UpdateProvisionerJobWithCompleteByID(ctx, database.UpdateProvisionerJobWithCompleteByIDParams{ - ID: job.ID, +// failAcquiredJob marks a job that was claimed but never delivered to the +// daemon as failed. It uses a fresh context so the update succeeds even when +// the session context is canceled. +func (s *server) failAcquiredJob(logger slog.Logger, jobID uuid.UUID, message string) { + now := s.timeNow() + err := s.Database.UpdateProvisionerJobWithCompleteByID( + //nolint:gocritic // Provisionerd has specific authz rules. + dbauthz.AsProvisionerd(context.Background()), + database.UpdateProvisionerJobWithCompleteByIDParams{ + ID: jobID, CompletedAt: sql.NullTime{ - Time: s.timeNow(), + Time: now, Valid: true, }, + UpdatedAt: now, Error: sql.NullString{ - String: errorMessage, + String: message, Valid: true, }, - ErrorCode: job.ErrorCode, - UpdatedAt: s.timeNow(), + ErrorCode: sql.NullString{}, }) + if err != nil { + logger.Error(s.lifecycleCtx, "error updating failed job", slog.Error(err)) + } +} + +func (s *server) acquireProtoJob(ctx context.Context, job database.ProvisionerJob) (*proto.AcquiredJob, error) { + // Marks the acquired job as failed with the error message provided. The + // update runs on a fresh context so it succeeds even when the session + // context is canceled; otherwise the claimed job would stay assigned to + // this worker until the job reaper. + failJob := func(errorMessage string) error { + err := s.Database.UpdateProvisionerJobWithCompleteByID( + //nolint:gocritic // Provisionerd has specific authz rules. + dbauthz.AsProvisionerd(context.Background()), + database.UpdateProvisionerJobWithCompleteByIDParams{ + ID: job.ID, + CompletedAt: sql.NullTime{ + Time: s.timeNow(), + Valid: true, + }, + Error: sql.NullString{ + String: errorMessage, + Valid: true, + }, + ErrorCode: job.ErrorCode, + UpdatedAt: s.timeNow(), + }) if err != nil { return xerrors.Errorf("update provisioner job: %w", err) } @@ -1385,6 +1549,7 @@ func (s *server) FailJob(ctx context.Context, failJob *proto.FailedJob) (*proto. s.Logger.Error(ctx, "failed to publish end of job logs", slog.F("job_id", jobID), slog.Error(err)) return nil, xerrors.Errorf("publish end of job logs: %w", err) } + s.jobFinished(jobID) return &proto.Empty{}, nil } @@ -1702,6 +1867,7 @@ func (s *server) CompleteJob(ctx context.Context, completed *proto.CompletedJob) } s.Logger.Debug(ctx, "stage CompleteJob done", slog.F("job_id", jobID)) + s.jobFinished(jobID) return &proto.Empty{}, nil } diff --git a/coderd/provisionerdserver/provisionerdserver_internal_test.go b/coderd/provisionerdserver/provisionerdserver_internal_test.go index 7e6aa80f9b..185f9790a4 100644 --- a/coderd/provisionerdserver/provisionerdserver_internal_test.go +++ b/coderd/provisionerdserver/provisionerdserver_internal_test.go @@ -9,10 +9,12 @@ import ( "github.com/stretchr/testify/require" "golang.org/x/oauth2" + "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/database" "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/codersdk" "github.com/coder/coder/v2/testutil" ) @@ -167,3 +169,88 @@ func TestObtainOIDCAccessToken(t *testing.T) { require.Equal(t, "token", link.OAuthAccessToken) }) } + +// TestNewServer_SessionCancelRequired verifies that constructing a server for +// a deletable provisioner key without a SessionCancel fails, while reserved +// keys do not require one. +func TestNewServer_SessionCancelRequired(t *testing.T) { + t.Parallel() + + // The SessionCancel validation runs before the remaining nil-pointer + // checks, so the other arguments can be zero values. + newServer := func(keyID uuid.UUID) error { + _, err := NewServer( + context.Background(), "", nil, uuid.Nil, uuid.Nil, slog.Logger{}, + nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + Options{KeyID: keyID}, + nil, nil, nil, codersdk.Experiments{}, + ) + return err + } + + require.ErrorContains(t, newServer(uuid.New()), "SessionCancel is required") + // A reserved key passes the SessionCancel check; the error comes from the + // next validation instead. + require.ErrorContains(t, newServer(codersdk.ProvisionerKeyUUIDPSK), "quotaCommitter is nil") +} + +// TestTerminateSession_Deferral verifies that session cancellation is +// immediate when no job is active and deferred until the last active job +// finishes otherwise. +func TestTerminateSession_Deferral(t *testing.T) { + t.Parallel() + + newTestServer := func(canceled chan struct{}) *server { + return &server{ + lifecycleCtx: context.Background(), + Logger: testutil.Logger(t), + sessionCancel: func() { close(canceled) }, + activeJobs: map[uuid.UUID]struct{}{}, + } + } + assertCanceled := func(t *testing.T, canceled chan struct{}, want bool) { + t.Helper() + select { + case <-canceled: + require.True(t, want, "session canceled unexpectedly") + default: + require.False(t, want, "expected session to be canceled") + } + } + + t.Run("ImmediateWhenIdle", func(t *testing.T) { + t.Parallel() + canceled := make(chan struct{}) + s := newTestServer(canceled) + s.TerminateSession() + assertCanceled(t, canceled, true) + }) + + t.Run("DeferredUntilJobsFinish", func(t *testing.T) { + t.Parallel() + canceled := make(chan struct{}) + s := newTestServer(canceled) + job1, job2 := uuid.New(), uuid.New() + s.jobStarted(job1) + s.jobStarted(job2) + + s.TerminateSession() + assertCanceled(t, canceled, false) + + s.jobFinished(job1) + assertCanceled(t, canceled, false) + + s.jobFinished(job2) + assertCanceled(t, canceled, true) + }) + + t.Run("NoPendingTerminationNoCancel", func(t *testing.T) { + t.Parallel() + canceled := make(chan struct{}) + s := newTestServer(canceled) + jobID := uuid.New() + s.jobStarted(jobID) + s.jobFinished(jobID) + assertCanceled(t, canceled, false) + }) +} diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index 4713dbe399..c5eb376c8d 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -288,6 +288,77 @@ func TestAcquireJobWithCancel_Cancel(t *testing.T) { require.Equal(t, "", job.JobId) } +// TestAcquireJob_ProvisionerKeyDeleted verifies that acquiring a job fails and +// the session is canceled once the provisioner key the daemon authenticated +// with is deleted. +func TestAcquireJob_ProvisionerKeyDeleted(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + acquire func(context.Context, proto.DRPCProvisionerDaemonServer) error + }{ + {name: "Deprecated", acquire: func(ctx context.Context, srv proto.DRPCProvisionerDaemonServer) error { + _, err := srv.AcquireJob(ctx, nil) + return err + }}, + {name: "WithCancel", acquire: func(ctx context.Context, srv proto.DRPCProvisionerDaemonServer) error { + fs := newFakeStream(ctx) + errCh := make(chan error, 1) + go func() { errCh <- srv.AcquireJobWithCancel(fs) }() + // Cancel so the present-key acquire returns an empty job promptly; on + // the deleted-key path the key check returns before this is read. + fs.cancel() + return <-errCh + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + + // setup ties the daemon to a deletable key with this ID; sessionCancel + // records the teardown. A short poll keeps the present-key acquire from + // blocking. + keyID := uuid.New() + sessionCanceled := make(chan struct{}) + srv, srvDB, _, _ := setup(t, false, &overrides{ + keyID: keyID, + acquireJobLongPollDuration: testutil.IntervalFast, + sessionCancel: sync.OnceFunc(func() { close(sessionCanceled) }), + }) + + // While the key exists, acquiring returns without error. + require.NoError(t, tc.acquire(ctx, srv)) + + // Once the key is deleted, acquiring must fail and the session must be + // canceled rather than left polling a deleted key. + err := srvDB.DeleteProvisionerKey(dbauthz.AsProvisionerd(ctx), keyID) + require.NoError(t, err) + + err = tc.acquire(ctx, srv) + require.ErrorIs(t, err, provisionerdserver.ErrProvisionerKeyDeleted) + testutil.TryReceive(ctx, t, sessionCanceled) + }) + } +} + +// TestAcquireJob_ReservedProvisionerKey verifies that daemons using a reserved +// provisioner key, which cannot be deleted, are not blocked by the key check. +func TestAcquireJob_ReservedProvisionerKey(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + + //nolint:dogsled + srv, _, _, _ := setup(t, false, &overrides{ + keyID: codersdk.ProvisionerKeyUUIDPSK, + acquireJobLongPollDuration: testutil.IntervalFast, + }) + job, err := srv.AcquireJob(ctx, nil) + require.NoError(t, err) + require.Equal(t, &proto.AcquiredJob{}, job) +} + func TestHeartbeat(t *testing.T) { t.Parallel() @@ -316,6 +387,36 @@ func TestHeartbeat(t *testing.T) { // goleak.VerifyTestMain ensures that the heartbeat goroutine does not leak } +// TestHeartbeat_ProvisionerKeyDeleted verifies that the heartbeat loop cancels +// the session once the daemon's deletable key no longer exists. +func TestHeartbeat_ProvisionerKeyDeleted(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + + keyID := uuid.New() + sessionCanceled := make(chan struct{}) + //nolint:dogsled + _, db, _, _ := setup(t, false, &overrides{ + keyID: keyID, + heartbeatInterval: testutil.IntervalFast, + sessionCancel: sync.OnceFunc(func() { close(sessionCanceled) }), + }) + + // While the key exists, heartbeats must not cancel the session. + select { + case <-sessionCanceled: + t.Fatal("session canceled while key exists") + default: + } + + err := db.DeleteProvisionerKey(dbauthz.AsProvisionerd(ctx), keyID) + require.NoError(t, err) + + // A subsequent heartbeat tick must observe the deletion and cancel the + // session. + testutil.TryReceive(ctx, t, sessionCanceled) +} + func TestAcquireJob(t *testing.T) { t.Parallel() @@ -5178,6 +5279,8 @@ type overrides struct { notificationEnqueuer notifications.Enqueuer prebuildsOrchestrator agplprebuilds.ReconciliationOrchestrator provisionerdLogger *slog.Logger + keyID uuid.UUID + sessionCancel context.CancelFunc } func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisionerDaemonServer, database.Store, pubsub.Pubsub, database.ProvisionerDaemon) { @@ -5259,6 +5362,19 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi provisionerdLogger = *ov.provisionerdLogger } + keyID := codersdk.ProvisionerKeyUUIDBuiltIn + if ov.keyID != uuid.Nil { + keyID = ov.keyID + // The daemon's key_id is a foreign key to provisioner_keys, so a + // non-reserved key must exist before the daemon is created. + if !codersdk.IsReservedProvisionerKey(keyID) { + dbgen.ProvisionerKey(t, db, database.ProvisionerKey{ + ID: keyID, + OrganizationID: defOrg.ID, + }) + } + } + daemon, err := db.UpsertProvisionerDaemon(ov.ctx, database.UpsertProvisionerDaemonParams{ Name: "test", CreatedAt: dbtime.Now(), @@ -5268,7 +5384,7 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi Version: buildinfo.Version(), APIVersion: proto.CurrentVersion.String(), OrganizationID: defOrg.ID, - KeyID: codersdk.ProvisionerKeyUUIDBuiltIn, + KeyID: keyID, }) require.NoError(t, err) @@ -5316,6 +5432,8 @@ func setup(t *testing.T, ignoreLogErrors bool, ov *overrides) (proto.DRPCProvisi AcquireJobLongPollDur: pollDur, HeartbeatInterval: ov.heartbeatInterval, HeartbeatFn: ov.heartbeatFn, + KeyID: keyID, + SessionCancel: ov.sessionCancel, }, notifEnq, &op, diff --git a/coderd/pubsub/provisionerkeydeleted.go b/coderd/pubsub/provisionerkeydeleted.go new file mode 100644 index 0000000000..e04346e4d7 --- /dev/null +++ b/coderd/pubsub/provisionerkeydeleted.go @@ -0,0 +1,9 @@ +package pubsub + +import "github.com/google/uuid" + +// ProvisionerKeyDeletedChannel returns the pubsub channel that carries a +// notification when the provisioner key with the given ID is deleted. +func ProvisionerKeyDeletedChannel(keyID uuid.UUID) string { + return "provisioner_key_deleted:" + keyID.String() +} diff --git a/codersdk/provisionerdaemons.go b/codersdk/provisionerdaemons.go index 9dbd54abe2..df0eb990a4 100644 --- a/codersdk/provisionerdaemons.go +++ b/codersdk/provisionerdaemons.go @@ -408,6 +408,25 @@ func ReservedProvisionerKeyNames() []string { } } +// IsReservedProvisionerKey reports whether the given ID is one of the reserved +// provisioner keys (built-in, user-auth, PSK). Reserved keys are created by the +// system and cannot be deleted. +func IsReservedProvisionerKey(id uuid.UUID) bool { + switch id { + case ProvisionerKeyUUIDBuiltIn, ProvisionerKeyUUIDUserAuth, ProvisionerKeyUUIDPSK: + return true + default: + return false + } +} + +// IsDeletableProvisionerKey reports whether the given ID identifies a +// provisioner key that can be deleted. The zero value and reserved keys cannot +// be deleted. +func IsDeletableProvisionerKey(id uuid.UUID) bool { + return id != uuid.Nil && !IsReservedProvisionerKey(id) +} + type CreateProvisionerKeyRequest struct { Name string `json:"name"` Tags map[string]string `json:"tags"` diff --git a/enterprise/coderd/provisionerdaemons.go b/enterprise/coderd/provisionerdaemons.go index c2ba568b96..b6d0658433 100644 --- a/enterprise/coderd/provisionerdaemons.go +++ b/enterprise/coderd/provisionerdaemons.go @@ -21,10 +21,12 @@ import ( "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/coderd/database/dbtime" + dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub" "github.com/coder/coder/v2/coderd/httpapi" "github.com/coder/coder/v2/coderd/httpmw" "github.com/coder/coder/v2/coderd/httpmw/loggermw" "github.com/coder/coder/v2/coderd/provisionerdserver" + "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/rbac/policy" "github.com/coder/coder/v2/coderd/telemetry" @@ -356,6 +358,8 @@ func (api *API) provisionerDaemonServe(rw http.ResponseWriter, r *http.Request) OIDCConfig: api.OIDCConfig, AISeatTracker: api.AGPL.AISeatTracker, Clock: api.Clock, + KeyID: authRes.keyID, + SessionCancel: srvCancel, }, api.NotificationsEnqueuer, &api.AGPL.PrebuildsReconciler, @@ -389,8 +393,58 @@ func (api *API) provisionerDaemonServe(rw http.ResponseWriter, r *http.Request) rl.WriteLog(ctx, http.StatusAccepted) } - err = server.Serve(ctx, session) - srvCancel() + if codersdk.IsDeletableProvisionerKey(authRes.keyID) { + keyDeleted := func(ctx context.Context) (deleted bool, err error) { + _, err = api.Database.GetProvisionerKeyByID(ctx, authRes.keyID) + if xerrors.Is(err, sql.ErrNoRows) { + return true, nil + } + return false, err + } + + closeSubscribe, err := api.Pubsub.SubscribeWithErr( + pubsub.ProvisionerKeyDeletedChannel(authRes.keyID), + func(_ context.Context, _ []byte, subErr error) { + // ErrDroppedMessages means the Postgres listener reconnected; a + // deletion published during the outage may not have been + // delivered, so query the key directly instead of relying on the + // notification. + if xerrors.Is(subErr, dbpubsub.ErrDroppedMessages) { + deleted, err := keyDeleted(authCtx) + if err != nil { + logger.Warn(ctx, "failed to re-check provisioner key after dropped messages", + slog.F("provisioner_key_id", authRes.keyID), slog.Error(err)) + return + } + if !deleted { + return + } + } + logger.Info(ctx, "provisioner key deleted, terminating session", + slog.F("provisioner_key_id", authRes.keyID)) + srv.TerminateSession() + }, + ) + if err != nil { + _ = conn.Close(websocket.StatusInternalError, httpapi.WebsocketCloseSprintf("subscribe to provisioner key deletion: %s", err)) + return + } + defer closeSubscribe() + + // Postgres LISTEN/NOTIFY does not deliver notifications published before + // registration, so re-check after subscribing. + if deleted, err := keyDeleted(authCtx); err != nil { + _ = conn.Close(websocket.StatusInternalError, httpapi.WebsocketCloseSprintf("check provisioner key: %s", err)) + return + } else if deleted { + logger.Info(ctx, "provisioner key no longer exists, closing connection", + slog.F("provisioner_key_id", authRes.keyID)) + _ = conn.Close(websocket.StatusGoingAway, "provisioner key deleted") + return + } + } + + err = server.Serve(srvCtx, session) logger.Info(ctx, "provisioner daemon disconnected", slog.Error(err)) if err != nil && !xerrors.Is(err, io.EOF) { _ = conn.Close(websocket.StatusInternalError, httpapi.WebsocketCloseSprintf("serve: %s", err)) diff --git a/enterprise/coderd/provisionerdaemons_test.go b/enterprise/coderd/provisionerdaemons_test.go index 3d9347bcbf..a1f3f7f38e 100644 --- a/enterprise/coderd/provisionerdaemons_test.go +++ b/enterprise/coderd/provisionerdaemons_test.go @@ -7,12 +7,15 @@ import ( "fmt" "io" "net/http" + "sync" + "sync/atomic" "testing" "time" "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/xerrors" "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/apiversion" @@ -21,7 +24,10 @@ import ( "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/dbtestutil" + dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub" "github.com/coder/coder/v2/coderd/provisionerkey" + "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/util/ptr" "github.com/coder/coder/v2/codersdk" @@ -171,6 +177,360 @@ func TestProvisionerDaemonServe(t *testing.T) { require.Contains(t, string(b), fmt.Sprintf("server is at version %s, behind requested major version %s", proto.CurrentVersion.String(), v.String())) }) + t.Run("KeyDeletionClosesSession", func(t *testing.T) { + t.Parallel() + client, _ := coderdenttest.New(t, &coderdenttest.Options{LicenseOptions: &coderdenttest.LicenseOptions{ + Features: license.Features{ + codersdk.FeatureExternalProvisionerDaemons: 1, + codersdk.FeatureMultipleOrganizations: 1, + }, + }}) + org := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + orgAdmin, _ := coderdtest.CreateAnotherUser(t, client, org.ID, rbac.ScopedRoleOrgAdmin(org.ID)) + ctx := testutil.Context(t, testutil.WaitLong) + + res, err := orgAdmin.CreateProvisionerKey(ctx, org.ID, codersdk.CreateProvisionerKeyRequest{ + Name: "my-key", + }) + require.NoError(t, err) + + srv, err := orgAdmin.ServeProvisionerDaemon(ctx, codersdk.ServeProvisionerDaemonRequest{ + Name: testutil.MustRandString(t, 63), + Organization: org.ID, + Provisioners: []codersdk.ProvisionerType{ + codersdk.ProvisionerTypeEcho, + }, + Tags: map[string]string{}, + ProvisionerKey: res.Key, + }) + require.NoError(t, err) + defer srv.DRPCConn().Close() + + // The session is established and open. + select { + case <-srv.DRPCConn().Closed(): + t.Fatal("connection closed before key deletion") + default: + } + + // Deleting the key must tear down the active daemon session. + err = orgAdmin.DeleteProvisionerKey(ctx, org.ID, "my-key") + require.NoError(t, err) + + select { + case <-srv.DRPCConn().Closed(): + case <-ctx.Done(): + t.Fatal("timed out waiting for provisioner daemon session to close") + } + }) + + t.Run("KeyDeletedDuringSetupClosesSession", func(t *testing.T) { + t.Parallel() + // Provisioner key auth fetches the key by name, so the only + // GetProvisionerKeyByID reads in the serve path are the post-subscribe + // re-check and the heartbeat watchdog, whose first beat fires + // immediately at session start. Deleting the key on the first such + // read reproduces a key deleted between authentication and + // subscription; either reader must close the session, so this test + // pins the setup-window behavior rather than the re-check in + // isolation. No pubsub notification is delivered. + db, ps := dbtestutil.NewDB(t) + store := &deleteKeyOnReadStore{Store: db} + client, _ := coderdenttest.New(t, &coderdenttest.Options{ + Options: &coderdtest.Options{ + Database: store, + Pubsub: ps, + }, + LicenseOptions: &coderdenttest.LicenseOptions{ + Features: license.Features{ + codersdk.FeatureExternalProvisionerDaemons: 1, + codersdk.FeatureMultipleOrganizations: 1, + }, + }, + }) + org := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + orgAdmin, _ := coderdtest.CreateAnotherUser(t, client, org.ID, rbac.ScopedRoleOrgAdmin(org.ID)) + ctx := testutil.Context(t, testutil.WaitLong) + + res, err := orgAdmin.CreateProvisionerKey(ctx, org.ID, codersdk.CreateProvisionerKeyRequest{ + Name: "my-key", + }) + require.NoError(t, err) + keys, err := orgAdmin.ListProvisionerKeys(ctx, org.ID) + require.NoError(t, err) + require.Len(t, keys, 1) + keyID := keys[0].ID + store.keyID.Store(&keyID) + + srv, err := orgAdmin.ServeProvisionerDaemon(ctx, codersdk.ServeProvisionerDaemonRequest{ + Name: testutil.MustRandString(t, 63), + Organization: org.ID, + Provisioners: []codersdk.ProvisionerType{ + codersdk.ProvisionerTypeEcho, + }, + Tags: map[string]string{}, + ProvisionerKey: res.Key, + }) + require.NoError(t, err) + defer srv.DRPCConn().Close() + + select { + case <-srv.DRPCConn().Closed(): + case <-ctx.Done(): + t.Fatal("timed out waiting for re-check to close the session") + } + // Confirm the close was driven by a key lookup, not another path (no + // pubsub notification is published in this test). + require.True(t, store.deleted.Load()) + }) + + t.Run("DroppedMessageClosesSession", func(t *testing.T) { + t.Parallel() + // A dropped-messages signal is delivered when the Postgres listener + // reconnects. If the key was deleted while the listener was down, the + // deletion notification is never delivered, so the serve handler must + // re-check the key on the dropped-messages signal. This test deletes + // the key directly in the database (no notification published) and then + // drives the captured listener with ErrDroppedMessages. + // + // The heartbeat watchdog also detects a deleted key, but its first + // beat fires at session start while the key still exists and later + // beats (1m default) exceed the test deadline (testutil.WaitLong), so + // a close after the signal is attributable to the dropped-messages + // re-check. If either timing changes, the heartbeat could close the + // session instead and mask removal of the re-check. + db, ps := dbtestutil.NewDB(t) + capturePS := newCaptureKeyDeletePubsub(ps) + client, _ := coderdenttest.New(t, &coderdenttest.Options{ + Options: &coderdtest.Options{ + Database: db, + Pubsub: capturePS, + // The wrapper is not a *PGPubsub, so provide the real one for + // replica sync. + ReplicaSyncPubsub: ps.(*dbpubsub.PGPubsub), + }, + LicenseOptions: &coderdenttest.LicenseOptions{ + Features: license.Features{ + codersdk.FeatureExternalProvisionerDaemons: 1, + codersdk.FeatureMultipleOrganizations: 1, + }, + }, + }) + org := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + orgAdmin, _ := coderdtest.CreateAnotherUser(t, client, org.ID, rbac.ScopedRoleOrgAdmin(org.ID)) + ctx := testutil.Context(t, testutil.WaitLong) + + res, err := orgAdmin.CreateProvisionerKey(ctx, org.ID, codersdk.CreateProvisionerKeyRequest{ + Name: "my-key", + }) + require.NoError(t, err) + keys, err := orgAdmin.ListProvisionerKeys(ctx, org.ID) + require.NoError(t, err) + require.Len(t, keys, 1) + keyID := keys[0].ID + capturePS.target.Store(ptr.Ref(pubsub.ProvisionerKeyDeletedChannel(keyID))) + + srv, err := orgAdmin.ServeProvisionerDaemon(ctx, codersdk.ServeProvisionerDaemonRequest{ + Name: testutil.MustRandString(t, 63), + Organization: org.ID, + Provisioners: []codersdk.ProvisionerType{ + codersdk.ProvisionerTypeEcho, + }, + Tags: map[string]string{}, + ProvisionerKey: res.Key, + }) + require.NoError(t, err) + defer srv.DRPCConn().Close() + + // Capture the listener the serve handler registered for this key. + listener := capturePS.waitListener(ctx, t) + + // The session is established and open. + select { + case <-srv.DRPCConn().Closed(): + t.Fatal("connection closed before key deletion") + default: + } + + // Delete the key without publishing, simulating a deletion missed while + // the listener was down. + //nolint:gocritic // The test deletes the key outside the request actor. + err = db.DeleteProvisionerKey(dbauthz.AsSystemRestricted(ctx), keyID) + require.NoError(t, err) + + // Deliver the dropped-messages signal the reconnect would have produced. + listener(ctx, nil, dbpubsub.ErrDroppedMessages) + + select { + case <-srv.DRPCConn().Closed(): + case <-ctx.Done(): + t.Fatal("timed out waiting for dropped-message re-check to close the session") + } + }) + + t.Run("DroppedMessageKeyExistsKeepsSession", func(t *testing.T) { + t.Parallel() + // A dropped-messages signal with the key still present must leave the + // session running; only a confirmed missing key may terminate it. + db, ps := dbtestutil.NewDB(t) + capturePS := newCaptureKeyDeletePubsub(ps) + client, _ := coderdenttest.New(t, &coderdenttest.Options{ + Options: &coderdtest.Options{ + Database: db, + Pubsub: capturePS, + // The wrapper is not a *PGPubsub, so provide the real one for + // replica sync. + ReplicaSyncPubsub: ps.(*dbpubsub.PGPubsub), + }, + LicenseOptions: &coderdenttest.LicenseOptions{ + Features: license.Features{ + codersdk.FeatureExternalProvisionerDaemons: 1, + codersdk.FeatureMultipleOrganizations: 1, + }, + }, + }) + org := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + orgAdmin, _ := coderdtest.CreateAnotherUser(t, client, org.ID, rbac.ScopedRoleOrgAdmin(org.ID)) + ctx := testutil.Context(t, testutil.WaitLong) + + res, err := orgAdmin.CreateProvisionerKey(ctx, org.ID, codersdk.CreateProvisionerKeyRequest{ + Name: "my-key", + }) + require.NoError(t, err) + keys, err := orgAdmin.ListProvisionerKeys(ctx, org.ID) + require.NoError(t, err) + require.Len(t, keys, 1) + keyID := keys[0].ID + capturePS.target.Store(ptr.Ref(pubsub.ProvisionerKeyDeletedChannel(keyID))) + + srv, err := orgAdmin.ServeProvisionerDaemon(ctx, codersdk.ServeProvisionerDaemonRequest{ + Name: testutil.MustRandString(t, 63), + Organization: org.ID, + Provisioners: []codersdk.ProvisionerType{ + codersdk.ProvisionerTypeEcho, + }, + Tags: map[string]string{}, + ProvisionerKey: res.Key, + }) + require.NoError(t, err) + defer srv.DRPCConn().Close() + + listener := capturePS.waitListener(ctx, t) + + // Deliver the dropped-messages signal without deleting the key. The + // re-check finds the key and the session must stay open. + listener(ctx, nil, dbpubsub.ErrDroppedMessages) + select { + case <-srv.DRPCConn().Closed(): + t.Fatal("session closed although provisioner key still exists") + case <-time.After(testutil.IntervalMedium): + } + + // The same signal after the key is gone closes the session, confirming + // the session was still fully functional above. + //nolint:gocritic // The test deletes the key outside the request actor. + err = db.DeleteProvisionerKey(dbauthz.AsSystemRestricted(ctx), keyID) + require.NoError(t, err) + listener(ctx, nil, dbpubsub.ErrDroppedMessages) + select { + case <-srv.DRPCConn().Closed(): + case <-ctx.Done(): + t.Fatal("timed out waiting for dropped-message re-check to close the session") + } + }) + + t.Run("DroppedMessageKeyCheckErrorKeepsSession", func(t *testing.T) { + t.Parallel() + // A dropped-messages signal whose key re-check fails must leave the + // session running; a transient database error is not evidence of + // deletion. + db, ps := dbtestutil.NewDB(t) + store := &failKeyReadStore{Store: db} + capturePS := newCaptureKeyDeletePubsub(ps) + client, _ := coderdenttest.New(t, &coderdenttest.Options{ + Options: &coderdtest.Options{ + Database: store, + Pubsub: capturePS, + // The wrapper is not a *PGPubsub, so provide the real one for + // replica sync. + ReplicaSyncPubsub: ps.(*dbpubsub.PGPubsub), + }, + LicenseOptions: &coderdenttest.LicenseOptions{ + Features: license.Features{ + codersdk.FeatureExternalProvisionerDaemons: 1, + codersdk.FeatureMultipleOrganizations: 1, + }, + }, + }) + org := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + orgAdmin, _ := coderdtest.CreateAnotherUser(t, client, org.ID, rbac.ScopedRoleOrgAdmin(org.ID)) + ctx := testutil.Context(t, testutil.WaitLong) + + res, err := orgAdmin.CreateProvisionerKey(ctx, org.ID, codersdk.CreateProvisionerKeyRequest{ + Name: "my-key", + }) + require.NoError(t, err) + keys, err := orgAdmin.ListProvisionerKeys(ctx, org.ID) + require.NoError(t, err) + require.Len(t, keys, 1) + keyID := keys[0].ID + store.keyID.Store(&keyID) + capturePS.target.Store(ptr.Ref(pubsub.ProvisionerKeyDeletedChannel(keyID))) + + srv, err := orgAdmin.ServeProvisionerDaemon(ctx, codersdk.ServeProvisionerDaemonRequest{ + Name: testutil.MustRandString(t, 63), + Organization: org.ID, + Provisioners: []codersdk.ProvisionerType{ + codersdk.ProvisionerTypeEcho, + }, + Tags: map[string]string{}, + ProvisionerKey: res.Key, + }) + require.NoError(t, err) + defer srv.DRPCConn().Close() + + listener := capturePS.waitListener(ctx, t) + + // A completed RPC proves the handler finished connection setup, which + // includes the post-subscribe key re-check. That re-check closes the + // connection both when the key is gone and when the read errors, so it + // has to run before the key is touched below. The error must come from + // the job lookup: a transport error would mean the session is already + // gone, which proves nothing about setup. + _, err = srv.UpdateJob(ctx, &proto.UpdateJobRequest{JobId: uuid.NewString()}) + require.ErrorContains(t, err, "get job") + + // Make the key re-check fail before deleting the key, so no reader + // (including the heartbeat key check) can observe the bare deletion; + // then delete the key so a successful re-check would terminate. The + // error path must leave the session running. The write gate holds off + // key reads for both steps, so none is in flight across the deletion. + func() { + store.gate.Lock() + defer store.gate.Unlock() + store.fail.Store(true) + //nolint:gocritic // The test deletes the key outside the request actor. + err = db.DeleteProvisionerKey(dbauthz.AsSystemRestricted(ctx), keyID) + }() + require.NoError(t, err) + listener(ctx, nil, dbpubsub.ErrDroppedMessages) + select { + case <-srv.DRPCConn().Closed(): + t.Fatal("session closed although the key re-check failed") + case <-time.After(testutil.IntervalMedium): + } + + // Once the re-check succeeds again it observes the deletion and closes + // the session. + store.fail.Store(false) + listener(ctx, nil, dbpubsub.ErrDroppedMessages) + select { + case <-srv.DRPCConn().Closed(): + case <-ctx.Done(): + t.Fatal("timed out waiting for dropped-message re-check to close the session") + } + }) + t.Run("NoLicense", func(t *testing.T) { t.Parallel() client, user := coderdenttest.New(t, &coderdenttest.Options{DontAddLicense: true}) @@ -986,3 +1346,96 @@ func TestGetProvisionerDaemons(t *testing.T) { } }) } + +// failKeyReadStore returns an error from GetProvisionerKeyByID for the key +// identified by keyID while fail is set. All other reads pass through. keyID +// is set after the key is created so earlier lookups are unaffected. +// +// Reads of that key hold gate for reading while they check fail and query the +// database. A caller that takes gate for writing therefore runs with no such +// read in flight, and any read that arrives afterwards observes the writer's +// changes to fail and to the key itself. gate is not reentrant: a writer must +// not read through this store while holding it. +type failKeyReadStore struct { + database.Store + keyID atomic.Pointer[uuid.UUID] + fail atomic.Bool + gate sync.RWMutex +} + +func (s *failKeyReadStore) GetProvisionerKeyByID(ctx context.Context, id uuid.UUID) (database.ProvisionerKey, error) { + if target := s.keyID.Load(); target != nil && *target == id { + s.gate.RLock() + defer s.gate.RUnlock() + if s.fail.Load() { + return database.ProvisionerKey{}, xerrors.New("transient database error") + } + } + return s.Store.GetProvisionerKeyByID(ctx, id) +} + +// deleteKeyOnReadStore deletes the provisioner key identified by keyID the +// first time it is fetched by ID, simulating a key deleted during connection +// setup. keyID is set after the key is created so earlier lookups are +// unaffected. deleted records that the interception fired. +type deleteKeyOnReadStore struct { + database.Store + keyID atomic.Pointer[uuid.UUID] + once sync.Once + deleted atomic.Bool +} + +func (s *deleteKeyOnReadStore) GetProvisionerKeyByID(ctx context.Context, id uuid.UUID) (database.ProvisionerKey, error) { + if target := s.keyID.Load(); target != nil && *target == id { + s.once.Do(func() { + //nolint:gocritic // The test deletes the key outside the request actor. + _ = s.Store.DeleteProvisionerKey(dbauthz.AsSystemRestricted(context.Background()), id) + s.deleted.Store(true) + }) + } + return s.Store.GetProvisionerKeyByID(ctx, id) +} + +// captureKeyDeletePubsub records the ListenerWithErr registered for the channel +// named by target so a test can invoke it directly, e.g. with +// ErrDroppedMessages. All other pubsub operations pass through to the embedded +// Pubsub unchanged. target is set before the subscription is expected so the +// capturing subscribe observes it. +type captureKeyDeletePubsub struct { + dbpubsub.Pubsub + target atomic.Pointer[string] + mu sync.Mutex + listener dbpubsub.ListenerWithErr + once sync.Once + got chan struct{} +} + +func newCaptureKeyDeletePubsub(ps dbpubsub.Pubsub) *captureKeyDeletePubsub { + return &captureKeyDeletePubsub{Pubsub: ps, got: make(chan struct{})} +} + +func (p *captureKeyDeletePubsub) SubscribeWithErr(event string, listener dbpubsub.ListenerWithErr) (func(), error) { + cancel, err := p.Pubsub.SubscribeWithErr(event, listener) + if err != nil { + return cancel, err + } + if target := p.target.Load(); target != nil && *target == event { + p.mu.Lock() + p.listener = listener + p.mu.Unlock() + p.once.Do(func() { close(p.got) }) + } + return cancel, nil +} + +func (p *captureKeyDeletePubsub) waitListener(ctx context.Context, t *testing.T) dbpubsub.ListenerWithErr { + t.Helper() + select { + case <-p.got: + case <-ctx.Done(): + t.Fatal("timed out waiting for provisioner key deletion subscription") + } + p.mu.Lock() + defer p.mu.Unlock() + return p.listener +} diff --git a/enterprise/coderd/provisionerkeys.go b/enterprise/coderd/provisionerkeys.go index 49640042d4..546b073585 100644 --- a/enterprise/coderd/provisionerkeys.go +++ b/enterprise/coderd/provisionerkeys.go @@ -7,12 +7,14 @@ import ( "strings" "time" + "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/db2sdk" "github.com/coder/coder/v2/coderd/httpapi" "github.com/coder/coder/v2/coderd/httpmw" "github.com/coder/coder/v2/coderd/provisionerdserver" "github.com/coder/coder/v2/coderd/provisionerkey" + "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/codersdk" ) @@ -196,9 +198,7 @@ func (api *API) deleteProvisionerKey(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() provisionerKey := httpmw.ProvisionerKeyParam(r) - if provisionerKey.ID.String() == codersdk.ProvisionerKeyIDBuiltIn || - provisionerKey.ID.String() == codersdk.ProvisionerKeyIDUserAuth || - provisionerKey.ID.String() == codersdk.ProvisionerKeyIDPSK { + if codersdk.IsReservedProvisionerKey(provisionerKey.ID) { httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ Message: fmt.Sprintf("Cannot delete reserved '%s' provisioner key", provisionerKey.Name), }) @@ -211,6 +211,13 @@ func (api *API) deleteProvisionerKey(rw http.ResponseWriter, r *http.Request) { return } + // Notify subscribers that this key was deleted so active sessions tear down. + // Publishing is best effort; a failure does not leave the key usable. + if err := api.Pubsub.Publish(pubsub.ProvisionerKeyDeletedChannel(provisionerKey.ID), nil); err != nil { + api.Logger.Warn(ctx, "failed to publish provisioner key deletion", + slog.F("provisioner_key_id", provisionerKey.ID), slog.Error(err)) + } + httpapi.Write(ctx, rw, http.StatusNoContent, nil) }