mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: invalidate provisioner daemon sessions on key deletion (#26532)
## Summary Closes PLAT-305. When a provisioner key is deleted, the associated daemon kept operating on its existing WebSocket connection, because authentication was only checked at connection establishment and deletion was a bare `DELETE` with no session invalidation. This adds four layers of defense so a deleted key promptly stops doing work: 1. **Publish on delete.** `deleteProvisionerKey` publishes to a new per-key pubsub channel (`coderd/pubsub.ProvisionerKeyDeletedChannel`) after a successful delete. Publish errors are logged but still return `204`, since layer 3 is the durable backstop. 2. **Subscribe and tear down.** The daemon serve handler subscribes to its key's channel and terminates the DRPC session on a deletion event. Termination is deferred while a job claimed by the session is active: the daemon may finish and report the in-flight job (`UpdateJob`/`CompleteJob` have no key check), and the last active job's completion performs the cancellation. Because Postgres `LISTEN`/`NOTIFY` does not buffer for non-listeners, the handler also performs a synchronous key-existence re-check immediately after subscribing to close the race between auth and subscription. The subscription uses `SubscribeWithErr` so that an `ErrDroppedMessages` signal (emitted when the pubsub listener reconnects) triggers the same key re-check, closing the listener-outage window in which a deletion notification could be missed. 3. **Backstop on acquire.** `AcquireJob` and `AcquireJobWithCancel` verify the key still exists before waiting for a job, and the `Acquirer` claims jobs in a transaction that first locks the worker's deletable key (`LockProvisionerKeyByIDForShare`, a `FOR KEY SHARE` row lock held until commit) before running the `AcquireProvisionerJob` claim, so a claim cannot commit after the key's deletion. This guards against a missed pubsub message. A missing key row surfaces as its own result rather than overloading the claim query's no-rows response: the acquire terminates with `ErrProvisionerKeyDeleted` (terminating the session, with the same active-job deferral) and hands the consumed wakeup to another waiting daemon in the same domain, rather than silently re-parking and starving peers of job postings. 4. **Heartbeat watchdog.** The per-session heartbeat loop (1m interval) also re-checks the key, so even a session whose deletion notification was silently lost terminates within one heartbeat interval instead of living until the connection breaks (same active-job deferral as layer 2). Reserved keys skip the check. A job that is claimed but never delivered (the session or connection dies between the database claim and the stream send) is marked failed immediately on a fresh context, instead of staying assigned to the worker until the job reaper. Reserved keys (built-in, user-auth, PSK) are exempt throughout, since they are not deletable rows. The acquire-time lookup runs as `dbauthz.AsSystemReadProvisionerDaemons`, because the provisionerd role cannot read provisioner keys and a provisioner key's RBAC object is a provisioner daemon. A single key can back many daemons (and span HA replicas), so the per-key channel fans out to invalidate all of them at once. Per-key channels keep the `LISTEN` count proportional to distinct keys rather than waking every daemon on unrelated deletions. ### Known limitations - **`UpdateJob`/`CompleteJob` intentionally have no key check.** By the time those RPCs arrive the work has already run; rejecting completion would strand a build in "running" (until the job reaper fails it) with real infrastructure left unreconciled. Session termination is deferred while a job is active so the completion can be reported; the daemon may not receive the final RPC response when the deferred termination fires, but the job's outcome is already persisted. - **After termination, the daemon process redials and receives 401s until restarted.** The dial-time exit logic only triggers on 403, and the auth middleware returns 401 for an invalid key; this dial behavior predates this PR and is tracked as a follow-up in [PLAT-452](https://linear.app/codercom/issue/PLAT-452) (return 403 for invalid provisioner keys). ## Tests - `coderd/provisionerdserver`: `TestAcquireJob_ProvisionerKeyDeleted` (both RPC variants), `TestAcquireJob_ReservedProvisionerKey`, `TestHeartbeat_ProvisionerKeyDeleted` (heartbeat watchdog cancels the session after key deletion), `TestAcquirer_ProvisionerKeyDeleted` (a dead-key acquiree exits terminally and its clearance is promoted to a peer in the same domain), and `TestTerminateSession_Deferral` (termination is immediate when idle and deferred until the last active job finishes). - `coderd/database`: `TestAcquireProvisionerJob/ProvisionerKeyLock` covers the lock query against real Postgres: it returns the key ID while the row exists and no rows once it is deleted. The lock-then-claim composition is pinned by `TestAcquirer_ProvisionerKeyDeleted`. - `enterprise/coderd`: `TestProvisionerDaemonServe/KeyDeletionClosesSession` asserts an active session closes after its key is deleted. `KeyDeletedDuringSetupClosesSession` covers the post-subscribe re-check when a key is deleted between auth and subscription, and `DroppedMessageClosesSession` covers the `ErrDroppedMessages` re-check when a deletion is missed during a listener outage. ## Validation - `make` pre-commit (gen/fmt/lint/build) passed via git hooks. - Targeted tests pass; existing acquire tests pass with no regression. - Manual: brought up a dev deployment (coder-in-coder) with a Premium license, created a deletable provisioner key, and started an external daemon with `coder provisionerd start`. Confirmed it authenticated via the key and connected, appearing as `idle` in both `coder provisioner list` (with the key name) and the organization Provisioners UI. - Manual, idle teardown: deleted the key while the daemon was idle. The server logged `provisioner key deleted, terminating session`, the daemon's session closed immediately, and it dropped from `coder provisioner list` (then entered the known 401 redial loop, PLAT-452). - Manual, deferred termination: ran a workspace build (tagged template, `sleep 45` in `local-exec`) pinned to the external daemon and deleted the key mid-build. The server logged `deferring session cancellation until active jobs finish`; the heartbeat watchdog re-checked mid-build and re-deferred rather than force-killing. The build ran to completion (`Apply complete`, workspace `Started`) and only then did `canceling session after job completion` fire. The documented caveat reproduced: the daemon lost the final `CompleteJob` ack, and the build outcome was still persisted correctly. <details> <summary>Implementation plan and design decisions</summary> ### Design - **Per-key vs global channel:** chose per-key (`provisioner_key_deleted:<keyID>`) so daemons do not wake on unrelated deletions. The cost is one `LISTEN` per distinct key per replica on the shared listener connection, which is negligible against Coder's existing channels. - **Missing-key behavior on acquire:** returns an error that tears down the acquire rather than silently returning an empty job. - **Subscribe-startup race:** ordering is `authorize -> UpsertProvisionerDaemon -> Subscribe -> GetProvisionerKeyByID`. The post-subscribe re-check handles a deletion that committed before the `LISTEN` registered (Postgres does not buffer notifications for non-listeners; the in-process buffer only smooths bursts and drops on overflow). - **`NewServer` change:** `KeyID` was added to `provisionerdserver.Options` to avoid a positional signature change across call sites. The in-memory (built-in) daemon leaves it unset and is therefore exempt. ### Files - `coderd/pubsub/provisionerkeydeleted.go` (new) — channel helper. - `enterprise/coderd/provisionerkeys.go` — publish on delete. - `enterprise/coderd/provisionerdaemons.go` — subscribe, re-check, cancel session; pass `KeyID`. - `coderd/provisionerdserver/provisionerdserver.go` — `KeyID` option and acquire-time existence check. </details> --- This pull request was created by Coder Agents on behalf of @jscottmiller.
This commit is contained in:
@@ -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())
|
||||
|
||||
|
||||
@@ -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")})
|
||||
|
||||
+8
@@ -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)
|
||||
|
||||
Generated
+15
@@ -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()
|
||||
|
||||
Generated
+5
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
Generated
+21
@@ -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
|
||||
|
||||
@@ -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
|
||||
*
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user