diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index d9e715ca1a1..f6009c5423b 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -2115,6 +2115,10 @@ func (q *querier) DeleteAPIKeyByID(ctx context.Context, id string) error { return deleteQ(q.log, q.auth, q.db.GetAPIKeyByID, q.db.DeleteAPIKeyByID)(ctx, id) } +func (q *querier) DeleteAPIKeyByIDReturningRow(ctx context.Context, id string) (database.APIKey, error) { + return fetchAndQuery(q.log, q.auth, policy.ActionDelete, q.db.GetAPIKeyByID, q.db.DeleteAPIKeyByIDReturningRow)(ctx, id) +} + func (q *querier) DeleteAPIKeysByUserID(ctx context.Context, userID uuid.UUID) error { // TODO: This is not 100% correct because it omits apikey IDs. err := q.authorizeContext(ctx, policy.ActionDelete, diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 4a7ba7e3a11..8551585cf50 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -360,6 +360,12 @@ func (s *MethodTestSuite) TestAPIKey() { dbm.EXPECT().DeleteAPIKeyByID(gomock.Any(), key.ID).Return(nil).AnyTimes() check.Args(key.ID).Asserts(key, policy.ActionDelete).Returns() })) + s.Run("DeleteAPIKeyByIDReturningRow", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + key := testutil.Fake(s.T(), faker, database.APIKey{}) + dbm.EXPECT().GetAPIKeyByID(gomock.Any(), key.ID).Return(key, nil).AnyTimes() + dbm.EXPECT().DeleteAPIKeyByIDReturningRow(gomock.Any(), key.ID).Return(key, nil).AnyTimes() + check.Args(key.ID).Asserts(key, policy.ActionDelete).Returns(key) + })) s.Run("DeleteExpiredAPIKeys", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { args := database.DeleteExpiredAPIKeysParams{ Before: time.Date(2025, 11, 21, 0, 0, 0, 0, time.UTC), diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 227536761b1..887c3b5ad3a 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -440,6 +440,14 @@ func (m queryMetricsStore) DeleteAPIKeyByID(ctx context.Context, id string) erro return r0 } +func (m queryMetricsStore) DeleteAPIKeyByIDReturningRow(ctx context.Context, id string) (database.APIKey, error) { + start := time.Now() + r0, r1 := m.s.DeleteAPIKeyByIDReturningRow(ctx, id) + m.queryLatencies.WithLabelValues("DeleteAPIKeyByIDReturningRow").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteAPIKeyByIDReturningRow").Inc() + return r0, r1 +} + func (m queryMetricsStore) DeleteAPIKeysByUserID(ctx context.Context, userID uuid.UUID) error { start := time.Now() r0 := m.s.DeleteAPIKeysByUserID(ctx, userID) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index ce99316869b..3c92f7d5f58 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -704,6 +704,21 @@ func (mr *MockStoreMockRecorder) DeleteAPIKeyByID(ctx, id any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAPIKeyByID", reflect.TypeOf((*MockStore)(nil).DeleteAPIKeyByID), ctx, id) } +// DeleteAPIKeyByIDReturningRow mocks base method. +func (m *MockStore) DeleteAPIKeyByIDReturningRow(ctx context.Context, id string) (database.APIKey, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteAPIKeyByIDReturningRow", ctx, id) + ret0, _ := ret[0].(database.APIKey) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// DeleteAPIKeyByIDReturningRow indicates an expected call of DeleteAPIKeyByIDReturningRow. +func (mr *MockStoreMockRecorder) DeleteAPIKeyByIDReturningRow(ctx, id any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAPIKeyByIDReturningRow", reflect.TypeOf((*MockStore)(nil).DeleteAPIKeyByIDReturningRow), ctx, id) +} + // DeleteAPIKeysByUserID mocks base method. func (m *MockStore) DeleteAPIKeysByUserID(ctx context.Context, userID uuid.UUID) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 7ee24d1c4de..57a452ed5e1 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -119,6 +119,9 @@ type sqlcQuerier interface { DeleteAIProviderByID(ctx context.Context, id uuid.UUID) error DeleteAIProviderKey(ctx context.Context, id uuid.UUID) error DeleteAPIKeyByID(ctx context.Context, id string) error + // Returns sql.ErrNoRows when the key is already gone, which lets a caller + // enforce single use by racing this delete instead of reading first. + DeleteAPIKeyByIDReturningRow(ctx context.Context, id string) (APIKey, error) DeleteAPIKeysByUserID(ctx context.Context, userID uuid.UUID) error // Deletes all heartbeat rows for the chat. Used during ownership // transitions that abandon a lease. diff --git a/coderd/database/querier_test.go b/coderd/database/querier_test.go index 69a9b1073a7..89ecf1e0b59 100644 --- a/coderd/database/querier_test.go +++ b/coderd/database/querier_test.go @@ -19088,6 +19088,22 @@ func TestSingleUseDeleteByIDReturningRow(t *testing.T) { _, err = db.DeleteOAuth2ProviderAppCodeByIDReturningRow(ctx, code.ID) require.ErrorIs(t, err, sql.ErrNoRows) }) + + t.Run("APIKey", func(t *testing.T) { + t.Parallel() + db, _ := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitLong) + + user := dbgen.User(t, db, database.User{}) + key, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) + + deleted, err := db.DeleteAPIKeyByIDReturningRow(ctx, key.ID) + require.NoError(t, err) + require.Equal(t, key, deleted) + + _, err = db.DeleteAPIKeyByIDReturningRow(ctx, key.ID) + require.ErrorIs(t, err, sql.ErrNoRows) + }) } func TestGetAIModelPriceByProviderModel(t *testing.T) { diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 6dc0aa217a3..02622d4d33f 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -3745,6 +3745,37 @@ func (q *sqlQuerier) DeleteAPIKeyByID(ctx context.Context, id string) error { return err } +const deleteAPIKeyByIDReturningRow = `-- name: DeleteAPIKeyByIDReturningRow :one +DELETE FROM + api_keys +WHERE + id = $1 +RETURNING id, hashed_secret, user_id, last_used, expires_at, created_at, updated_at, login_type, lifetime_seconds, ip_address, token_name, scopes, allow_list +` + +// Returns sql.ErrNoRows when the key is already gone, which lets a caller +// enforce single use by racing this delete instead of reading first. +func (q *sqlQuerier) DeleteAPIKeyByIDReturningRow(ctx context.Context, id string) (APIKey, error) { + row := q.db.QueryRowContext(ctx, deleteAPIKeyByIDReturningRow, id) + var i APIKey + err := row.Scan( + &i.ID, + &i.HashedSecret, + &i.UserID, + &i.LastUsed, + &i.ExpiresAt, + &i.CreatedAt, + &i.UpdatedAt, + &i.LoginType, + &i.LifetimeSeconds, + &i.IPAddress, + &i.TokenName, + &i.Scopes, + &i.AllowList, + ) + return i, err +} + const deleteAPIKeysByUserID = `-- name: DeleteAPIKeysByUserID :exec DELETE FROM api_keys diff --git a/coderd/database/queries/apikeys.sql b/coderd/database/queries/apikeys.sql index 90e7610cf06..df62cd7a66f 100644 --- a/coderd/database/queries/apikeys.sql +++ b/coderd/database/queries/apikeys.sql @@ -92,6 +92,15 @@ DELETE FROM WHERE id = $1; +-- name: DeleteAPIKeyByIDReturningRow :one +-- Returns sql.ErrNoRows when the key is already gone, which lets a caller +-- enforce single use by racing this delete instead of reading first. +DELETE FROM + api_keys +WHERE + id = $1 +RETURNING *; + -- name: DeleteApplicationConnectAPIKeysByUserID :exec DELETE FROM api_keys diff --git a/coderd/oauth2provider/tokens.go b/coderd/oauth2provider/tokens.go index aa221530ed0..1e4f1919f2b 100644 --- a/coderd/oauth2provider/tokens.go +++ b/coderd/oauth2provider/tokens.go @@ -674,7 +674,12 @@ func refreshTokenGrant(ctx context.Context, db database.Store, logger slog.Logge err = db.InTx(func(tx database.Store) error { ctx := dbauthz.As(ctx, actor) - err = tx.DeleteAPIKeyByID(ctx, prevKey.ID) // This cascades to the token. + // The delete decides the race: only the refresh that removes the key may + // mint a replacement, and the loser sees the token as already spent. + _, err = tx.DeleteAPIKeyByIDReturningRow(ctx, prevKey.ID) // This cascades to the token. + if errors.Is(err, sql.ErrNoRows) { + return errBadToken + } if err != nil { return xerrors.Errorf("delete oauth2 app token: %w", err) } diff --git a/coderd/oauth2provider/tokens_test.go b/coderd/oauth2provider/tokens_test.go index bf8a7c05cd1..752c574cd3b 100644 --- a/coderd/oauth2provider/tokens_test.go +++ b/coderd/oauth2provider/tokens_test.go @@ -295,25 +295,55 @@ func TestOAuth2TokenExchangeSingleUse(t *testing.T) { app := seedAppWithSecret(t, db, sql.NullString{String: scopeInCatalog, Valid: true}) code, verifier := authorizeCode(ctx, t, client, app.ID.String(), "workspace:ssh") - form := tokenExchangeForm(app, code, verifier) - type exchange struct { + requireExactlyOneMinted(ctx, t, client, tokenExchangeForm(app, code, verifier), + "a code may mint at most one token") +} + +// A refresh mints a replacement token and deletes the key the presented one +// hangs off, so the same race as the exchange applies: the deletion has to +// arbitrate, or both requests mint from one refresh token. +func TestOAuth2RefreshSingleUse(t *testing.T) { + t.Parallel() + + db, pubsub := dbtestutil.NewDB(t) + client := coderdtest.New(t, &coderdtest.Options{ + Database: db, + Pubsub: pubsub, + }) + coderdtest.CreateFirstUser(t, client) + ctx := testutil.Context(t, testutil.WaitLong) + + app := seedAppWithSecret(t, db, sql.NullString{String: scopeInCatalog, Valid: true}) + code, verifier := authorizeCode(ctx, t, client, app.ID.String(), "workspace:ssh") + token := exchangeCode(ctx, t, client, app, code, verifier) + + requireExactlyOneMinted(ctx, t, client, refreshForm(app, token.RefreshToken), + "a refresh token may mint at most one replacement") +} + +// requireExactlyOneMinted posts form twice concurrently and requires one 200 +// and one `invalid_grant`. +func requireExactlyOneMinted(ctx context.Context, t *testing.T, client *codersdk.Client, form url.Values, msg string) { + t.Helper() + + type attempt struct { status int body string } var barrier sync.WaitGroup barrier.Add(2) - redeem := func() exchange { + redeem := func() attempt { barrier.Done() barrier.Wait() status, body := postTokenRequest(ctx, t, client, form) - return exchange{status: status, body: body} + return attempt{status: status, body: body} } - other := make(chan exchange, 1) + other := make(chan attempt, 1) go func() { other <- redeem() }() - results := []exchange{redeem(), <-other} + results := []attempt{redeem(), <-other} var minted, rejected int for _, result := range results { @@ -327,7 +357,7 @@ func TestOAuth2TokenExchangeSingleUse(t *testing.T) { t.Fatalf("unexpected status %d: %s", result.status, result.body) } } - require.Equal(t, 1, minted, "a code may mint at most one token") + require.Equal(t, 1, minted, msg) require.Equal(t, 1, rejected) }