From 70cd604f0601fd0d400c0e03a4318a62497b9900 Mon Sep 17 00:00:00 2001 From: Asher Date: Thu, 23 Jul 2026 14:31:21 -0800 Subject: [PATCH 01/12] fix: add distributed protection for external auth refresh Before refreshing, the caller has to put a lease on the row to prevent any other replicas from trying to refresh using the same token. Callers unable to get a lease enter a holding pattern waiting for the results. Also use this same lease query to re-read the link. This more or less functions the same as the re-read we used to have after a refresh failure, except it also allows us to avoid making the failed request in the first place. The downside is that we have to refetch the link every time even if the caller just fetched it, but since we have to make a database call anyway... --- coderd/database/dbauthz/dbauthz.go | 21 +- coderd/database/dbauthz/dbauthz_test.go | 24 +- coderd/database/dbmetrics/querymetrics.go | 24 +- coderd/database/dbmock/dbmock.go | 43 +- coderd/database/dump.sql | 5 +- .../000571_external_auth_lock.down.sql | 1 + .../000571_external_auth_lock.up.sql | 2 + coderd/database/models.go | 2 + coderd/database/querier.go | 10 +- coderd/database/queries.sql.go | 171 +++--- coderd/database/queries/externalauth.sql | 62 ++- coderd/externalauth.go | 40 +- coderd/externalauth/externalauth.go | 213 ++++--- coderd/externalauth/externalauth_test.go | 518 +++++++++++++----- .../provisionerdserver_test.go | 4 +- enterprise/dbcrypt/cliutil.go | 40 +- enterprise/dbcrypt/dbcrypt.go | 62 +-- enterprise/dbcrypt/dbcrypt_internal_test.go | 81 ++- 18 files changed, 854 insertions(+), 469 deletions(-) create mode 100644 coderd/database/migrations/000571_external_auth_lock.down.sql create mode 100644 coderd/database/migrations/000571_external_auth_lock.up.sql diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index c6618192877..f68a7a1b2f3 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -1731,6 +1731,13 @@ func scopedOrgRoleIdentifiers(names []string, orgID uuid.UUID) []rbac.RoleIdenti return out } +func (q *querier) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + fetch := func(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) + } + return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.AcquireExternalAuthLinkRefreshLease)(ctx, arg) +} + func (q *querier) AcquireLock(ctx context.Context, id int64) error { return q.db.AcquireLock(ctx, id) } @@ -7037,6 +7044,13 @@ func (q *querier) RegisterWorkspaceProxy(ctx context.Context, arg database.Regis return updateWithReturn(q.log, q.auth, fetch, q.db.RegisterWorkspaceProxy)(ctx, arg) } +func (q *querier) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { + fetch := func(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) + } + return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.ReleaseExternalAuthLinkRefreshLease)(ctx, arg) +} + func (q *querier) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { // This is a system function to clear user groups in group sync. if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil { @@ -7594,13 +7608,6 @@ func (q *querier) UpdateExternalAuthLink(ctx context.Context, arg database.Updat return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.UpdateExternalAuthLink)(ctx, arg) } -func (q *querier) UpdateExternalAuthLinkRefreshToken(ctx context.Context, arg database.UpdateExternalAuthLinkRefreshTokenParams) error { - fetch := func(ctx context.Context, arg database.UpdateExternalAuthLinkRefreshTokenParams) (database.ExternalAuthLink, error) { - return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) - } - return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.UpdateExternalAuthLinkRefreshToken)(ctx, arg) -} - func (q *querier) UpdateGitSSHKey(ctx context.Context, arg database.UpdateGitSSHKeyParams) (database.GitSSHKey, error) { fetch := func(ctx context.Context, arg database.UpdateGitSSHKeyParams) (database.GitSSHKey, error) { return q.db.GetGitSSHKey(ctx, arg.UserID) diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index a9de44c4601..975ba83a22f 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -3195,13 +3195,6 @@ func (s *MethodTestSuite) TestUser() { dbm.EXPECT().InsertExternalAuthLink(gomock.Any(), arg).Return(database.ExternalAuthLink{}, nil).AnyTimes() check.Args(arg).Asserts(u, policy.ActionUpdatePersonal) })) - s.Run("UpdateExternalAuthLinkRefreshToken", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{}) - arg := database.UpdateExternalAuthLinkRefreshTokenParams{OAuthRefreshToken: "", OAuthRefreshTokenKeyID: "", ProviderID: link.ProviderID, UserID: link.UserID, UpdatedAt: link.UpdatedAt, OldOauthRefreshToken: link.OAuthRefreshToken} - dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() - dbm.EXPECT().UpdateExternalAuthLinkRefreshToken(gomock.Any(), arg).Return(nil).AnyTimes() - check.Args(arg).Asserts(link, policy.ActionUpdatePersonal) - })) s.Run("UpdateExternalAuthLink", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{}) arg := database.UpdateExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID, OAuthAccessToken: link.OAuthAccessToken, OAuthRefreshToken: link.OAuthRefreshToken, OAuthExpiry: link.OAuthExpiry, UpdatedAt: link.UpdatedAt} @@ -3209,6 +3202,23 @@ func (s *MethodTestSuite) TestUser() { dbm.EXPECT().UpdateExternalAuthLink(gomock.Any(), arg).Return(link, nil).AnyTimes() check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) })) + s.Run("AcquireExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{}) + dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() + link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Minute), Valid: true} + arg := database.AcquireExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} + dbm.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(link, nil).AnyTimes() + check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) + })) + s.Run("ReleaseExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{ + RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now().Add(time.Minute), Valid: true}, + }) + dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() + arg := database.ReleaseExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} + dbm.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(nil).AnyTimes() + check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns() + })) s.Run("UpdateUserLink", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { link := testutil.Fake(s.T(), faker, database.UserLink{}) arg := database.UpdateUserLinkParams{OAuthAccessToken: link.OAuthAccessToken, OAuthRefreshToken: link.OAuthRefreshToken, OAuthExpiry: link.OAuthExpiry, UserID: link.UserID, LoginType: link.LoginType, Claims: database.UserLinkClaims{}} diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 4a986ad1eb1..e295484cdad 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -105,6 +105,14 @@ func (m queryMetricsStore) DeleteOrganization(ctx context.Context, id uuid.UUID) return r0 } +func (m queryMetricsStore) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + start := time.Now() + r0, r1 := m.s.AcquireExternalAuthLinkRefreshLease(ctx, arg) + m.queryLatencies.WithLabelValues("AcquireExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "AcquireExternalAuthLinkRefreshLease").Inc() + return r0, r1 +} + func (m queryMetricsStore) AcquireLock(ctx context.Context, pgAdvisoryXactLock int64) error { start := time.Now() r0 := m.s.AcquireLock(ctx, pgAdvisoryXactLock) @@ -4953,6 +4961,14 @@ func (m queryMetricsStore) RegisterWorkspaceProxy(ctx context.Context, arg datab return r0, r1 } +func (m queryMetricsStore) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { + start := time.Now() + r0 := m.s.ReleaseExternalAuthLinkRefreshLease(ctx, arg) + m.queryLatencies.WithLabelValues("ReleaseExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ReleaseExternalAuthLinkRefreshLease").Inc() + return r0 +} + func (m queryMetricsStore) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { start := time.Now() r0, r1 := m.s.RemoveUserFromGroups(ctx, arg) @@ -5353,14 +5369,6 @@ func (m queryMetricsStore) UpdateExternalAuthLink(ctx context.Context, arg datab return r0, r1 } -func (m queryMetricsStore) UpdateExternalAuthLinkRefreshToken(ctx context.Context, arg database.UpdateExternalAuthLinkRefreshTokenParams) error { - start := time.Now() - r0 := m.s.UpdateExternalAuthLinkRefreshToken(ctx, arg) - m.queryLatencies.WithLabelValues("UpdateExternalAuthLinkRefreshToken").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateExternalAuthLinkRefreshToken").Inc() - return r0 -} - func (m queryMetricsStore) UpdateGitSSHKey(ctx context.Context, arg database.UpdateGitSSHKeyParams) (database.GitSSHKey, error) { start := time.Now() r0, r1 := m.s.UpdateGitSSHKey(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 4f09e5a4b7f..3c14163b0b6 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -45,6 +45,21 @@ func (m *MockStore) EXPECT() *MockStoreMockRecorder { return m.recorder } +// AcquireExternalAuthLinkRefreshLease mocks base method. +func (m *MockStore) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "AcquireExternalAuthLinkRefreshLease", ctx, arg) + ret0, _ := ret[0].(database.ExternalAuthLink) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// AcquireExternalAuthLinkRefreshLease indicates an expected call of AcquireExternalAuthLinkRefreshLease. +func (mr *MockStoreMockRecorder) AcquireExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AcquireExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).AcquireExternalAuthLinkRefreshLease), ctx, arg) +} + // AcquireLock mocks base method. func (m *MockStore) AcquireLock(ctx context.Context, pgAdvisoryXactLock int64) error { m.ctrl.T.Helper() @@ -9385,6 +9400,20 @@ func (mr *MockStoreMockRecorder) RegisterWorkspaceProxy(ctx, arg any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterWorkspaceProxy", reflect.TypeOf((*MockStore)(nil).RegisterWorkspaceProxy), ctx, arg) } +// ReleaseExternalAuthLinkRefreshLease mocks base method. +func (m *MockStore) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReleaseExternalAuthLinkRefreshLease", ctx, arg) + ret0, _ := ret[0].(error) + return ret0 +} + +// ReleaseExternalAuthLinkRefreshLease indicates an expected call of ReleaseExternalAuthLinkRefreshLease. +func (mr *MockStoreMockRecorder) ReleaseExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReleaseExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).ReleaseExternalAuthLinkRefreshLease), ctx, arg) +} + // RemoveUserFromGroups mocks base method. func (m *MockStore) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { m.ctrl.T.Helper() @@ -10118,20 +10147,6 @@ func (mr *MockStoreMockRecorder) UpdateExternalAuthLink(ctx, arg any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateExternalAuthLink", reflect.TypeOf((*MockStore)(nil).UpdateExternalAuthLink), ctx, arg) } -// UpdateExternalAuthLinkRefreshToken mocks base method. -func (m *MockStore) UpdateExternalAuthLinkRefreshToken(ctx context.Context, arg database.UpdateExternalAuthLinkRefreshTokenParams) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "UpdateExternalAuthLinkRefreshToken", ctx, arg) - ret0, _ := ret[0].(error) - return ret0 -} - -// UpdateExternalAuthLinkRefreshToken indicates an expected call of UpdateExternalAuthLinkRefreshToken. -func (mr *MockStoreMockRecorder) UpdateExternalAuthLinkRefreshToken(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateExternalAuthLinkRefreshToken", reflect.TypeOf((*MockStore)(nil).UpdateExternalAuthLinkRefreshToken), ctx, arg) -} - // UpdateGitSSHKey mocks base method. func (m *MockStore) UpdateGitSSHKey(ctx context.Context, arg database.UpdateGitSSHKeyParams) (database.GitSSHKey, error) { m.ctrl.T.Helper() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index d70220b91ee..bd3aa4f9676 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -2335,7 +2335,8 @@ CREATE TABLE external_auth_links ( oauth_access_token_key_id text, oauth_refresh_token_key_id text, oauth_extra jsonb, - oauth_refresh_failure_reason text DEFAULT ''::text NOT NULL + oauth_refresh_failure_reason text DEFAULT ''::text NOT NULL, + refresh_lease_expires_at timestamp with time zone ); COMMENT ON COLUMN external_auth_links.oauth_access_token_key_id IS 'The ID of the key used to encrypt the OAuth access token. If this is NULL, the access token is not encrypted'; @@ -2344,6 +2345,8 @@ COMMENT ON COLUMN external_auth_links.oauth_refresh_token_key_id IS 'The ID of t COMMENT ON COLUMN external_auth_links.oauth_refresh_failure_reason IS 'This error means the refresh token is invalid. Cached so we can avoid calling the external provider again for the same error.'; +COMMENT ON COLUMN external_auth_links.refresh_lease_expires_at IS 'Indicates a replica is refreshing the token; prevents concurrent refreshes.'; + CREATE TABLE files ( hash character varying(64) NOT NULL, created_at timestamp with time zone NOT NULL, diff --git a/coderd/database/migrations/000571_external_auth_lock.down.sql b/coderd/database/migrations/000571_external_auth_lock.down.sql new file mode 100644 index 00000000000..c51dc71bbc7 --- /dev/null +++ b/coderd/database/migrations/000571_external_auth_lock.down.sql @@ -0,0 +1 @@ +ALTER TABLE external_auth_links DROP COLUMN IF EXISTS refresh_lease_expires_at; diff --git a/coderd/database/migrations/000571_external_auth_lock.up.sql b/coderd/database/migrations/000571_external_auth_lock.up.sql new file mode 100644 index 00000000000..6853d9d8791 --- /dev/null +++ b/coderd/database/migrations/000571_external_auth_lock.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE external_auth_links ADD COLUMN IF NOT EXISTS refresh_lease_expires_at timestamp WITH time zone DEFAULT NULL; +COMMENT ON COLUMN external_auth_links.refresh_lease_expires_at IS 'Indicates a replica is refreshing the token; prevents concurrent refreshes.'; diff --git a/coderd/database/models.go b/coderd/database/models.go index c80a23665c0..ea7cbd35e9f 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -5316,6 +5316,8 @@ type ExternalAuthLink struct { OAuthExtra pqtype.NullRawMessage `db:"oauth_extra" json:"oauth_extra"` // This error means the refresh token is invalid. Cached so we can avoid calling the external provider again for the same error. OauthRefreshFailureReason string `db:"oauth_refresh_failure_reason" json:"oauth_refresh_failure_reason"` + // Indicates a replica is refreshing the token; prevents concurrent refreshes. + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` } type File struct { diff --git a/coderd/database/querier.go b/coderd/database/querier.go index bbc90e02859..87647dc6d11 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -13,6 +13,8 @@ import ( ) type sqlcQuerier interface { + // Only set the lease if there is not already a non-expired one. + AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg AcquireExternalAuthLinkRefreshLeaseParams) (ExternalAuthLink, error) // Blocks until the lock is acquired. // // This must be called from within a transaction. The lock will be automatically @@ -1315,6 +1317,8 @@ type sqlcQuerier interface { PopNextQueuedMessage(ctx context.Context, chatID uuid.UUID) (ChatQueuedMessage, error) ReduceWorkspaceAgentShareLevelToAuthenticatedByTemplate(ctx context.Context, templateID uuid.UUID) error RegisterWorkspaceProxy(ctx context.Context, arg RegisterWorkspaceProxyParams) (WorkspaceProxy, error) + // Only unset the lease if it matches the one passed in. + ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg ReleaseExternalAuthLinkRefreshLeaseParams) error RemoveUserFromGroups(ctx context.Context, arg RemoveUserFromGroupsParams) ([]uuid.UUID, error) // Mutates only created_at on the target row; ids are unchanged so // consumers can keep tracking queued messages by id. @@ -1481,12 +1485,8 @@ type sqlcQuerier interface { // rows in place. UpdateEncryptedAIProviderSettings(ctx context.Context, arg UpdateEncryptedAIProviderSettingsParams) (AIProvider, error) UpdateEncryptedUserAIProviderKey(ctx context.Context, arg UpdateEncryptedUserAIProviderKeyParams) (UserAIProviderKey, error) + // If a refresh lease is provided, the row is only updated if the lease matches. UpdateExternalAuthLink(ctx context.Context, arg UpdateExternalAuthLinkParams) (ExternalAuthLink, error) - // Optimistic lock: only update the row if the refresh token in the database - // still matches the one we read before attempting the refresh. This prevents - // a concurrent caller that lost a token-refresh race from overwriting a valid - // token stored by the winner. - UpdateExternalAuthLinkRefreshToken(ctx context.Context, arg UpdateExternalAuthLinkRefreshTokenParams) error UpdateGitSSHKey(ctx context.Context, arg UpdateGitSSHKeyParams) (GitSSHKey, error) UpdateGroupByID(ctx context.Context, arg UpdateGroupByIDParams) (Group, error) UpdateInactiveUsersToDormant(ctx context.Context, arg UpdateInactiveUsersToDormantParams) ([]UpdateInactiveUsersToDormantRow, error) diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index f1029f44a15..d71b495b258 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -13956,6 +13956,45 @@ func (q *sqlQuerier) RevokeDBCryptKey(ctx context.Context, activeKeyDigest strin return err } +const acquireExternalAuthLinkRefreshLease = `-- name: AcquireExternalAuthLinkRefreshLease :one +UPDATE + external_auth_links +SET + refresh_lease_expires_at = $1 +WHERE + provider_id = $2 + AND user_id = $3 + AND (refresh_lease_expires_at IS NULL OR refresh_lease_expires_at < NOW()) +RETURNING provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at +` + +type AcquireExternalAuthLinkRefreshLeaseParams struct { + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` + ProviderID string `db:"provider_id" json:"provider_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` +} + +// Only set the lease if there is not already a non-expired one. +func (q *sqlQuerier) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg AcquireExternalAuthLinkRefreshLeaseParams) (ExternalAuthLink, error) { + row := q.db.QueryRowContext(ctx, acquireExternalAuthLinkRefreshLease, arg.RefreshLeaseExpiresAt, arg.ProviderID, arg.UserID) + var i ExternalAuthLink + err := row.Scan( + &i.ProviderID, + &i.UserID, + &i.CreatedAt, + &i.UpdatedAt, + &i.OAuthAccessToken, + &i.OAuthRefreshToken, + &i.OAuthExpiry, + &i.OAuthAccessTokenKeyID, + &i.OAuthRefreshTokenKeyID, + &i.OAuthExtra, + &i.OauthRefreshFailureReason, + &i.RefreshLeaseExpiresAt, + ) + return i, err +} + const deleteExternalAuthLink = `-- name: DeleteExternalAuthLink :exec DELETE FROM external_auth_links WHERE provider_id = $1 AND user_id = $2 ` @@ -13971,7 +14010,7 @@ func (q *sqlQuerier) DeleteExternalAuthLink(ctx context.Context, arg DeleteExter } const getExternalAuthLink = `-- name: GetExternalAuthLink :one -SELECT provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason FROM external_auth_links WHERE provider_id = $1 AND user_id = $2 +SELECT provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at FROM external_auth_links WHERE provider_id = $1 AND user_id = $2 ` type GetExternalAuthLinkParams struct { @@ -13994,12 +14033,13 @@ func (q *sqlQuerier) GetExternalAuthLink(ctx context.Context, arg GetExternalAut &i.OAuthRefreshTokenKeyID, &i.OAuthExtra, &i.OauthRefreshFailureReason, + &i.RefreshLeaseExpiresAt, ) return i, err } const getExternalAuthLinksByUserID = `-- name: GetExternalAuthLinksByUserID :many -SELECT provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason FROM external_auth_links WHERE user_id = $1 +SELECT provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at FROM external_auth_links WHERE user_id = $1 ` func (q *sqlQuerier) GetExternalAuthLinksByUserID(ctx context.Context, userID uuid.UUID) ([]ExternalAuthLink, error) { @@ -14023,6 +14063,7 @@ func (q *sqlQuerier) GetExternalAuthLinksByUserID(ctx context.Context, userID uu &i.OAuthRefreshTokenKeyID, &i.OAuthExtra, &i.OauthRefreshFailureReason, + &i.RefreshLeaseExpiresAt, ); err != nil { return nil, err } @@ -14060,7 +14101,7 @@ INSERT INTO external_auth_links ( $8, $9, $10 -) RETURNING provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason +) RETURNING provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at ` type InsertExternalAuthLinkParams struct { @@ -14102,42 +14143,71 @@ func (q *sqlQuerier) InsertExternalAuthLink(ctx context.Context, arg InsertExter &i.OAuthRefreshTokenKeyID, &i.OAuthExtra, &i.OauthRefreshFailureReason, + &i.RefreshLeaseExpiresAt, ) return i, err } +const releaseExternalAuthLinkRefreshLease = `-- name: ReleaseExternalAuthLinkRefreshLease :exec +UPDATE + external_auth_links +SET + refresh_lease_expires_at = NULL +WHERE + provider_id = $1 + AND user_id = $2 + AND refresh_lease_expires_at = $3 +` + +type ReleaseExternalAuthLinkRefreshLeaseParams struct { + ProviderID string `db:"provider_id" json:"provider_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` +} + +// Only unset the lease if it matches the one passed in. +func (q *sqlQuerier) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg ReleaseExternalAuthLinkRefreshLeaseParams) error { + _, err := q.db.ExecContext(ctx, releaseExternalAuthLinkRefreshLease, arg.ProviderID, arg.UserID, arg.RefreshLeaseExpiresAt) + return err +} + const updateExternalAuthLink = `-- name: UpdateExternalAuthLink :one UPDATE external_auth_links SET - updated_at = $3, - oauth_access_token = $4, - oauth_access_token_key_id = $5, - oauth_refresh_token = $6, - oauth_refresh_token_key_id = $7, - oauth_expiry = $8, - oauth_extra = $9, - -- Only 'UpdateExternalAuthLinkRefreshToken' supports updating the oauth_refresh_failure_reason. - -- Any updates to the external auth link, will be assumed to change the state and clear - -- any cached errors. - oauth_refresh_failure_reason = '' -WHERE provider_id = $1 AND user_id = $2 RETURNING provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason + updated_at = $4, + oauth_access_token = $5, + oauth_access_token_key_id = $6, + oauth_refresh_token = $7, + oauth_refresh_token_key_id = $8, + oauth_expiry = $9, + oauth_extra = $10, + oauth_refresh_failure_reason = $11 +WHERE + provider_id = $1 + AND user_id = $2 + AND (refresh_lease_expires_at = $3 OR $3 IS NULL) +RETURNING provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at ` type UpdateExternalAuthLinkParams struct { - ProviderID string `db:"provider_id" json:"provider_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` - UpdatedAt time.Time `db:"updated_at" json:"updated_at"` - OAuthAccessToken string `db:"oauth_access_token" json:"oauth_access_token"` - OAuthAccessTokenKeyID sql.NullString `db:"oauth_access_token_key_id" json:"oauth_access_token_key_id"` - OAuthRefreshToken string `db:"oauth_refresh_token" json:"oauth_refresh_token"` - OAuthRefreshTokenKeyID sql.NullString `db:"oauth_refresh_token_key_id" json:"oauth_refresh_token_key_id"` - OAuthExpiry time.Time `db:"oauth_expiry" json:"oauth_expiry"` - OAuthExtra pqtype.NullRawMessage `db:"oauth_extra" json:"oauth_extra"` -} - + ProviderID string `db:"provider_id" json:"provider_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` + UpdatedAt time.Time `db:"updated_at" json:"updated_at"` + OAuthAccessToken string `db:"oauth_access_token" json:"oauth_access_token"` + OAuthAccessTokenKeyID sql.NullString `db:"oauth_access_token_key_id" json:"oauth_access_token_key_id"` + OAuthRefreshToken string `db:"oauth_refresh_token" json:"oauth_refresh_token"` + OAuthRefreshTokenKeyID sql.NullString `db:"oauth_refresh_token_key_id" json:"oauth_refresh_token_key_id"` + OAuthExpiry time.Time `db:"oauth_expiry" json:"oauth_expiry"` + OAuthExtra pqtype.NullRawMessage `db:"oauth_extra" json:"oauth_extra"` + OauthRefreshFailureReason string `db:"oauth_refresh_failure_reason" json:"oauth_refresh_failure_reason"` +} + +// If a refresh lease is provided, the row is only updated if the lease matches. func (q *sqlQuerier) UpdateExternalAuthLink(ctx context.Context, arg UpdateExternalAuthLinkParams) (ExternalAuthLink, error) { row := q.db.QueryRowContext(ctx, updateExternalAuthLink, arg.ProviderID, arg.UserID, + arg.RefreshLeaseExpiresAt, arg.UpdatedAt, arg.OAuthAccessToken, arg.OAuthAccessTokenKeyID, @@ -14145,6 +14215,7 @@ func (q *sqlQuerier) UpdateExternalAuthLink(ctx context.Context, arg UpdateExter arg.OAuthRefreshTokenKeyID, arg.OAuthExpiry, arg.OAuthExtra, + arg.OauthRefreshFailureReason, ) var i ExternalAuthLink err := row.Scan( @@ -14159,57 +14230,11 @@ func (q *sqlQuerier) UpdateExternalAuthLink(ctx context.Context, arg UpdateExter &i.OAuthRefreshTokenKeyID, &i.OAuthExtra, &i.OauthRefreshFailureReason, + &i.RefreshLeaseExpiresAt, ) return i, err } -const updateExternalAuthLinkRefreshToken = `-- name: UpdateExternalAuthLinkRefreshToken :exec -UPDATE - external_auth_links -SET - -- oauth_refresh_failure_reason can be set to cache the failure reason - -- for subsequent refresh attempts. - oauth_refresh_failure_reason = $1, - oauth_refresh_token = $2, - updated_at = $3 -WHERE - provider_id = $4 -AND - user_id = $5 -AND - oauth_refresh_token = $6 -AND - -- Required for sqlc to generate a parameter for the oauth_refresh_token_key_id - $7 :: text = $7 :: text -` - -type UpdateExternalAuthLinkRefreshTokenParams struct { - OauthRefreshFailureReason string `db:"oauth_refresh_failure_reason" json:"oauth_refresh_failure_reason"` - OAuthRefreshToken string `db:"oauth_refresh_token" json:"oauth_refresh_token"` - UpdatedAt time.Time `db:"updated_at" json:"updated_at"` - ProviderID string `db:"provider_id" json:"provider_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` - OldOauthRefreshToken string `db:"old_oauth_refresh_token" json:"old_oauth_refresh_token"` - OAuthRefreshTokenKeyID string `db:"oauth_refresh_token_key_id" json:"oauth_refresh_token_key_id"` -} - -// Optimistic lock: only update the row if the refresh token in the database -// still matches the one we read before attempting the refresh. This prevents -// a concurrent caller that lost a token-refresh race from overwriting a valid -// token stored by the winner. -func (q *sqlQuerier) UpdateExternalAuthLinkRefreshToken(ctx context.Context, arg UpdateExternalAuthLinkRefreshTokenParams) error { - _, err := q.db.ExecContext(ctx, updateExternalAuthLinkRefreshToken, - arg.OauthRefreshFailureReason, - arg.OAuthRefreshToken, - arg.UpdatedAt, - arg.ProviderID, - arg.UserID, - arg.OldOauthRefreshToken, - arg.OAuthRefreshTokenKeyID, - ) - return err -} - const getFileByHashAndCreator = `-- name: GetFileByHashAndCreator :one SELECT hash, created_at, created_by, mimetype, data, id diff --git a/coderd/database/queries/externalauth.sql b/coderd/database/queries/externalauth.sql index e5d0ec548bf..d92e6c6c561 100644 --- a/coderd/database/queries/externalauth.sql +++ b/coderd/database/queries/externalauth.sql @@ -33,39 +33,41 @@ INSERT INTO external_auth_links ( ) RETURNING *; -- name: UpdateExternalAuthLink :one +-- If a refresh lease is provided, the row is only updated if the lease matches. UPDATE external_auth_links SET - updated_at = $3, - oauth_access_token = $4, - oauth_access_token_key_id = $5, - oauth_refresh_token = $6, - oauth_refresh_token_key_id = $7, - oauth_expiry = $8, - oauth_extra = $9, - -- Only 'UpdateExternalAuthLinkRefreshToken' supports updating the oauth_refresh_failure_reason. - -- Any updates to the external auth link, will be assumed to change the state and clear - -- any cached errors. - oauth_refresh_failure_reason = '' -WHERE provider_id = $1 AND user_id = $2 RETURNING *; + updated_at = $4, + oauth_access_token = $5, + oauth_access_token_key_id = $6, + oauth_refresh_token = $7, + oauth_refresh_token_key_id = $8, + oauth_expiry = $9, + oauth_extra = $10, + oauth_refresh_failure_reason = $11 +WHERE + provider_id = $1 + AND user_id = $2 + AND (refresh_lease_expires_at = $3 OR $3 IS NULL) +RETURNING *; + +-- name: AcquireExternalAuthLinkRefreshLease :one +-- Only set the lease if there is not already a non-expired one. +UPDATE + external_auth_links +SET + refresh_lease_expires_at = @refresh_lease_expires_at +WHERE + provider_id = @provider_id + AND user_id = @user_id + AND (refresh_lease_expires_at IS NULL OR refresh_lease_expires_at < NOW()) +RETURNING *; --- name: UpdateExternalAuthLinkRefreshToken :exec --- Optimistic lock: only update the row if the refresh token in the database --- still matches the one we read before attempting the refresh. This prevents --- a concurrent caller that lost a token-refresh race from overwriting a valid --- token stored by the winner. +-- name: ReleaseExternalAuthLinkRefreshLease :exec +-- Only unset the lease if it matches the one passed in. UPDATE external_auth_links SET - -- oauth_refresh_failure_reason can be set to cache the failure reason - -- for subsequent refresh attempts. - oauth_refresh_failure_reason = @oauth_refresh_failure_reason, - oauth_refresh_token = @oauth_refresh_token, - updated_at = @updated_at + refresh_lease_expires_at = NULL WHERE - provider_id = @provider_id -AND - user_id = @user_id -AND - oauth_refresh_token = @old_oauth_refresh_token -AND - -- Required for sqlc to generate a parameter for the oauth_refresh_token_key_id - @oauth_refresh_token_key_id :: text = @oauth_refresh_token_key_id :: text; + provider_id = @provider_id + AND user_id = @user_id + AND refresh_lease_expires_at = @refresh_lease_expires_at; diff --git a/coderd/externalauth.go b/coderd/externalauth.go index 51b7727c00d..04ab62d3af6 100644 --- a/coderd/externalauth.go +++ b/coderd/externalauth.go @@ -203,15 +203,17 @@ func (api *API) postExternalAuthDeviceByID(rw http.ResponseWriter, r *http.Reque } } else { _, err = api.Database.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ - ProviderID: config.ID, - UserID: apiKey.UserID, - UpdatedAt: dbtime.Now(), - OAuthAccessToken: token.AccessToken, - OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthRefreshToken: token.RefreshToken, - OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthExpiry: token.Expiry, - OAuthExtra: pqtype.NullRawMessage{}, + ProviderID: config.ID, + UserID: apiKey.UserID, + UpdatedAt: dbtime.Now(), + OAuthAccessToken: token.AccessToken, + OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthRefreshToken: token.RefreshToken, + OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthExpiry: token.Expiry, + OAuthExtra: pqtype.NullRawMessage{}, + OauthRefreshFailureReason: "", + RefreshLeaseExpiresAt: sql.NullTime{}, }) if err != nil { httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ @@ -305,15 +307,17 @@ func (api *API) externalAuthCallback(externalAuthConfig *externalauth.Config) ht } } else { _, err = api.Database.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ - ProviderID: externalAuthConfig.ID, - UserID: apiKey.UserID, - UpdatedAt: dbtime.Now(), - OAuthAccessToken: state.Token.AccessToken, - OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthRefreshToken: state.Token.RefreshToken, - OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthExpiry: state.Token.Expiry, - OAuthExtra: extra, + ProviderID: externalAuthConfig.ID, + UserID: apiKey.UserID, + UpdatedAt: dbtime.Now(), + OAuthAccessToken: state.Token.AccessToken, + OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthRefreshToken: state.Token.RefreshToken, + OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthExpiry: state.Token.Expiry, + OAuthExtra: extra, + OauthRefreshFailureReason: "", + RefreshLeaseExpiresAt: sql.NullTime{}, }) if err != nil { httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index 4f0b9bdb5af..fe63781b946 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "encoding/json" + "errors" "fmt" "io" "mime" @@ -55,6 +56,14 @@ const ( // defaultRefreshRetryTimeout bounds the total time spent retrying a // transient refresh failure across all attempts. defaultRefreshRetryTimeout = 10 * time.Second + + // defaultRefreshLeaseInitialBackoff is the starting wait between polls to + // check whether another replica has finished refreshing. + defaultRefreshLeaseInitialBackoff = 50 * time.Millisecond + + // defaultRefreshLeaseInitialBackoff is the maximum wait between polls to + // check whether another replica has finished refreshing. + defaultRefreshLeaseMaxBackoff = 500 * time.Millisecond ) // SingleflightGroup exposes a subset of singleflight.Group for easier testing. @@ -159,6 +168,14 @@ type Config struct { // RefreshGroup deduplicates concurrent requests. RefreshGroup SingleflightGroup + + // RefreshLeaseInitialBackoff is the starting wait between polls to check + // whether another replica has finished refreshing. + RefreshLeaseInitialBackoff time.Duration + + // RefreshLeaseInitialBackoff is the maximum wait between polls to check + // whether another replica has finished refreshing. + RefreshLeaseMaxBackoff time.Duration } // Git returns a Provider for this config if the provider type is a @@ -207,11 +224,28 @@ func IsInvalidTokenError(err error) bool { // RefreshToken automatically refreshes the token if expired and permitted. func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink) (database.ExternalAuthLink, error) { + // If the token is not expired and there is no validation URL, we can + // short-circuit since there is nothing to do. + token := &oauth2.Token{ + AccessToken: externalAuthLink.OAuthAccessToken, + RefreshToken: externalAuthLink.OAuthRefreshToken, + Expiry: externalAuthLink.OAuthExpiry, + } + if c.ValidateURL == "" && token.Valid() { + return externalAuthLink, nil + } + + // If the token is expired and refresh is disabled, we prompt the user to + // authenticate manually again. + if c.NoRefresh && !token.Valid() { + return externalAuthLink, InvalidTokenError("token expired and refreshing is disabled") + } + // Prevent parallel refreshes by waiting for the result of any already // in-flight refresh. Otherwise, the parallel calls will fail with a bad // refresh token error as they can only be used once. key := c.ID + ":" + externalAuthLink.UserID.String() - ch := c.RefreshGroup.DoChan(key, func() (any, error) { + ch := c.RefreshGroup.DoChan(key, func() (newLink any, refreshErr error) { // Use a detached context so if a request is canceled or times out it does // not cancel all the other requests as well. The deadline is arbitrary but // we give at least enough time for the refresh timeout then another 10 @@ -220,9 +254,64 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu if c.RefreshRetryTimeout > 0 { timeout += c.RefreshRetryTimeout } - rctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), timeout) + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), timeout) defer cancel() - return c.innerRefreshToken(rctx, db, externalAuthLink) + + lease := dbtime.Now().Add(timeout) + + // There may be other replicas also wanting to refresh; try to get a lease + // on the row. This also ensures we have the latest link. + var dblink database.ExternalAuthLink + initial := defaultRefreshLeaseInitialBackoff + if c.RefreshLeaseInitialBackoff > 0 { + initial = c.RefreshLeaseInitialBackoff + } + maximum := defaultRefreshLeaseMaxBackoff + if c.RefreshLeaseMaxBackoff > 0 { + maximum = c.RefreshLeaseMaxBackoff + } + r := retry.New(initial, maximum) + for { + var err error + dblink, err = db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + }) + switch { + case errors.Is(err, sql.ErrNoRows): + // Something still holds the lock; keep waiting. + if !r.Wait(ctx) { + return externalAuthLink, ctx.Err() + } + case err != nil: + // Some kind of DB error. + return externalAuthLink, err + default: + // We now hold the lock. + goto refresh + } + } + + refresh: + defer func() { + refreshErr = errors.Join(refreshErr, db.ReleaseExternalAuthLinkRefreshLease(ctx, database.ReleaseExternalAuthLinkRefreshLeaseParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + })) + }() + + // If we got a different token, it means something else refreshed either + // while we were waiting or before we got a hold of the lease but after we + // initially fetched the link. + if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { + return dblink, nil + } + + // Otherwise the token has still not been updated; refresh it now. + newLink, refreshErr = c.innerRefreshToken(ctx, db, dblink, lease) + return newLink, refreshErr }) select { case results := <-ch: @@ -237,37 +326,27 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu } } -func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink) (database.ExternalAuthLink, error) { - // If the token is expired and refresh is disabled, we prompt - // the user to authenticate again. - if c.NoRefresh && - // If the time is set to 0, then it should never expire. - // This is true for github, which has no expiry. - !externalAuthLink.OAuthExpiry.IsZero() && - externalAuthLink.OAuthExpiry.Before(dbtime.Now()) { - return externalAuthLink, InvalidTokenError("token expired, refreshing is either disabled or refreshing failed and will not be retried") - } - - refreshToken := externalAuthLink.OAuthRefreshToken - - // This is additional defensive programming. Because TokenSource is an interface, - // we cannot be sure that the implementation will treat an 'IsZero' time - // as "not-expired". The default implementation does, but a custom implementation - // might not. Removing the refreshToken will guarantee a refresh will fail. - if c.NoRefresh { - refreshToken = "" - } - +func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, lease time.Time) (database.ExternalAuthLink, error) { existingToken := &oauth2.Token{ AccessToken: externalAuthLink.OAuthAccessToken, - RefreshToken: refreshToken, + RefreshToken: externalAuthLink.OAuthRefreshToken, Expiry: externalAuthLink.OAuthExpiry, } + // This is additional defensive programming. Because TokenSource is an + // interface, we cannot be sure that the implementation will treat an 'IsZero' + // time as "not-expired". The default implementation does, but a custom + // implementation might not. Removing the refresh token will guarantee a + // refresh will fail. + if c.NoRefresh { + existingToken.RefreshToken = "" + } + // NOTE: TokenSource(...).Token() will short-circuit if the token: // - is not expired (returns original token) // - is expired and has no refresh token (returns error) - // This means we will avoid making useless HTTP requests. + // This means we will avoid making useless HTTP requests, and only get errors + // when an actual refresh attempt is made. // // External providers (GitHub in particular) intermittently fail token // refreshes with transient errors such as 5xx responses, network timeouts, @@ -278,54 +357,43 @@ func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, exter // will never succeed and retrying wastes the refresh quota. token, err := c.refreshTokenWithRetry(ctx, existingToken) if err != nil { - // A refresh attempt can fail for numerous reasons. If it fails because - // of a bad refresh token, then the refresh token is invalid, and we - // should get rid of it. Keeping it around will cause additional refresh - // attempts that will fail and cost us api rate limits. - // - // The error message is saved for debugging purposes. + // A refresh attempt can fail for numerous reasons. If it fails because of a + // bad refresh token, then the refresh token is invalid, and we should get + // rid of it. Keeping it around will cause additional refresh attempts that + // will fail and cost us api rate limits. Also save the error message for + // debugging purposes. if isFailedRefresh(existingToken, err) { - // Before caching the failure, re-read the external auth link from the - // database. A nearly-concurrent request may have already refreshed the - // token successfully, consuming the single-use refresh token (e.g., - // GitHub App tokens). In that case our "bad_refresh_token" error is a - // false positive from losing the race, and we should use the winner's - // updated token instead of poisoning the database with a cached failure. - currentLink, readErr := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - }) - if readErr == nil && currentLink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { - return currentLink, nil - } - reason := err.Error() if len(reason) > failureReasonLimit { // Limit the length of the error message to prevent // spamming the database with long error messages. reason = reason[:failureReasonLimit] } - dbExecErr := db.UpdateExternalAuthLinkRefreshToken(ctx, database.UpdateExternalAuthLinkRefreshTokenParams{ - // Adding a reason will prevent further attempts to try and refresh the token. - OauthRefreshFailureReason: reason, - // Remove the invalid refresh token so it is never used again. The cached - // `reason` can be used to know why this field was zeroed out. + dblink, updateErr := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + UpdatedAt: dbtime.Now(), + // Remove the invalid refresh token so it is never used again. OAuthRefreshToken: "", - OAuthRefreshTokenKeyID: externalAuthLink.OAuthRefreshTokenKeyID.String, - UpdatedAt: dbtime.Now(), - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - // Optimistic lock: only clear the token if it hasn't been - // updated by a concurrent caller that won the refresh race. - OldOauthRefreshToken: externalAuthLink.OAuthRefreshToken, + OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required + // The cached reason can be used to know why the token was zeroed out. + OauthRefreshFailureReason: reason, + // Preserve the access token, expiry, and extra info as they are. + OAuthAccessToken: externalAuthLink.OAuthAccessToken, + OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthExpiry: externalAuthLink.OAuthExpiry, + OAuthExtra: externalAuthLink.OAuthExtra, + // The row will only update if we hold the current lease. This is + // somewhat redundant since if we lost the lease our context would be + // expired anyway, so it is not actually possible to get the + // sql.ErrNoRows that would result from this. + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, }) - if dbExecErr != nil { + if updateErr != nil { // This error should be rare. - return externalAuthLink, InvalidTokenError(fmt.Sprintf("refresh token failed: %q, then removing refresh token failed: %q", err.Error(), dbExecErr.Error())) + return externalAuthLink, InvalidTokenError(fmt.Sprintf("refresh token failed: %q, then removing refresh token failed: %q", err.Error(), updateErr.Error())) } - // The refresh token was cleared - externalAuthLink.OAuthRefreshToken = "" - externalAuthLink.UpdatedAt = dbtime.Now() + externalAuthLink = dblink } // Unfortunately have to match exactly on the error message string. @@ -347,11 +415,9 @@ func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, exter )) } - return externalAuthLink, InvalidTokenError("token expired, refreshing is either disabled or refreshing failed and will not be retried") + return externalAuthLink, InvalidTokenError("token expired, refreshing failed and will not be retried") } - // Non-expired tokens are short-circuited as noted above; reaching here - // means refresh failed. return externalAuthLink, InvalidTokenError(fmt.Sprintf("refresh token: %s", err.Error())) } @@ -369,7 +435,7 @@ func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, exter originalAccessToken := externalAuthLink.OAuthAccessToken if token.AccessToken != originalAccessToken { updatedAuthLink, err := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ - ProviderID: c.ID, + ProviderID: externalAuthLink.ProviderID, UserID: externalAuthLink.UserID, UpdatedAt: dbtime.Now(), OAuthAccessToken: token.AccessToken, @@ -378,6 +444,13 @@ func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, exter OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required OAuthExpiry: token.Expiry, OAuthExtra: extra, + // If there was any failure before, we can clear it now. + OauthRefreshFailureReason: "", + // The row will only update if we hold the current lease. This is + // somewhat redundant since if we lost the lease our context would be + // expired anyway, so it is not actually possible to get the + // sql.ErrNoRows that would result from this. + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, }) if err != nil { return updatedAuthLink, xerrors.Errorf("persist refreshed token: %w", err) @@ -463,12 +536,8 @@ func (c *Config) refreshTokenWithRetry(ctx context.Context, existingToken *oauth defer retryCancel() backoff := retry.New(initial, maximum) - var ( - token *oauth2.Token - err error - ) for { - token, err = c.TokenSource(ctx, existingToken).Token() + token, err := c.TokenSource(ctx, existingToken).Token() if err == nil || isFailedRefresh(existingToken, err) { return token, err } diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index f29d1a3af61..554c2a55f3c 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -3,6 +3,7 @@ package externalauth_test import ( "bytes" "context" + "database/sql" "encoding/json" "fmt" "io" @@ -48,6 +49,10 @@ func TestRefreshToken(t *testing.T) { t.Run("NoRefreshExpired", func(t *testing.T) { t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -63,17 +68,20 @@ func TestRefreshToken(t *testing.T) { }, ExternalAuthOpt: func(cfg *externalauth.Config) { cfg.NoRefresh = true + // Should abort before entering the group. + cfg.RefreshGroup = nil + }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - // Expire the link - link.OAuthExpiry = expired - - _, err := config.RefreshToken(ctx, nil, link) + // There should be no database calls since we return early. + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + _, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) require.True(t, externalauth.IsInvalidTokenError(err)) - require.Contains(t, err.Error(), "refreshing is either disabled or refreshing failed") + require.Contains(t, err.Error(), "token expired and refreshing is disabled") }) // NoRefreshNoExpiry tests that an oauth token without an expiry is always valid. @@ -81,6 +89,9 @@ func TestRefreshToken(t *testing.T) { t.Run("NoRefreshNoExpiry", func(t *testing.T) { t.Parallel() + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + validated := false fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ @@ -98,18 +109,28 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) // Zero time used link.OAuthExpiry = time.Time{} - _, err := config.RefreshToken(ctx, nil, link) + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + + _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.True(t, validated, "token should have been validated") }) t.Run("FalseIfTokenSourceFails", func(t *testing.T) { t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + config := &externalauth.Config{ InstrumentedOAuth2Config: &testutil.OAuth2Config{ TokenSourceFunc: func() (*oauth2.Token, error) { @@ -119,9 +140,18 @@ func TestRefreshToken(t *testing.T) { RefreshGroup: new(singleflight.Group), } - _, err := config.RefreshToken(context.Background(), nil, database.ExternalAuthLink{ + link := database.ExternalAuthLink{ OAuthExpiry: expired, - }) + } + + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + + ctx := testutil.Context(t, testutil.WaitLong) + _, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) require.True(t, externalauth.IsInvalidTokenError(err)) require.Contains(t, err.Error(), "failure") @@ -148,9 +178,15 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) link.OAuthExpiry = expired + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + _, err := config.RefreshToken(ctx, mDB, link) require.ErrorContains(t, err, staticError) // Unsure if this should be the correct behavior. It's an invalid token because @@ -199,10 +235,17 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) // Expire the link link.OAuthExpiry = expired + // Allow getting the lease for all the temporary error attempts and the + // first bad refresh token attempt. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(4) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(5) + // Make the failure a server internal error. Not related to the token // This should be retried since this error is temporary. refreshErr = &oauth2.RetrieveError{ @@ -222,10 +265,21 @@ func TestRefreshToken(t *testing.T) { } // Try again with a bad refresh token error. This will invalidate the - // refresh token, and not retry again. Expect DB calls to check for - // concurrent refresh (GetExternalAuthLink) and then remove the refresh token. - mDB.EXPECT().GetExternalAuthLink(gomock.Any(), gomock.Any()).Return(link, nil).Times(1) - mDB.EXPECT().UpdateExternalAuthLinkRefreshToken(gomock.Any(), gomock.Any()).Return(nil).Times(1) + // refresh token, and not retry again. + mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). + DoAndReturn(func(_ context.Context, p database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { + link = database.ExternalAuthLink{ + ProviderID: p.ProviderID, + UserID: p.UserID, + OAuthAccessToken: p.OAuthAccessToken, + // This should be zeroed out. + OAuthRefreshToken: p.OAuthRefreshToken, + OAuthExpiry: p.OAuthExpiry, + OauthRefreshFailureReason: p.OauthRefreshFailureReason, + } + return link, nil + }).Times(1) + refreshErr = &oauth2.RetrieveError{ // github error Response: &http.Response{ StatusCode: http.StatusOK, @@ -238,8 +292,11 @@ func TestRefreshToken(t *testing.T) { require.True(t, externalauth.IsInvalidTokenError(err)) require.Equal(t, refreshCount, totalRefreshes) + // Update mock database with the zeroed-out link.. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + // When the refresh token is empty, no api calls should be made - link.OAuthRefreshToken = "" // mock'd db, so manually set the token to '' _, err = config.RefreshToken(ctx, mDB, link) require.Error(t, err) require.True(t, externalauth.IsInvalidTokenError(err)) @@ -279,11 +336,13 @@ func TestRefreshToken(t *testing.T) { cfg.RefreshRetryTimeout = 5 * time.Second }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) oldAccessToken := link.OAuthAccessToken - link.OAuthExpiry = expired updated, err := config.RefreshToken(ctx, db, link) require.NoError(t, err, "transient errors should be retried until success") @@ -301,8 +360,7 @@ func TestRefreshToken(t *testing.T) { t.Run("RefreshTokenBackoffPermanentError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) + db, _ := dbtestutil.NewDB(t) var refreshCalls atomic.Int64 fake, config, link := setupOauth2Test(t, testConfig{ @@ -324,20 +382,14 @@ func TestRefreshToken(t *testing.T) { cfg.RefreshRetryMaxBackoff = 5 * time.Millisecond cfg.RefreshRetryTimeout = time.Second }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, + DB: db, }) - // The race-detection re-read returns the same refresh token so it - // does not look like a concurrent winner. The cached-failure write - // then proceeds. Each runs exactly once for a single refresh attempt. - mDB.EXPECT().GetExternalAuthLink(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().UpdateExternalAuthLinkRefreshToken(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - link.OAuthExpiry = expired - - _, err := config.RefreshToken(ctx, mDB, link) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + _, err := config.RefreshToken(ctx, db, link) require.Error(t, err) require.True(t, externalauth.IsInvalidTokenError(err)) require.Equal(t, int64(1), refreshCalls.Load(), @@ -394,6 +446,12 @@ func TestRefreshToken(t *testing.T) { OAuthExpiry: refreshedToken.Expiry, } + // Only one call should try to get a lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + // The single winning call will update the link. mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Cond(func(params database.UpdateExternalAuthLinkParams) bool { return params.ProviderID == link.ProviderID && params.UserID == link.UserID @@ -426,13 +484,13 @@ func TestRefreshToken(t *testing.T) { // ConcurrentRefreshRace tests what happens a request reads the refresh token // from the database, then another request finishes and updates the token and - // releases the refresh group lock before this request can join. + // releases the refresh group lock before this request can join that group. // - // This request will then fail with `bad_refresh_token` for providers that - // have single-use refresh tokens. It should re-read the token from the - // database after making this failed request to check whether the token was - // updated by another request and returns that rather than incorrectly - // recording in the database that the request failed. + // This request would fail with `bad_refresh_token` for providers that have + // single-use refresh tokens. It should instead re-read the token from the + // database to check whether the token was updated by another request and + // returns that rather than incorrectly recording in the database that the + // request failed. t.Run("ConcurrentRefreshRace", func(t *testing.T) { t.Parallel() @@ -442,35 +500,28 @@ func TestRefreshToken(t *testing.T) { fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { - return &oauth2.RetrieveError{ - Response: &http.Response{ - StatusCode: http.StatusOK, - }, - ErrorCode: "bad_refresh_token", - } + return xerrors.New("should not reach this") }), }, - ExternalAuthOpt: func(cfg *externalauth.Config) {}, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - link.OAuthExpiry = time.Now().Add(time.Hour * -1) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - // Simulate a concurrent winner: when the loser re-reads the - // DB, the refresh token has changed (the winner stored a new - // one). The loser should return the updated link instead of - // caching the failure. + // Simulate that another caller updated the link. winnerLink := link winnerLink.OAuthRefreshToken = "winner-refresh-token" winnerLink.OAuthAccessToken = "winner-access-token" - mDB.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ - ProviderID: link.ProviderID, - UserID: link.UserID, - }).Return(winnerLink, nil).Times(1) - - // UpdateExternalAuthLinkRefreshToken should NOT be called - // because the re-read detected the concurrent refresh. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(winnerLink, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + // UpdateExternalAuthLinkRefreshToken should NOT be called because trying to + // get the lease detected the nearly-concurrent refresh. It should instead + // return the winning token. result, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err, "loser should succeed using the winner's token") require.Equal(t, "winner-access-token", result.OAuthAccessToken) @@ -514,8 +565,7 @@ func TestRefreshToken(t *testing.T) { } } } - // Should never reach here. - return xerrors.New("bad_refresh_token") + return xerrors.New("should not reach this") }), oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { return jwt.MapClaims{}, nil @@ -526,11 +576,13 @@ func TestRefreshToken(t *testing.T) { cfg.RefreshGroup = &group{notify: ch} }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) oldAccessToken := link.OAuthAccessToken oldRefreshToken := link.OAuthRefreshToken - link.OAuthExpiry = expired var wg sync.WaitGroup // Start the first call with the cancelable context. @@ -563,7 +615,7 @@ func TestRefreshToken(t *testing.T) { wg.Wait() // DB link should have been updated. - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + dbLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ ProviderID: link.ProviderID, UserID: link.UserID, }) @@ -577,14 +629,75 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, int64(1), refreshCalls.Load()) }) + t.Run("ReturnsReleaseError", func(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + return nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) {}, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, + }) + + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(xerrors.New("release error")).Times(1) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + _, err := config.RefreshToken(ctx, mDB, link) + require.Error(t, err) + require.ErrorContains(t, err, "release error") + }) + + t.Run("ReturnsCombinedWithReleaseError", func(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + return nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) {}, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, + }) + + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). + Return(database.ExternalAuthLink{}, xerrors.New("update error")).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(xerrors.New("release error")).Times(1) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + _, err := config.RefreshToken(ctx, mDB, link) + require.Error(t, err) + require.ErrorContains(t, err, "release error") + require.ErrorContains(t, err, "update error") + }) + // ValidateFailure tests if the token is no longer valid with a 401 response. t.Run("ValidateFailure", func(t *testing.T) { t.Parallel() ctrl := gomock.NewController(t) mDB := dbmock.NewMockStore(ctrl) - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - Return(database.ExternalAuthLink{}, nil).AnyTimes() const staticError = "static error" validated := false @@ -599,9 +712,18 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) link.OAuthExpiry = expired + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + + mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + _, err := config.RefreshToken(ctx, mDB, link) require.ErrorContains(t, err, "token failed to validate") require.True(t, externalauth.IsInvalidTokenError(err)) @@ -611,6 +733,9 @@ func TestRefreshToken(t *testing.T) { t.Run("ValidateRetryGitHub", func(t *testing.T) { t.Parallel() + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + const staticError = "static error" validateCalls := 0 fake, config, link := setupOauth2Test(t, testConfig{ @@ -631,13 +756,21 @@ func TestRefreshToken(t *testing.T) { ExternalAuthOpt: func(cfg *externalauth.Config) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + // Unlimited lifetime, this is what GitHub returns tokens as. + link.OAuthExpiry = time.Time{} + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - // Unlimited lifetime, this is what GitHub returns tokens as - link.OAuthExpiry = time.Time{} + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) - _, err := config.RefreshToken(ctx, nil, link) + _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, 2, validateCalls, "token should have been attempted to be validated more than once") }) @@ -645,6 +778,9 @@ func TestRefreshToken(t *testing.T) { t.Run("ValidateNoUpdate", func(t *testing.T) { t.Parallel() + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + validateCalls := 0 fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ @@ -662,9 +798,15 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).AnyTimes() + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) - _, err := config.RefreshToken(ctx, nil, link) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, 1, validateCalls, "token is validated") }) @@ -691,24 +833,25 @@ func TestRefreshToken(t *testing.T) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - // Force a refresh - link.OAuthExpiry = expired - + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) updated, err := config.RefreshToken(ctx, db, link) require.NoError(t, err) require.Equal(t, 1, validateCalls, "token is validated") require.Equal(t, 1, refreshCalls, "token is refreshed") require.NotEqualf(t, link.OAuthAccessToken, updated.OAuthAccessToken, "token is updated") - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + dbLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ ProviderID: link.ProviderID, UserID: link.UserID, }) require.NoError(t, err) require.Equal(t, updated.OAuthAccessToken, dbLink.OAuthAccessToken, "token is updated in the DB") }) + t.Run("WithExtra", func(t *testing.T) { t.Parallel() @@ -727,11 +870,12 @@ func TestRefreshToken(t *testing.T) { cfg.ValidateURL = "" }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - // Force a refresh - link.OAuthExpiry = expired + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) updated, err := config.RefreshToken(ctx, db, link) require.NoError(t, err) @@ -777,16 +921,16 @@ func TestRefreshToken(t *testing.T) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) oldAccessToken := link.OAuthAccessToken oldRefreshToken := link.OAuthRefreshToken - // Expire the token to force a refresh. - link.OAuthExpiry = expired - // First call: refresh succeeds, validation fails (403). _, err := config.RefreshToken(ctx, db, link) require.Error(t, err, "expected error because validation returned 403") @@ -795,7 +939,7 @@ func TestRefreshToken(t *testing.T) { // Critical assertion: the DB must contain the NEW tokens from the // successful refresh, not the old (now-stale) ones. - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + dbLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ ProviderID: link.ProviderID, UserID: link.UserID, }) @@ -827,7 +971,8 @@ func TestRefreshToken(t *testing.T) { db, _ := dbtestutil.NewDB(t) var refreshCalls atomic.Int64 - cancelOnRefresh, cancel := context.WithCancel(context.Background()) + ctx := testutil.Context(t, testutil.WaitLong) + cancelOnRefresh, cancel := context.WithCancel(ctx) defer cancel() fake, config, link := setupOauth2Test(t, testConfig{ @@ -847,20 +992,20 @@ func TestRefreshToken(t *testing.T) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(cancelOnRefresh, fake.HTTPClient(nil)) - oldAccessToken := link.OAuthAccessToken oldRefreshToken := link.OAuthRefreshToken - link.OAuthExpiry = expired - _, err := config.RefreshToken(ctx, db, link) + octx := oidc.ClientContext(cancelOnRefresh, fake.HTTPClient(nil)) + _, err := config.RefreshToken(octx, db, link) require.ErrorIs(t, err, context.Canceled) - require.Equal(t, int64(1), refreshCalls.Load()) require.Eventually(t, func() bool { - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + dbLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ ProviderID: link.ProviderID, UserID: link.UserID, }) @@ -871,6 +1016,8 @@ func TestRefreshToken(t *testing.T) { dbLink.OAuthAccessToken != oldAccessToken && dbLink.OAuthRefreshToken != oldRefreshToken }, testutil.WaitShort, testutil.IntervalFast, "never saw refresh token db updated") + + require.Equal(t, int64(1), refreshCalls.Load()) }) // SaveBeforeValidate_RateLimited tests the full path: refresh @@ -905,20 +1052,20 @@ func TestRefreshToken(t *testing.T) { cfg.ValidateURL = rateLimitValidate.URL }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) // Use a real HTTP transport for non-IDP requests so the // validate request can reach the httptest server. - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(&http.Client{ + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(&http.Client{ Transport: http.DefaultTransport, })) oldAccessToken := link.OAuthAccessToken oldRefreshToken := link.OAuthRefreshToken - // Expire the token to force a refresh. - link.OAuthExpiry = expired - // RefreshToken should succeed: the IDP refresh works, the // early save persists the token, and ValidateToken returns // (true, nil, nil) because the 403 has rate-limit headers. @@ -929,7 +1076,7 @@ func TestRefreshToken(t *testing.T) { "returned token should be the new one from the refresh") // Verify the DB has the new token. - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + dbLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ ProviderID: link.ProviderID, UserID: link.UserID, }) @@ -960,10 +1107,18 @@ func TestRefreshToken(t *testing.T) { ExternalAuthOpt: func(cfg *externalauth.Config) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) - link.OAuthExpiry = expired + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) mDB.EXPECT(). UpdateExternalAuthLink(gomock.Any(), gomock.Any()). @@ -976,18 +1131,67 @@ func TestRefreshToken(t *testing.T) { "DB errors should not be treated as invalid token") }) + t.Run("WaitsIfLeased", func(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + return xerrors.New("should not be called") + }), + oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { + return jwt.MapClaims{}, nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) { + cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() + // Faster polling for faster tests. + cfg.RefreshLeaseInitialBackoff = time.Millisecond + cfg.RefreshLeaseMaxBackoff = time.Millisecond + }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, + }) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + refreshed := link + refreshed.OAuthAccessToken = "winner-access-token" + refreshed.OAuthRefreshToken = "winner-refresh-token" + + // Simulate another replica already having a lease the first time, then + // resolve with an updated link the second time. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(database.ExternalAuthLink{}, sql.ErrNoRows).Times(1) + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(refreshed, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + + updated, err := config.RefreshToken(ctx, mDB, link) + require.NoError(t, err) + require.Equal(t, updated.OAuthAccessToken, "winner-access-token") + require.Equal(t, updated.OAuthRefreshToken, "winner-refresh-token") + }) + // OptimisticLockPreventsStaleOverwrite verifies that the - // UpdateExternalAuthLinkRefreshToken WHERE clause prevents a - // stale caller from overwriting a valid refresh token saved - // by a concurrent winner. + // UpdateExternalAuthLink WHERE clause prevents a stale caller from + // overwriting a valid refresh token saved by a concurrent winner. t.Run("OptimisticLockPreventsStaleOverwrite", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) + wait := make(chan struct{}) fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { + wait <- struct{}{} + <-wait return nil }), oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { @@ -998,39 +1202,57 @@ func TestRefreshToken(t *testing.T) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() }, DB: db, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(context.Background(), fake.HTTPClient(nil)) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) // Snapshot the original tokens before any refresh. oldRefreshToken := link.OAuthRefreshToken - // Expire the token to force a refresh. - link.OAuthExpiry = expired + var ( + updated database.ExternalAuthLink + err error + ) - // Caller A: refresh and save successfully. - updated, err := config.RefreshToken(ctx, db, link) - require.NoError(t, err) - require.NotEqual(t, oldRefreshToken, updated.OAuthRefreshToken, - "caller A should have a new refresh token") - - // Caller B had a stale read of the original link. It tries to - // destroy the refresh token using the OLD refresh token in the - // optimistic lock. Because caller A already wrote a different - // refresh token, this WHERE clause matches nothing. - err = db.UpdateExternalAuthLinkRefreshToken(ctx, database.UpdateExternalAuthLinkRefreshTokenParams{ + // Caller A begins a refresh and takes the lease. + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + updated, err = config.RefreshToken(ctx, db, link) + assert.NoError(t, err) + assert.NotEqual(t, oldRefreshToken, updated.OAuthRefreshToken, + "caller A should have a new refresh token") + }() + + // Once caller A has the lease, simulate a caller B trying to update the + // link with an error. It should fail because caller A has the lease. + <-wait + _, err = db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + // Write the error an 8lear the token. OauthRefreshFailureReason: "simulated failure from stale caller B", OAuthRefreshToken: "", - OAuthRefreshTokenKeyID: "", UpdatedAt: dbtime.Now(), - ProviderID: link.ProviderID, - UserID: link.UserID, - OldOauthRefreshToken: oldRefreshToken, + // This shou* prevent the write because it does not match. + RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now().Add(time.Hour), Valid: true}, + // Preserve then. + OAuthAccessToken: link.OAuthAccessToken, + OAuthExpiry: link.OAuthExpiry, + OAuthExtra: link.OAuthExtra, }) - require.NoError(t, err, "optimistic lock write should not error, it is a no-op") + require.ErrorIs(t, err, sql.ErrNoRows) + + // Let caller A finish. + close(wait) + wg.Wait() // Verify DB still has caller A's valid token. - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + dbLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ ProviderID: link.ProviderID, UserID: link.UserID, }) @@ -1091,13 +1313,18 @@ func TestRefreshTokenWithScopes(t *testing.T) { expired := dbtime.Now().Add(-time.Hour) - // mockDBPassthrough returns a mock store that echoes the - // UpdateExternalAuthLink params back as a populated ExternalAuthLink, - // letting the test read what RefreshToken decided to persist. - mockDBPassthrough := func(t *testing.T) database.Store { + // mockDBPassthrough returns a mock store that uses the provided link for the + // initial fetch then echoes the UpdateExternalAuthLink params back as a + // populated ExternalAuthLink, letting the test read what RefreshToken decided + // to persist. + mockDBPassthrough := func(t *testing.T, link database.ExternalAuthLink) database.Store { t.Helper() ctrl := gomock.NewController(t) mDB := dbmock.NewMockStore(ctrl) + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). DoAndReturn(func(_ context.Context, p database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { return database.ExternalAuthLink{ @@ -1117,12 +1344,13 @@ func TestRefreshTokenWithScopes(t *testing.T) { []byte(`{"access_token":"new","refresh_token":"new-r","token_type":"bearer","expires_in":3600}`)) cfg := newConfig(t, []string{"openid", "offline_access", "api://app/session:role-any"}) - ctx := context.WithValue(context.Background(), oauth2.HTTPClient, client) - _, err := cfg.RefreshToken(ctx, mockDBPassthrough(t), database.ExternalAuthLink{ + ctx := context.WithValue(testutil.Context(t, testutil.WaitLong), oauth2.HTTPClient, client) + link := database.ExternalAuthLink{ OAuthAccessToken: "old", OAuthRefreshToken: "old-r", OAuthExpiry: expired, - }) + } + _, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) require.NoError(t, err) require.Equal(t, "refresh_token", captured.Get("grant_type")) @@ -1137,12 +1365,13 @@ func TestRefreshTokenWithScopes(t *testing.T) { []byte(`{"access_token":"new","refresh_token":"new-r","token_type":"bearer","expires_in":3600}`)) cfg := newConfig(t, nil) - ctx := context.WithValue(context.Background(), oauth2.HTTPClient, client) - _, err := cfg.RefreshToken(ctx, mockDBPassthrough(t), database.ExternalAuthLink{ + ctx := context.WithValue(testutil.Context(t, testutil.WaitLong), oauth2.HTTPClient, client) + link := database.ExternalAuthLink{ OAuthAccessToken: "old", OAuthRefreshToken: "old-r", OAuthExpiry: expired, - }) + } + _, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) require.NoError(t, err) require.Equal(t, "refresh_token", captured.Get("grant_type")) @@ -1158,12 +1387,13 @@ func TestRefreshTokenWithScopes(t *testing.T) { []byte(`{"access_token":"new","token_type":"bearer","expires_in":3600}`)) cfg := newConfig(t, nil) - ctx := context.WithValue(context.Background(), oauth2.HTTPClient, client) - link, err := cfg.RefreshToken(ctx, mockDBPassthrough(t), database.ExternalAuthLink{ + ctx := context.WithValue(testutil.Context(t, testutil.WaitLong), oauth2.HTTPClient, client) + link := database.ExternalAuthLink{ OAuthAccessToken: "old", OAuthRefreshToken: "prior-r", OAuthExpiry: expired, - }) + } + link, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) require.NoError(t, err) require.Equal(t, "prior-r", link.OAuthRefreshToken, "prior refresh_token must be preserved when AS omits a new one (RFC 6749 §6)") @@ -1175,12 +1405,13 @@ func TestRefreshTokenWithScopes(t *testing.T) { []byte(`{"access_token":"new","refresh_token":"rotated-r","token_type":"bearer","expires_in":3600}`)) cfg := newConfig(t, nil) - ctx := context.WithValue(context.Background(), oauth2.HTTPClient, client) - link, err := cfg.RefreshToken(ctx, mockDBPassthrough(t), database.ExternalAuthLink{ + ctx := context.WithValue(testutil.Context(t, testutil.WaitLong), oauth2.HTTPClient, client) + link := database.ExternalAuthLink{ OAuthAccessToken: "old", OAuthRefreshToken: "prior-r", OAuthExpiry: expired, - }) + } + link, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) require.NoError(t, err) require.Equal(t, "rotated-r", link.OAuthRefreshToken, "rotated refresh_token from AS must be persisted") @@ -1272,7 +1503,7 @@ func TestValidateToken(t *testing.T) { t.Helper() tp := &http.Transport{} t.Cleanup(tp.CloseIdleConnections) - return oidc.ClientContext(context.Background(), &http.Client{Transport: tp}) + return oidc.ClientContext(testutil.Context(t, testutil.WaitLong), &http.Client{Transport: tp}) } // RateLimitRemaining: 403 with X-RateLimit-Remaining: 0 should be @@ -1634,7 +1865,8 @@ func TestExchangeWithClientSecret(t *testing.T) { }), } - _, err = config.Exchange(context.WithValue(context.Background(), oauth2.HTTPClient, client), "code") + ctx := testutil.Context(t, testutil.WaitLong) + _, err = config.Exchange(context.WithValue(ctx, oauth2.HTTPClient, client), "code") require.NoError(t, err) } @@ -1869,6 +2101,9 @@ type testConfig struct { ExternalAuthOpt func(cfg *externalauth.Config) // If DB is passed in, the link will be inserted into the DB. DB database.Store + // ExternalAuthLinkOpts can be used to manipulate the link inserted into the + // DB when the DB is provided. + ExternalAuthLinkOpts func(link *database.ExternalAuthLink) } // setupTest will configure a fake IDP and a externalauth.Config for testing. @@ -1920,6 +2155,9 @@ func setupOauth2Test(t *testing.T, settings testConfig) (*oidctest.FakeIDP, *ext // The caller can manually expire this if they want. OAuthExpiry: now.Add(time.Hour), } + if settings.ExternalAuthLinkOpts != nil { + settings.ExternalAuthLinkOpts(&link) + } if settings.DB != nil { // Feel free to insert additional things like the user, etc if required. diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index c5eb376c8d3..fb28ad2ecd1 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -522,7 +522,7 @@ func TestAcquireJob(t *testing.T) { OAuthExpiry: dbtime.Now().Add(time.Hour), OAuthAccessToken: "access-token", }) - dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{ + ealink := dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{ ProviderID: gitAuthProvider.Id, UserID: user.ID, }) @@ -768,7 +768,7 @@ func TestAcquireJob(t *testing.T) { }, ExternalAuthProviders: []*sdkproto.ExternalAuthProvider{{ Id: gitAuthProvider.Id, - AccessToken: "access_token", + AccessToken: ealink.OAuthAccessToken, }}, Metadata: wantedMetadata, }, diff --git a/enterprise/dbcrypt/cliutil.go b/enterprise/dbcrypt/cliutil.go index 086c9ba6b4f..a6656356ccc 100644 --- a/enterprise/dbcrypt/cliutil.go +++ b/enterprise/dbcrypt/cliutil.go @@ -60,15 +60,17 @@ func Rotate(ctx context.Context, log slog.Logger, sqlDB *sql.DB, ciphers []Ciphe continue } if _, err := cryptTx.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: uid, - UpdatedAt: externalAuthLink.UpdatedAt, - OAuthAccessToken: externalAuthLink.OAuthAccessToken, - OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthRefreshToken: externalAuthLink.OAuthRefreshToken, - OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthExpiry: externalAuthLink.OAuthExpiry, - OAuthExtra: externalAuthLink.OAuthExtra, + ProviderID: externalAuthLink.ProviderID, + UserID: uid, + UpdatedAt: externalAuthLink.UpdatedAt, + OAuthAccessToken: externalAuthLink.OAuthAccessToken, + OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthRefreshToken: externalAuthLink.OAuthRefreshToken, + OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthExpiry: externalAuthLink.OAuthExpiry, + OAuthExtra: externalAuthLink.OAuthExtra, + OauthRefreshFailureReason: "", + RefreshLeaseExpiresAt: sql.NullTime{}, }); err != nil { return xerrors.Errorf("update external auth link user_id=%s provider_id=%s: %w", externalAuthLink.UserID, externalAuthLink.ProviderID, err) } @@ -274,15 +276,17 @@ func Decrypt(ctx context.Context, log slog.Logger, sqlDB *sql.DB, ciphers []Ciph continue } if _, err := tx.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: uid, - UpdatedAt: externalAuthLink.UpdatedAt, - OAuthAccessToken: externalAuthLink.OAuthAccessToken, - OAuthAccessTokenKeyID: sql.NullString{}, // we explicitly want to clear the key id - OAuthRefreshToken: externalAuthLink.OAuthRefreshToken, - OAuthRefreshTokenKeyID: sql.NullString{}, // we explicitly want to clear the key id - OAuthExpiry: externalAuthLink.OAuthExpiry, - OAuthExtra: externalAuthLink.OAuthExtra, + ProviderID: externalAuthLink.ProviderID, + UserID: uid, + UpdatedAt: externalAuthLink.UpdatedAt, + OAuthAccessToken: externalAuthLink.OAuthAccessToken, + OAuthAccessTokenKeyID: sql.NullString{}, // we explicitly want to clear the key id + OAuthRefreshToken: externalAuthLink.OAuthRefreshToken, + OAuthRefreshTokenKeyID: sql.NullString{}, // we explicitly want to clear the key id + OAuthExpiry: externalAuthLink.OAuthExpiry, + OAuthExtra: externalAuthLink.OAuthExtra, + OauthRefreshFailureReason: "", + RefreshLeaseExpiresAt: sql.NullTime{}, }); err != nil { return xerrors.Errorf("update external auth link user_id=%s provider_id=%s: %w", externalAuthLink.UserID, externalAuthLink.ProviderID, err) } diff --git a/enterprise/dbcrypt/dbcrypt.go b/enterprise/dbcrypt/dbcrypt.go index 6c9150f17a3..b871204e244 100644 --- a/enterprise/dbcrypt/dbcrypt.go +++ b/enterprise/dbcrypt/dbcrypt.go @@ -242,6 +242,20 @@ func (db *dbCrypt) GetExternalAuthLinksByUserID(ctx context.Context, userID uuid return links, nil } +func (db *dbCrypt) AcquireExternalAuthLinkRefreshLease(ctx context.Context, params database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + link, err := db.Store.AcquireExternalAuthLinkRefreshLease(ctx, params) + if err != nil { + return database.ExternalAuthLink{}, err + } + if err := db.decryptField(&link.OAuthAccessToken, link.OAuthAccessTokenKeyID); err != nil { + return database.ExternalAuthLink{}, err + } + if err := db.decryptField(&link.OAuthRefreshToken, link.OAuthRefreshTokenKeyID); err != nil { + return database.ExternalAuthLink{}, err + } + return link, nil +} + func (db *dbCrypt) UpdateExternalAuthLink(ctx context.Context, params database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { if err := db.encryptField(¶ms.OAuthAccessToken, ¶ms.OAuthAccessTokenKeyID); err != nil { return database.ExternalAuthLink{}, err @@ -262,54 +276,6 @@ func (db *dbCrypt) UpdateExternalAuthLink(ctx context.Context, params database.U return link, nil } -func (db *dbCrypt) UpdateExternalAuthLinkRefreshToken(ctx context.Context, params database.UpdateExternalAuthLinkRefreshTokenParams) error { - // The SQL query uses an optimistic lock: - // WHERE oauth_refresh_token = @old_oauth_refresh_token - // The caller supplies the plaintext old token (since dbcrypt - // decrypts on read), but the DB stores the encrypted value. - // Because AES-GCM is non-deterministic, we cannot simply - // re-encrypt the old token — the ciphertext would differ. - // Instead, read the current row from the inner (raw) store - // and use the actual encrypted value for the WHERE clause. - if params.OldOauthRefreshToken != "" && db.ciphers != nil && db.primaryCipherDigest != "" { - raw, err := db.Store.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ - ProviderID: params.ProviderID, - UserID: params.UserID, - }) - if err != nil { - return err - } - // Decrypt the stored token so we can compare with the - // caller-supplied plaintext. - decrypted := raw.OAuthRefreshToken - if err := db.decryptField(&decrypted, raw.OAuthRefreshTokenKeyID); err != nil { - return err - } - if decrypted != params.OldOauthRefreshToken { - // The token has changed since the caller read it; - // the optimistic lock should fail (no rows updated). - // Return nil to match the :exec semantics of the SQL - // query, which silently updates zero rows. - return nil - } - // Use the raw encrypted value so the WHERE clause matches. - params.OldOauthRefreshToken = raw.OAuthRefreshToken - } - - // We would normally use a sql.NullString here, but sqlc does not want to make - // a params struct with a nullable string. - var digest sql.NullString - if params.OAuthRefreshTokenKeyID != "" { - digest.String = params.OAuthRefreshTokenKeyID - digest.Valid = true - } - if err := db.encryptField(¶ms.OAuthRefreshToken, &digest); err != nil { - return err - } - - return db.Store.UpdateExternalAuthLinkRefreshToken(ctx, params) -} - func (db *dbCrypt) GetCryptoKeys(ctx context.Context) ([]database.CryptoKey, error) { keys, err := db.Store.GetCryptoKeys(ctx) if err != nil { diff --git a/enterprise/dbcrypt/dbcrypt_internal_test.go b/enterprise/dbcrypt/dbcrypt_internal_test.go index d37fbbacf88..2721973fccd 100644 --- a/enterprise/dbcrypt/dbcrypt_internal_test.go +++ b/enterprise/dbcrypt/dbcrypt_internal_test.go @@ -98,32 +98,6 @@ func TestUserLinks(t *testing.T) { require.EqualValues(t, expectedClaims, rawLink.Claims) }) - t.Run("UpdateExternalAuthLinkRefreshToken", func(t *testing.T) { - t.Parallel() - db, crypt, ciphers := setup(t) - user := dbgen.User(t, crypt, database.User{}) - link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ - UserID: user.ID, - }) - - err := crypt.UpdateExternalAuthLinkRefreshToken(ctx, database.UpdateExternalAuthLinkRefreshTokenParams{ - OAuthRefreshToken: "", - OAuthRefreshTokenKeyID: link.OAuthRefreshTokenKeyID.String, - OldOauthRefreshToken: link.OAuthRefreshToken, - UpdatedAt: dbtime.Now(), - ProviderID: link.ProviderID, - UserID: link.UserID, - }) - require.NoError(t, err) - - rawLink, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ - ProviderID: link.ProviderID, - UserID: link.UserID, - }) - require.NoError(t, err) - requireEncryptedEquals(t, ciphers[0], rawLink.OAuthRefreshToken, "") - }) - t.Run("GetUserLinkByLinkedID", func(t *testing.T) { t.Parallel() t.Run("OK", func(t *testing.T) { @@ -358,6 +332,61 @@ func TestExternalAuthLinks(t *testing.T) { }) }) + t.Run("AcquireExternalAuthLinkRefreshLease", func(t *testing.T) { + t.Run("OK", func(t *testing.T) { + t.Parallel() + db, crypt, ciphers := setup(t) + link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ + OAuthAccessToken: "access", + OAuthRefreshToken: "refresh", + }) + link, err := db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + UserID: link.UserID, + ProviderID: link.ProviderID, + RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now(), Valid: true}, + }) + require.NoError(t, err) + requireEncryptedEquals(t, ciphers[0], link.OAuthAccessToken, "access") + requireEncryptedEquals(t, ciphers[0], link.OAuthRefreshToken, "refresh") + }) + t.Run("Decrypt", func(t *testing.T) { + t.Parallel() + _, crypt, ciphers := setup(t) + link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ + OAuthAccessToken: "access", + OAuthRefreshToken: "refresh", + }) + link, err := crypt.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + UserID: link.UserID, + ProviderID: link.ProviderID, + RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now(), Valid: true}, + }) + require.NoError(t, err) + require.Equal(t, "access", link.OAuthAccessToken) + require.Equal(t, "refresh", link.OAuthRefreshToken) + require.Equal(t, ciphers[0].HexDigest(), link.OAuthAccessTokenKeyID.String) + require.Equal(t, ciphers[0].HexDigest(), link.OAuthRefreshTokenKeyID.String) + }) + t.Run("DecryptErr", func(t *testing.T) { + t.Parallel() + db, crypt, ciphers := setup(t) + link := dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{ + OAuthAccessToken: fakeBase64RandomData(t, 32), + OAuthRefreshToken: fakeBase64RandomData(t, 32), + OAuthAccessTokenKeyID: sql.NullString{String: ciphers[0].HexDigest(), Valid: true}, + OAuthRefreshTokenKeyID: sql.NullString{String: ciphers[0].HexDigest(), Valid: true}, + }) + link, err := crypt.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + UserID: link.UserID, + ProviderID: link.ProviderID, + RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now(), Valid: true}, + }) + require.Error(t, err, "expected an error") + var derr *DecryptFailedError + require.ErrorAs(t, err, &derr, "expected a decrypt error") + }) + }) + t.Run("GetExternalAuthLinksByUserID", func(t *testing.T) { t.Parallel() From 3fae32327031e8df46ae92df582f905f2ea7cfa4 Mon Sep 17 00:00:00 2001 From: Asher Date: Wed, 12 Aug 2026 10:50:53 -0800 Subject: [PATCH 02/12] Break validation out of the refresh lease If we have to refresh, we do still validate while holding the lease, since we have it anyway. But if we do not need to refresh, validate without getting a lease. The assumption is that the overhead to get the lease is greater than the overhead of concurrent validations. --- coderd/externalauth/externalauth.go | 282 ++++++++++++----------- coderd/externalauth/externalauth_test.go | 26 +-- 2 files changed, 155 insertions(+), 153 deletions(-) diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index fe63781b946..b2638c6b603 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -223,20 +223,11 @@ func IsInvalidTokenError(err error) bool { } // RefreshToken automatically refreshes the token if expired and permitted. +// Tokens are then validated, whether or not they were refreshed. func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink) (database.ExternalAuthLink, error) { - // If the token is not expired and there is no validation URL, we can - // short-circuit since there is nothing to do. - token := &oauth2.Token{ - AccessToken: externalAuthLink.OAuthAccessToken, - RefreshToken: externalAuthLink.OAuthRefreshToken, - Expiry: externalAuthLink.OAuthExpiry, - } - if c.ValidateURL == "" && token.Valid() { - return externalAuthLink, nil - } - - // If the token is expired and refresh is disabled, we prompt the user to + // If the token is expired and refresh is disabled, prompt the user to // authenticate manually again. + token := externalAuthLink.OAuthToken() if c.NoRefresh && !token.Valid() { return externalAuthLink, InvalidTokenError("token expired and refreshing is disabled") } @@ -245,7 +236,7 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu // in-flight refresh. Otherwise, the parallel calls will fail with a bad // refresh token error as they can only be used once. key := c.ID + ":" + externalAuthLink.UserID.String() - ch := c.RefreshGroup.DoChan(key, func() (newLink any, refreshErr error) { + ch := c.RefreshGroup.DoChan(key, func() (any, error) { // Use a detached context so if a request is canceled or times out it does // not cancel all the other requests as well. The deadline is arbitrary but // we give at least enough time for the refresh timeout then another 10 @@ -257,61 +248,21 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), timeout) defer cancel() - lease := dbtime.Now().Add(timeout) - - // There may be other replicas also wanting to refresh; try to get a lease - // on the row. This also ensures we have the latest link. - var dblink database.ExternalAuthLink - initial := defaultRefreshLeaseInitialBackoff - if c.RefreshLeaseInitialBackoff > 0 { - initial = c.RefreshLeaseInitialBackoff - } - maximum := defaultRefreshLeaseMaxBackoff - if c.RefreshLeaseMaxBackoff > 0 { - maximum = c.RefreshLeaseMaxBackoff - } - r := retry.New(initial, maximum) - for { - var err error - dblink, err = db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, - }) - switch { - case errors.Is(err, sql.ErrNoRows): - // Something still holds the lock; keep waiting. - if !r.Wait(ctx) { - return externalAuthLink, ctx.Err() - } - case err != nil: - // Some kind of DB error. - return externalAuthLink, err - default: - // We now hold the lock. - goto refresh - } + // Although TokenSource().Token() will also check if the token is expired, + // do so ahead of time here to avoid the overhead of acquiring the lease + // when no refresh is required. Also validate the token under this lock, + // since we already have it anyway. + if !token.Valid() { + return c.refreshAndValidateWithLease(ctx, db, externalAuthLink, timeout) } - refresh: - defer func() { - refreshErr = errors.Join(refreshErr, db.ReleaseExternalAuthLinkRefreshLease(ctx, database.ReleaseExternalAuthLinkRefreshLeaseParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, - })) - }() - - // If we got a different token, it means something else refreshed either - // while we were waiting or before we got a hold of the lease but after we - // initially fetched the link. - if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { - return dblink, nil - } - - // Otherwise the token has still not been updated; refresh it now. - newLink, refreshErr = c.innerRefreshToken(ctx, db, dblink, lease) - return newLink, refreshErr + // Validate the token even if we did not refresh. This is done within the + // group but outside the lease, meaning multiple instances may validate at + // the same time, but the lock overhead seems greater than the overhead of + // occasional concurrent validation, and is not strictly necessary like it + // is with refreshing, so avoid the lock when not refreshing. + _, err := c.validateWithRetry(ctx, token) + return externalAuthLink, err }) select { case results := <-ch: @@ -326,13 +277,72 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu } } -func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, lease time.Time) (database.ExternalAuthLink, error) { - existingToken := &oauth2.Token{ - AccessToken: externalAuthLink.OAuthAccessToken, - RefreshToken: externalAuthLink.OAuthRefreshToken, - Expiry: externalAuthLink.OAuthExpiry, +// refreshAndValidateWithLease wraps the refresh and subsequent validation with +// concurrency protection between multiple instances by using a lease column on +// the link's row. +func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, timeout time.Duration) (newLink database.ExternalAuthLink, refreshErr error) { + lease := dbtime.Now().Add(timeout) + + // There may be other replicas also wanting to refresh; try to get a lease + // on the row. This also ensures we have the latest link. + var dblink database.ExternalAuthLink + initial := defaultRefreshLeaseInitialBackoff + if c.RefreshLeaseInitialBackoff > 0 { + initial = c.RefreshLeaseInitialBackoff + } + maximum := defaultRefreshLeaseMaxBackoff + if c.RefreshLeaseMaxBackoff > 0 { + maximum = c.RefreshLeaseMaxBackoff + } + r := retry.New(initial, maximum) + for { + var err error + dblink, err = db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + }) + switch { + case errors.Is(err, sql.ErrNoRows): + // Something still holds the lock; keep waiting. + if !r.Wait(ctx) { + return externalAuthLink, ctx.Err() + } + case err != nil: + // Some kind of DB error. + return externalAuthLink, err + default: + // We now hold the lock. + goto refresh + } + } + +refresh: + defer func() { + refreshErr = errors.Join(refreshErr, db.ReleaseExternalAuthLinkRefreshLease(ctx, database.ReleaseExternalAuthLinkRefreshLeaseParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + })) + }() + + // If we got a different token, it means something else refreshed either + // while we were waiting or before we got a hold of the lease but after we + // initially fetched the link. + if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { + return dblink, nil } + // Otherwise the token has still not been updated; refresh it now. + newLink, refreshErr = c.refreshAndValidateToken(ctx, db, dblink, lease) + return newLink, refreshErr +} + +// refreshAndValidateToken does the actual token refresh, persists the result to +// the database, then validates the token. +func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, lease time.Time) (database.ExternalAuthLink, error) { + existingToken := externalAuthLink.OAuthToken() + // This is additional defensive programming. Because TokenSource is an // interface, we cannot be sure that the implementation will treat an 'IsZero' // time as "not-expired". The default implementation does, but a custom @@ -426,78 +436,53 @@ func (c *Config) innerRefreshToken(ctx context.Context, db database.Store, exter return externalAuthLink, xerrors.Errorf("generate token extra: %w", err) } - // Persist the refreshed token to the DB before validation. GitHub - // rotates refresh tokens on every use, so the old refresh token is - // already invalid on the IDP side. If we validated first and the - // validation endpoint was unavailable (e.g. rate-limited 403), the - // new token would be silently lost and the user would be forced to - // re-authenticate manually. - originalAccessToken := externalAuthLink.OAuthAccessToken - if token.AccessToken != originalAccessToken { - updatedAuthLink, err := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - UpdatedAt: dbtime.Now(), - OAuthAccessToken: token.AccessToken, - OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthRefreshToken: token.RefreshToken, - OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthExpiry: token.Expiry, - OAuthExtra: extra, - // If there was any failure before, we can clear it now. - OauthRefreshFailureReason: "", - // The row will only update if we hold the current lease. This is - // somewhat redundant since if we lost the lease our context would be - // expired anyway, so it is not actually possible to get the - // sql.ErrNoRows that would result from this. - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, - }) - if err != nil { - return updatedAuthLink, xerrors.Errorf("persist refreshed token: %w", err) - } - externalAuthLink = updatedAuthLink + // Persist the refreshed token to the DB before validation. GitHub rotates + // refresh tokens on every use, so the old refresh token is already invalid on + // the IDP side. If we validated first and the validation endpoint was + // unavailable (e.g. rate-limited 403), the new token would be silently lost + // and the user would be forced to re-authenticate manually. + updatedAuthLink, err := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + UpdatedAt: dbtime.Now(), + OAuthAccessToken: token.AccessToken, + OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthRefreshToken: token.RefreshToken, + OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required + OAuthExpiry: token.Expiry, + OAuthExtra: extra, + // If there was any failure before, we can clear it now. + OauthRefreshFailureReason: "", + // The row will only update if we hold the current lease. This is somewhat + // redundant since if we lost the lease our context would be expired anyway, + // so it is not actually possible to get the sql.ErrNoRows that would result + // from this. + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + }) + if err != nil { + return updatedAuthLink, xerrors.Errorf("persist refreshed token: %w", err) } - r := retry.New(50*time.Millisecond, 200*time.Millisecond) - // See the comment below why the retry and cancel is required. - retryCtx, retryCtxCancel := context.WithTimeout(ctx, time.Second) - defer retryCtxCancel() -validate: - valid, user, err := c.ValidateToken(ctx, token) + user, err := c.validateWithRetry(ctx, token) if err != nil { - return externalAuthLink, xerrors.Errorf("validate external auth token: %w", err) - } - if !valid { - // A customer using GitHub in Australia reported that validating immediately - // after refreshing the token would intermittently fail with a 401. Waiting - // a few milliseconds with the exact same token on the exact same request - // would resolve the issue. It seems likely that the write is not propagating - // to the read replica in time. - // - // We do an exponential backoff here to give the write time to propagate. - if c.Type == string(codersdk.EnhancedExternalAuthProviderGitHub) && r.Wait(retryCtx) { - goto validate - } - // The token is no longer valid! - return externalAuthLink, InvalidTokenError("token failed to validate") + return updatedAuthLink, err } - // Update the associated user's github.com user ID if the token - // is for github.com and validation returned user info. - if token.AccessToken != originalAccessToken && IsGithubDotComURL(c.AuthCodeURL("")) && user != nil { + // Update the associated user's github.com user ID if the token is for + // github.com and validation returned user info. + if IsGithubDotComURL(c.AuthCodeURL("")) && user != nil { err = db.UpdateUserGithubComUserID(ctx, database.UpdateUserGithubComUserIDParams{ - ID: externalAuthLink.UserID, + ID: updatedAuthLink.UserID, GithubComUserID: sql.NullInt64{ Int64: user.ID, Valid: true, }, }) if err != nil { - return externalAuthLink, xerrors.Errorf("update user github com user id: %w", err) + return updatedAuthLink, xerrors.Errorf("update user github com user id: %w", err) } } - - return externalAuthLink, nil + return updatedAuthLink, nil } // refreshTokenWithRetry exchanges the refresh token for a new access token, @@ -555,17 +540,48 @@ func (c *Config) refreshTokenWithRetry(ctx context.Context, existingToken *oauth } } +// validateWithRetry validates the provided link, retrying on failure. On +// success return the user info. +func (c *Config) validateWithRetry(ctx context.Context, token *oauth2.Token) (*codersdk.ExternalAuthUser, error) { + r := retry.New(50*time.Millisecond, 200*time.Millisecond) + // See the comment below why the retry and cancel is required. + retryCtx, retryCtxCancel := context.WithTimeout(ctx, time.Second) + defer retryCtxCancel() + +validate: + valid, user, err := c.ValidateToken(ctx, token) + if err != nil { + return nil, xerrors.Errorf("validate external auth token: %w", err) + } + if !valid { + // A customer using GitHub in Australia reported that validating immediately + // after refreshing the token would intermittently fail with a 401. Waiting + // a few milliseconds with the exact same token on the exact same request + // would resolve the issue. It seems likely that the write is not propagating + // to the read replica in time. + // + // We do an exponential backoff here to give the write time to propagate. + if c.Type == string(codersdk.EnhancedExternalAuthProviderGitHub) && r.Wait(retryCtx) { + goto validate + } + // The token is no longer valid! + return nil, InvalidTokenError("token failed to validate") + } + + return user, nil +} + // ValidateToken checks if the Git token provided is valid. // The user is optionally returned if the provider supports it. // Returns valid=true when: the provider confirmed the token, // no ValidateURL is configured, or the validation endpoint // returned a rate-limited response (403 with rate-limit headers // or 429). -func (c *Config) ValidateToken(ctx context.Context, link *oauth2.Token) (bool, *codersdk.ExternalAuthUser, error) { - if link == nil { +func (c *Config) ValidateToken(ctx context.Context, token *oauth2.Token) (bool, *codersdk.ExternalAuthUser, error) { + if token == nil { return false, nil, xerrors.New("validate external auth token: token is nil") } - if !link.Expiry.IsZero() && link.Expiry.Before(dbtime.Now()) { + if !token.Valid() { return false, nil, nil } @@ -578,7 +594,7 @@ func (c *Config) ValidateToken(ctx context.Context, link *oauth2.Token) (bool, * return false, nil, err } - req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", link.AccessToken)) + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", token.AccessToken)) res, err := c.InstrumentedOAuth2Config.Do(ctx, promoauth.SourceValidateToken, req) if err != nil { return false, nil, err diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index 554c2a55f3c..dde56939886 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -114,12 +114,8 @@ func TestRefreshToken(t *testing.T) { // Zero time used link.OAuthExpiry = time.Time{} - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - + // Since the token is not expired, no refresh lease will be acquired and + // it will only be validated. _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.True(t, validated, "token should have been validated") @@ -762,14 +758,9 @@ func TestRefreshToken(t *testing.T) { }, }) + // Since the token is not expired, no refresh lease will be acquired and + // it will only be validated. ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, 2, validateCalls, "token should have been attempted to be validated more than once") @@ -798,14 +789,9 @@ func TestRefreshToken(t *testing.T) { }, }) - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).AnyTimes() - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - + // Since the token is not expired, no refresh lease will be acquired and + // it will only be validated. ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, 1, validateCalls, "token is validated") From abab355c78b3e0313a84fedd6ab0c4fd952cb1d6 Mon Sep 17 00:00:00 2001 From: Asher Date: Thu, 13 Aug 2026 11:10:43 -0800 Subject: [PATCH 03/12] Fix require.Equal order --- coderd/externalauth/externalauth_test.go | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index dde56939886..cbe8cf45be8 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -1145,9 +1145,10 @@ func TestRefreshToken(t *testing.T) { ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - refreshed := link - refreshed.OAuthAccessToken = "winner-access-token" - refreshed.OAuthRefreshToken = "winner-refresh-token" + refreshed := database.ExternalAuthLink{ + OAuthAccessToken: "winner-access-token", + OAuthRefreshToken: "winner-refresh-token", + } // Simulate another replica already having a lease the first time, then // resolve with an updated link the second time. @@ -1160,8 +1161,8 @@ func TestRefreshToken(t *testing.T) { updated, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) - require.Equal(t, updated.OAuthAccessToken, "winner-access-token") - require.Equal(t, updated.OAuthRefreshToken, "winner-refresh-token") + require.Equal(t, refreshed.OAuthAccessToken, updated.OAuthAccessToken) + require.Equal(t, refreshed.OAuthRefreshToken, updated.OAuthRefreshToken) }) // OptimisticLockPreventsStaleOverwrite verifies that the From c644ef6f9d9f3ccac6bc21be960682072da6b7ef Mon Sep 17 00:00:00 2001 From: Asher Date: Thu, 13 Aug 2026 11:11:24 -0800 Subject: [PATCH 04/12] Fix variable names in comments --- coderd/externalauth/externalauth.go | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index b2638c6b603..1dd37b13b3b 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -61,8 +61,8 @@ const ( // check whether another replica has finished refreshing. defaultRefreshLeaseInitialBackoff = 50 * time.Millisecond - // defaultRefreshLeaseInitialBackoff is the maximum wait between polls to - // check whether another replica has finished refreshing. + // defaultRefreshLeaseMaxBackoff is the maximum wait between polls to check + // whether another replica has finished refreshing. defaultRefreshLeaseMaxBackoff = 500 * time.Millisecond ) @@ -173,8 +173,8 @@ type Config struct { // whether another replica has finished refreshing. RefreshLeaseInitialBackoff time.Duration - // RefreshLeaseInitialBackoff is the maximum wait between polls to check - // whether another replica has finished refreshing. + // RefreshLeaseMaxBackoff is the maximum wait between polls to check whether + // another replica has finished refreshing. RefreshLeaseMaxBackoff time.Duration } From 83d3c7436893dc9c0e63ac9d39cec07c6391eea3 Mon Sep 17 00:00:00 2001 From: Asher Date: Thu, 13 Aug 2026 11:42:44 -0800 Subject: [PATCH 05/12] Handle a cached error Previously if the token changed (including getting removed) we just returned it, but it could have errored. --- coderd/externalauth/externalauth.go | 26 ++++--- coderd/externalauth/externalauth_test.go | 88 +++++++++++++++++++++++- 2 files changed, 105 insertions(+), 9 deletions(-) diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index 1dd37b13b3b..211a20d979c 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -330,6 +330,9 @@ refresh: // while we were waiting or before we got a hold of the lease but after we // initially fetched the link. if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { + if dblink.OauthRefreshFailureReason != "" { + return externalAuthLink, refreshError(dblink, dblink.OauthRefreshFailureReason) + } return dblink, nil } @@ -338,6 +341,16 @@ refresh: return newLink, refreshErr } +// refreshError converts a failure reason to an error. +func refreshError(link database.ExternalAuthLink, reason string) error { + return InvalidTokenError(fmt.Sprintf("token expired and refreshing failed %s with: %s", + // Do not return the exact time, because then we have to know what timezone + // the user is in. This approximate time is good enough. + humanize.Time(link.UpdatedAt), + reason, + )) +} + // refreshAndValidateToken does the actual token refresh, persists the result to // the database, then validates the token. func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, lease time.Time) (database.ExternalAuthLink, error) { @@ -379,7 +392,7 @@ func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, // spamming the database with long error messages. reason = reason[:failureReasonLimit] } - dblink, updateErr := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ + _, updateErr := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ ProviderID: externalAuthLink.ProviderID, UserID: externalAuthLink.UserID, UpdatedAt: dbtime.Now(), @@ -403,7 +416,9 @@ func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, // This error should be rare. return externalAuthLink, InvalidTokenError(fmt.Sprintf("refresh token failed: %q, then removing refresh token failed: %q", err.Error(), updateErr.Error())) } - externalAuthLink = dblink + // The refresh token was cleared + externalAuthLink.OAuthRefreshToken = "" + externalAuthLink.UpdatedAt = dbtime.Now() } // Unfortunately have to match exactly on the error message string. @@ -417,12 +432,7 @@ func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, if externalAuthLink.OauthRefreshFailureReason != "" { // A cached refresh failure error exists. So the refresh token was set, but was invalid, and zeroed out. // Return this cached error for the original refresh attempt. - return externalAuthLink, InvalidTokenError(fmt.Sprintf("token expired and refreshing failed %s with: %s", - // Do not return the exact time, because then we have to know what timezone the - // user is in. This approximate time is good enough. - humanize.Time(externalAuthLink.UpdatedAt), - externalAuthLink.OauthRefreshFailureReason, - )) + return externalAuthLink, refreshError(externalAuthLink, externalAuthLink.OauthRefreshFailureReason) } return externalAuthLink, InvalidTokenError("token expired, refreshing failed and will not be retried") diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index cbe8cf45be8..baf126c8717 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -625,6 +625,46 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, int64(1), refreshCalls.Load()) }) + // CachedErrorForEmptyRefreshToken tests that when the refresh token is + // missing, we refer to the cached error, since likely the token was removed in + // a previous attempt and we want to return the original error rather than + // "token is missing". + t.Run("CachedErrorForEmptyRefreshToken", func(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + + var refreshCalls atomic.Int64 + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + refreshCalls.Add(1) + return xerrors.New("should not be called") + }), + }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + link.OauthRefreshFailureReason = "test cached error" + link.OAuthRefreshToken = "" + }, + }) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + // Allow getting the lease. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(link, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + + // It should return the cached error message. + _, err := config.RefreshToken(ctx, mDB, link) + require.Equal(t, int64(0), refreshCalls.Load()) + require.Error(t, err) + require.ErrorContains(t, err, "test cached error") + }) + t.Run("ReturnsReleaseError", func(t *testing.T) { t.Parallel() @@ -1117,7 +1157,7 @@ func TestRefreshToken(t *testing.T) { "DB errors should not be treated as invalid token") }) - t.Run("WaitsIfLeased", func(t *testing.T) { + t.Run("WaitsIfLeasedOK", func(t *testing.T) { t.Parallel() ctrl := gomock.NewController(t) @@ -1165,6 +1205,52 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, refreshed.OAuthRefreshToken, updated.OAuthRefreshToken) }) + t.Run("WaitsIfLeasedError", func(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + return xerrors.New("should not be called") + }), + oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { + return jwt.MapClaims{}, nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) { + cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() + // Faster polling for faster tests. + cfg.RefreshLeaseInitialBackoff = time.Millisecond + cfg.RefreshLeaseMaxBackoff = time.Millisecond + }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, + }) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + refreshed := database.ExternalAuthLink{ + OauthRefreshFailureReason: "failed to refresh", + } + + // Simulate another replica already having a lease the first time, then + // resolve with the failure the second time. + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(database.ExternalAuthLink{}, sql.ErrNoRows).Times(1) + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(refreshed, nil).Times(1) + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). + Return(nil).Times(1) + + _, err := config.RefreshToken(ctx, mDB, link) + require.Error(t, err) + require.ErrorContains(t, err, refreshed.OauthRefreshFailureReason) + }) + // OptimisticLockPreventsStaleOverwrite verifies that the // UpdateExternalAuthLink WHERE clause prevents a stale caller from // overwriting a valid refresh token saved by a concurrent winner. From b59ef04d534be6694f1a08b0ede1d7551e26047b Mon Sep 17 00:00:00 2001 From: Asher Date: Mon, 17 Aug 2026 13:15:23 -0800 Subject: [PATCH 06/12] Take a lock before updating the lease This lets us distinguish between the row not existing and not being able to take the lease because something else has it. --- coderd/database/dbauthz/dbauthz.go | 21 +- coderd/database/dbauthz/dbauthz_test.go | 17 +- coderd/database/dbmetrics/querymetrics.go | 24 +- coderd/database/dbmock/dbmock.go | 43 +- coderd/database/querier.go | 5 +- coderd/database/queries.sql.go | 57 +-- coderd/database/queries/externalauth.sql | 18 +- coderd/externalauth/externalauth.go | 85 +++- coderd/externalauth/externalauth_test.go | 486 +++++++++++--------- enterprise/dbcrypt/dbcrypt.go | 14 - enterprise/dbcrypt/dbcrypt_internal_test.go | 55 --- 11 files changed, 374 insertions(+), 451 deletions(-) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index f68a7a1b2f3..80c92ea6f14 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -1731,13 +1731,6 @@ func scopedOrgRoleIdentifiers(names []string, orgID uuid.UUID) []rbac.RoleIdenti return out } -func (q *querier) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - fetch := func(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) - } - return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.AcquireExternalAuthLinkRefreshLease)(ctx, arg) -} - func (q *querier) AcquireLock(ctx context.Context, id int64) error { return q.db.AcquireLock(ctx, id) } @@ -7044,13 +7037,6 @@ func (q *querier) RegisterWorkspaceProxy(ctx context.Context, arg database.Regis return updateWithReturn(q.log, q.auth, fetch, q.db.RegisterWorkspaceProxy)(ctx, arg) } -func (q *querier) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { - fetch := func(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) - } - return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.ReleaseExternalAuthLinkRefreshLease)(ctx, arg) -} - func (q *querier) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { // This is a system function to clear user groups in group sync. if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil { @@ -7108,6 +7094,13 @@ func (q *querier) SetChatContextSnapshot(ctx context.Context, arg database.SetCh return q.db.SetChatContextSnapshot(ctx, arg) } +func (q *querier) SetExternalAuthLinkRefreshLease(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) error { + fetch := func(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) + } + return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.SetExternalAuthLinkRefreshLease)(ctx, arg) +} + func (q *querier) SoftDeleteChatMessageByID(ctx context.Context, id int64) error { msg, err := q.db.GetChatMessageByID(ctx, id) if err != nil { diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 975ba83a22f..9601ae98b86 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -3202,22 +3202,13 @@ func (s *MethodTestSuite) TestUser() { dbm.EXPECT().UpdateExternalAuthLink(gomock.Any(), arg).Return(link, nil).AnyTimes() check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) })) - s.Run("AcquireExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + s.Run("SetExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{}) dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Minute), Valid: true} - arg := database.AcquireExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} - dbm.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(link, nil).AnyTimes() - check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) - })) - s.Run("ReleaseExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{ - RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now().Add(time.Minute), Valid: true}, - }) - dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() - arg := database.ReleaseExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} - dbm.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(nil).AnyTimes() - check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns() + arg := database.SetExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} + dbm.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(nil).AnyTimes() + check.Args(arg).Asserts(link, policy.ActionUpdatePersonal) })) s.Run("UpdateUserLink", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { link := testutil.Fake(s.T(), faker, database.UserLink{}) diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index e295484cdad..9d5a45d4740 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -105,14 +105,6 @@ func (m queryMetricsStore) DeleteOrganization(ctx context.Context, id uuid.UUID) return r0 } -func (m queryMetricsStore) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - start := time.Now() - r0, r1 := m.s.AcquireExternalAuthLinkRefreshLease(ctx, arg) - m.queryLatencies.WithLabelValues("AcquireExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "AcquireExternalAuthLinkRefreshLease").Inc() - return r0, r1 -} - func (m queryMetricsStore) AcquireLock(ctx context.Context, pgAdvisoryXactLock int64) error { start := time.Now() r0 := m.s.AcquireLock(ctx, pgAdvisoryXactLock) @@ -4961,14 +4953,6 @@ func (m queryMetricsStore) RegisterWorkspaceProxy(ctx context.Context, arg datab return r0, r1 } -func (m queryMetricsStore) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { - start := time.Now() - r0 := m.s.ReleaseExternalAuthLinkRefreshLease(ctx, arg) - m.queryLatencies.WithLabelValues("ReleaseExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ReleaseExternalAuthLinkRefreshLease").Inc() - return r0 -} - func (m queryMetricsStore) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { start := time.Now() r0, r1 := m.s.RemoveUserFromGroups(ctx, arg) @@ -5017,6 +5001,14 @@ func (m queryMetricsStore) SetChatContextSnapshot(ctx context.Context, arg datab return r0 } +func (m queryMetricsStore) SetExternalAuthLinkRefreshLease(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) error { + start := time.Now() + r0 := m.s.SetExternalAuthLinkRefreshLease(ctx, arg) + m.queryLatencies.WithLabelValues("SetExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "SetExternalAuthLinkRefreshLease").Inc() + return r0 +} + func (m queryMetricsStore) SoftDeleteChatMessageByID(ctx context.Context, id int64) error { start := time.Now() r0 := m.s.SoftDeleteChatMessageByID(ctx, id) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 3c14163b0b6..73ce4a58851 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -45,21 +45,6 @@ func (m *MockStore) EXPECT() *MockStoreMockRecorder { return m.recorder } -// AcquireExternalAuthLinkRefreshLease mocks base method. -func (m *MockStore) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "AcquireExternalAuthLinkRefreshLease", ctx, arg) - ret0, _ := ret[0].(database.ExternalAuthLink) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// AcquireExternalAuthLinkRefreshLease indicates an expected call of AcquireExternalAuthLinkRefreshLease. -func (mr *MockStoreMockRecorder) AcquireExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AcquireExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).AcquireExternalAuthLinkRefreshLease), ctx, arg) -} - // AcquireLock mocks base method. func (m *MockStore) AcquireLock(ctx context.Context, pgAdvisoryXactLock int64) error { m.ctrl.T.Helper() @@ -9400,20 +9385,6 @@ func (mr *MockStoreMockRecorder) RegisterWorkspaceProxy(ctx, arg any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterWorkspaceProxy", reflect.TypeOf((*MockStore)(nil).RegisterWorkspaceProxy), ctx, arg) } -// ReleaseExternalAuthLinkRefreshLease mocks base method. -func (m *MockStore) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "ReleaseExternalAuthLinkRefreshLease", ctx, arg) - ret0, _ := ret[0].(error) - return ret0 -} - -// ReleaseExternalAuthLinkRefreshLease indicates an expected call of ReleaseExternalAuthLinkRefreshLease. -func (mr *MockStoreMockRecorder) ReleaseExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReleaseExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).ReleaseExternalAuthLinkRefreshLease), ctx, arg) -} - // RemoveUserFromGroups mocks base method. func (m *MockStore) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { m.ctrl.T.Helper() @@ -9502,6 +9473,20 @@ func (mr *MockStoreMockRecorder) SetChatContextSnapshot(ctx, arg any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetChatContextSnapshot", reflect.TypeOf((*MockStore)(nil).SetChatContextSnapshot), ctx, arg) } +// SetExternalAuthLinkRefreshLease mocks base method. +func (m *MockStore) SetExternalAuthLinkRefreshLease(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "SetExternalAuthLinkRefreshLease", ctx, arg) + ret0, _ := ret[0].(error) + return ret0 +} + +// SetExternalAuthLinkRefreshLease indicates an expected call of SetExternalAuthLinkRefreshLease. +func (mr *MockStoreMockRecorder) SetExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).SetExternalAuthLinkRefreshLease), ctx, arg) +} + // SoftDeleteChatMessageByID mocks base method. func (m *MockStore) SoftDeleteChatMessageByID(ctx context.Context, id int64) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 87647dc6d11..0563a66dc87 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -13,8 +13,6 @@ import ( ) type sqlcQuerier interface { - // Only set the lease if there is not already a non-expired one. - AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg AcquireExternalAuthLinkRefreshLeaseParams) (ExternalAuthLink, error) // Blocks until the lock is acquired. // // This must be called from within a transaction. The lock will be automatically @@ -1317,8 +1315,6 @@ type sqlcQuerier interface { PopNextQueuedMessage(ctx context.Context, chatID uuid.UUID) (ChatQueuedMessage, error) ReduceWorkspaceAgentShareLevelToAuthenticatedByTemplate(ctx context.Context, templateID uuid.UUID) error RegisterWorkspaceProxy(ctx context.Context, arg RegisterWorkspaceProxyParams) (WorkspaceProxy, error) - // Only unset the lease if it matches the one passed in. - ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg ReleaseExternalAuthLinkRefreshLeaseParams) error RemoveUserFromGroups(ctx context.Context, arg RemoveUserFromGroupsParams) ([]uuid.UUID, error) // Mutates only created_at on the target row; ids are unchanged so // consumers can keep tracking queued messages by id. @@ -1337,6 +1333,7 @@ type sqlcQuerier interface { // refresh endpoint. Does not bump updated_at: context pinning is // background state and must not reorder chat lists. SetChatContextSnapshot(ctx context.Context, arg SetChatContextSnapshotParams) error + SetExternalAuthLinkRefreshLease(ctx context.Context, arg SetExternalAuthLinkRefreshLeaseParams) error SoftDeleteChatMessageByID(ctx context.Context, id int64) error SoftDeleteChatMessagesAfterID(ctx context.Context, arg SoftDeleteChatMessagesAfterIDParams) error SoftDeleteContextFileMessages(ctx context.Context, chatID uuid.UUID) error diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index d71b495b258..36a10be0ebd 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -13956,45 +13956,6 @@ func (q *sqlQuerier) RevokeDBCryptKey(ctx context.Context, activeKeyDigest strin return err } -const acquireExternalAuthLinkRefreshLease = `-- name: AcquireExternalAuthLinkRefreshLease :one -UPDATE - external_auth_links -SET - refresh_lease_expires_at = $1 -WHERE - provider_id = $2 - AND user_id = $3 - AND (refresh_lease_expires_at IS NULL OR refresh_lease_expires_at < NOW()) -RETURNING provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at -` - -type AcquireExternalAuthLinkRefreshLeaseParams struct { - RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` - ProviderID string `db:"provider_id" json:"provider_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` -} - -// Only set the lease if there is not already a non-expired one. -func (q *sqlQuerier) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg AcquireExternalAuthLinkRefreshLeaseParams) (ExternalAuthLink, error) { - row := q.db.QueryRowContext(ctx, acquireExternalAuthLinkRefreshLease, arg.RefreshLeaseExpiresAt, arg.ProviderID, arg.UserID) - var i ExternalAuthLink - err := row.Scan( - &i.ProviderID, - &i.UserID, - &i.CreatedAt, - &i.UpdatedAt, - &i.OAuthAccessToken, - &i.OAuthRefreshToken, - &i.OAuthExpiry, - &i.OAuthAccessTokenKeyID, - &i.OAuthRefreshTokenKeyID, - &i.OAuthExtra, - &i.OauthRefreshFailureReason, - &i.RefreshLeaseExpiresAt, - ) - return i, err -} - const deleteExternalAuthLink = `-- name: DeleteExternalAuthLink :exec DELETE FROM external_auth_links WHERE provider_id = $1 AND user_id = $2 ` @@ -14148,26 +14109,24 @@ func (q *sqlQuerier) InsertExternalAuthLink(ctx context.Context, arg InsertExter return i, err } -const releaseExternalAuthLinkRefreshLease = `-- name: ReleaseExternalAuthLinkRefreshLease :exec +const setExternalAuthLinkRefreshLease = `-- name: SetExternalAuthLinkRefreshLease :exec UPDATE external_auth_links SET - refresh_lease_expires_at = NULL + refresh_lease_expires_at = $1 WHERE - provider_id = $1 - AND user_id = $2 - AND refresh_lease_expires_at = $3 + provider_id = $2 + AND user_id = $3 ` -type ReleaseExternalAuthLinkRefreshLeaseParams struct { +type SetExternalAuthLinkRefreshLeaseParams struct { + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` ProviderID string `db:"provider_id" json:"provider_id"` UserID uuid.UUID `db:"user_id" json:"user_id"` - RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` } -// Only unset the lease if it matches the one passed in. -func (q *sqlQuerier) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg ReleaseExternalAuthLinkRefreshLeaseParams) error { - _, err := q.db.ExecContext(ctx, releaseExternalAuthLinkRefreshLease, arg.ProviderID, arg.UserID, arg.RefreshLeaseExpiresAt) +func (q *sqlQuerier) SetExternalAuthLinkRefreshLease(ctx context.Context, arg SetExternalAuthLinkRefreshLeaseParams) error { + _, err := q.db.ExecContext(ctx, setExternalAuthLinkRefreshLease, arg.RefreshLeaseExpiresAt, arg.ProviderID, arg.UserID) return err } diff --git a/coderd/database/queries/externalauth.sql b/coderd/database/queries/externalauth.sql index d92e6c6c561..93d19a078c9 100644 --- a/coderd/database/queries/externalauth.sql +++ b/coderd/database/queries/externalauth.sql @@ -49,25 +49,11 @@ WHERE AND (refresh_lease_expires_at = $3 OR $3 IS NULL) RETURNING *; --- name: AcquireExternalAuthLinkRefreshLease :one --- Only set the lease if there is not already a non-expired one. +-- name: SetExternalAuthLinkRefreshLease :exec UPDATE external_auth_links SET refresh_lease_expires_at = @refresh_lease_expires_at WHERE provider_id = @provider_id - AND user_id = @user_id - AND (refresh_lease_expires_at IS NULL OR refresh_lease_expires_at < NOW()) -RETURNING *; - --- name: ReleaseExternalAuthLinkRefreshLease :exec --- Only unset the lease if it matches the one passed in. -UPDATE - external_auth_links -SET - refresh_lease_expires_at = NULL -WHERE - provider_id = @provider_id - AND user_id = @user_id - AND refresh_lease_expires_at = @refresh_lease_expires_at; + AND user_id = @user_id; diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index 211a20d979c..eef0cf104bc 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -295,47 +295,82 @@ func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.St maximum = c.RefreshLeaseMaxBackoff } r := retry.New(initial, maximum) + lockID := database.GenLockID(fmt.Sprintf("external-auth-refresh:%s-%s", c.ID, externalAuthLink.UserID.String())) + gotLease := false for { - var err error - dblink, err = db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, - }) - switch { - case errors.Is(err, sql.ErrNoRows): - // Something still holds the lock; keep waiting. - if !r.Wait(ctx) { - return externalAuthLink, ctx.Err() + // A refresh can take an arbitrary amount of time, so to avoid holding a + // connection to the database during that time, we only lock temporarily to + // mark the row as being refreshed if not already being refreshed. + err := db.InTx(func(tx database.Store) error { + ok, err := tx.TryAcquireLock(ctx, lockID) + if err != nil { + return xerrors.Errorf("try acquire external auth lock: %w", err) + } + // Link is being refreshed by something else. + if !ok { + return nil } + // Fetch the latest link so we have up-to-date lease information. + dblink, err = tx.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + }) + if err != nil { + return err + } + // Link is being refreshed by something else. + if dblink.RefreshLeaseExpiresAt.Valid && + dblink.RefreshLeaseExpiresAt.Time.After(dbtime.Now()) { + return nil + } + // Link was refreshed while we waited. + if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { + return nil + } + // Acquire the lease. + err = tx.SetExternalAuthLinkRefreshLease(ctx, database.SetExternalAuthLinkRefreshLeaseParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + }) + if err != nil { + return err + } + gotLease = true + return nil + }, nil) + switch { + case gotLease: + // We now hold the lock. + goto refresh case err != nil: // Some kind of DB error. return externalAuthLink, err + // If we got a different token, it means something else refreshed either + // while we were waiting or before we got a hold of the lease but after we + // initially fetched the link. + case dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken: + if dblink.OauthRefreshFailureReason != "" { + return externalAuthLink, refreshError(dblink, dblink.OauthRefreshFailureReason) + } + return dblink, nil default: - // We now hold the lock. - goto refresh + // Something still holds the lock; keep waiting. + if !r.Wait(ctx) { + return externalAuthLink, ctx.Err() + } } } refresh: defer func() { - refreshErr = errors.Join(refreshErr, db.ReleaseExternalAuthLinkRefreshLease(ctx, database.ReleaseExternalAuthLinkRefreshLeaseParams{ + refreshErr = errors.Join(refreshErr, db.SetExternalAuthLinkRefreshLease(ctx, database.SetExternalAuthLinkRefreshLeaseParams{ ProviderID: externalAuthLink.ProviderID, UserID: externalAuthLink.UserID, - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + RefreshLeaseExpiresAt: sql.NullTime{}, })) }() - // If we got a different token, it means something else refreshed either - // while we were waiting or before we got a hold of the lease but after we - // initially fetched the link. - if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { - if dblink.OauthRefreshFailureReason != "" { - return externalAuthLink, refreshError(dblink, dblink.OauthRefreshFailureReason) - } - return dblink, nil - } - // Otherwise the token has still not been updated; refresh it now. newLink, refreshErr = c.refreshAndValidateToken(ctx, db, dblink, lease) return newLink, refreshErr diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index baf126c8717..f9b56bfb7bd 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -50,9 +50,6 @@ func TestRefreshToken(t *testing.T) { t.Run("NoRefreshExpired", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -76,6 +73,8 @@ func TestRefreshToken(t *testing.T) { }, }) + mDB := mockDB(t) + // There should be no database calls since we return early. ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) @@ -89,9 +88,6 @@ func TestRefreshToken(t *testing.T) { t.Run("NoRefreshNoExpiry", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - validated := false fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ @@ -109,6 +105,8 @@ func TestRefreshToken(t *testing.T) { }, }) + mDB := mockDB(t) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) // Zero time used @@ -124,9 +122,6 @@ func TestRefreshToken(t *testing.T) { t.Run("FalseIfTokenSourceFails", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - config := &externalauth.Config{ InstrumentedOAuth2Config: &testutil.OAuth2Config{ TokenSourceFunc: func() (*oauth2.Token, error) { @@ -140,11 +135,9 @@ func TestRefreshToken(t *testing.T) { OAuthExpiry: expired, } - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) + mDB := mockDB(t, + withLink(link), + withLease(link)) ctx := testutil.Context(t, testutil.WaitLong) _, err := config.RefreshToken(ctx, mDB, link) @@ -156,11 +149,6 @@ func TestRefreshToken(t *testing.T) { t.Run("ValidateServerError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - Return(database.ExternalAuthLink{}, nil).AnyTimes() - const staticError = "static error" validated := false fake, config, link := setupOauth2Test(t, testConfig{ @@ -174,15 +162,14 @@ func TestRefreshToken(t *testing.T) { }, }) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) link.OAuthExpiry = expired - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - _, err := config.RefreshToken(ctx, mDB, link) require.ErrorContains(t, err, staticError) // Unsure if this should be the correct behavior. It's an invalid token because @@ -192,10 +179,10 @@ func TestRefreshToken(t *testing.T) { require.True(t, validated, "token should have been attempted to be validated") }) - // RefreshRetries tests that refresh token retry behavior works as expected. - // If a refresh token fails because the token itself is invalid, no more - // refresh attempts should ever happen. An invalid refresh token does - // not magically become valid at some point in the future. + // RefreshRetries tests that refresh token external retry behavior works as + // expected. If a refresh token fails because the token itself is invalid, no + // more refresh attempts should ever happen. An invalid refresh token does not + // magically become valid at some point in the future. // // Internal retries are disabled in this subtest via a negative // RefreshRetryTimeout so each RefreshToken call results in exactly one @@ -205,10 +192,6 @@ func TestRefreshToken(t *testing.T) { t.Parallel() var refreshErr *oauth2.RetrieveError - - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - refreshCount := 0 fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ @@ -229,18 +212,19 @@ func TestRefreshToken(t *testing.T) { // (Windows). cfg.RefreshRetryTimeout = -1 }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, }) - ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - // Expire the link - link.OAuthExpiry = expired + // Allow acquiring and releasing the lease for all the temporary error + // attempts, the bad refresh token attempt, and then finally the last + // attempt with no refresh token set. + mDB := mockDB(t, + withLinkTimes(link, 4), + withLeaseTimes(link, 4)) - // Allow getting the lease for all the temporary error attempts and the - // first bad refresh token attempt. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(4) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(5) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) // Make the failure a server internal error. Not related to the token // This should be retried since this error is temporary. @@ -260,8 +244,10 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, refreshCount, totalRefreshes) } - // Try again with a bad refresh token error. This will invalidate the - // refresh token, and not retry again. + // The final attempt will be a permanent error and we should see the + // database update with the error. Need to extract it from the mock call + // like this rather than use the returned link from RefreshToken as it does + // not return the updated link in the error case. mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). DoAndReturn(func(_ context.Context, p database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { link = database.ExternalAuthLink{ @@ -282,21 +268,33 @@ func TestRefreshToken(t *testing.T) { }, ErrorCode: "bad_refresh_token", } - _, err := config.RefreshToken(ctx, mDB, link) + + zeroedLink, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) totalRefreshes++ require.True(t, externalauth.IsInvalidTokenError(err)) require.Equal(t, refreshCount, totalRefreshes) + // Although the fully updated link with the error is not returned, it does + // zero out the token. + require.Empty(t, zeroedLink.OAuthRefreshToken) + + // Databasae link should have a reason and no refresh token. + require.NotEmpty(t, link.OauthRefreshFailureReason) + require.Empty(t, link.OAuthRefreshToken) - // Update mock database with the zeroed-out link.. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) + // Once more, this time with the zeroed-out refresh token due to the bad + // refresh error. + withLink(link)(mDB) + withLease(link)(mDB) // When the refresh token is empty, no api calls should be made _, err = config.RefreshToken(ctx, mDB, link) require.Error(t, err) require.True(t, externalauth.IsInvalidTokenError(err)) require.Equal(t, refreshCount, totalRefreshes) + // It should return the original cached error, not "no refresh token". + require.ErrorContains(t, err, "a long while ago") + require.ErrorContains(t, err, "bad_refresh_token") }) // RefreshTokenWithBackoff tests that refreshes which fail with transient @@ -398,9 +396,6 @@ func TestRefreshToken(t *testing.T) { t.Run("ConcurrentRefreshGroup", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - parallelRequests := 5 ch := make(chan string) refreshedToken := &oauth2.Token{ @@ -436,22 +431,13 @@ func TestRefreshToken(t *testing.T) { } link := database.ExternalAuthLink{OAuthExpiry: expired} - refreshedLink := database.ExternalAuthLink{ - OAuthAccessToken: refreshedToken.AccessToken, - OAuthRefreshToken: refreshedToken.RefreshToken, - OAuthExpiry: refreshedToken.Expiry, - } - - // Only one call should try to get a lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - // The single winning call will update the link. - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Cond(func(params database.UpdateExternalAuthLinkParams) bool { - return params.ProviderID == link.ProviderID && params.UserID == link.UserID - })).Return(refreshedLink, nil).Times(1) + // Only one call should try to acquire and release a lease and only one call + // should update the link. + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) // When we fire off all requests in parallel... ctx := testutil.Context(t, testutil.WaitLong) @@ -471,7 +457,8 @@ func TestRefreshToken(t *testing.T) { // All calls should have picked up the winning token. for i := range parallelRequests { - require.Equal(t, refreshedLink, results[i]) + require.Equal(t, refreshedToken.AccessToken, results[i].OAuthAccessToken) + require.Equal(t, refreshedToken.RefreshToken, results[i].OAuthRefreshToken) } // Only one refresh call should have actually been made. @@ -490,9 +477,6 @@ func TestRefreshToken(t *testing.T) { t.Run("ConcurrentRefreshRace", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -506,22 +490,21 @@ func TestRefreshToken(t *testing.T) { ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - // Simulate that another caller updated the link. winnerLink := link winnerLink.OAuthRefreshToken = "winner-refresh-token" winnerLink.OAuthAccessToken = "winner-access-token" - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(winnerLink, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) + + // Simulate that another caller updated the link. This will also + // short-circuit acquiring the lease. + mDB := mockDB(t, withLink(winnerLink)) // UpdateExternalAuthLinkRefreshToken should NOT be called because trying to // get the lease detected the nearly-concurrent refresh. It should instead // return the winning token. result, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err, "loser should succeed using the winner's token") - require.Equal(t, "winner-access-token", result.OAuthAccessToken) - require.Equal(t, "winner-refresh-token", result.OAuthRefreshToken) + require.Equal(t, winnerLink.OAuthAccessToken, result.OAuthAccessToken) + require.Equal(t, winnerLink.OAuthRefreshToken, result.OAuthRefreshToken) }) // ConcurrentContextCancel tests that if one request is canceled, it does not @@ -625,52 +608,34 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, int64(1), refreshCalls.Load()) }) - // CachedErrorForEmptyRefreshToken tests that when the refresh token is - // missing, we refer to the cached error, since likely the token was removed in - // a previous attempt and we want to return the original error rather than - // "token is missing". - t.Run("CachedErrorForEmptyRefreshToken", func(t *testing.T) { + t.Run("LeaseAcquisitionError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - - var refreshCalls atomic.Int64 fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { - refreshCalls.Add(1) - return xerrors.New("should not be called") + return nil }), }, + ExternalAuthOpt: func(cfg *externalauth.Config) {}, ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { link.OAuthExpiry = expired - link.OauthRefreshFailureReason = "test cached error" - link.OAuthRefreshToken = "" }, }) - ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + mDB := mockDB(t, + withLink(link), + withLeaseErrors(link, 1, xerrors.New("acquire error"), nil)) - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - - // It should return the cached error message. + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) - require.Equal(t, int64(0), refreshCalls.Load()) require.Error(t, err) - require.ErrorContains(t, err, "test cached error") + require.ErrorContains(t, err, "acquire error") }) t.Run("ReturnsReleaseError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -683,25 +648,24 @@ func TestRefreshToken(t *testing.T) { }, }) - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(xerrors.New("release error")).Times(1) + mDB := mockDB(t, + withLink(link), + withLeaseErrors(link, 1, nil, xerrors.New("release error")), + withUpdatePassthrough()) + // Although the refresh was successful, an error is still returned due to + // the release having failed. ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - _, err := config.RefreshToken(ctx, mDB, link) + refreshed, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) require.ErrorContains(t, err, "release error") + require.NotEqual(t, link.OAuthAccessToken, refreshed.OAuthAccessToken) + require.NotEqual(t, link.OAuthRefreshToken, refreshed.OAuthRefreshToken) }) t.Run("ReturnsCombinedWithReleaseError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -714,13 +678,12 @@ func TestRefreshToken(t *testing.T) { }, }) - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - Return(database.ExternalAuthLink{}, xerrors.New("update error")).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(xerrors.New("release error")).Times(1) + mDB := mockDB(t, + withLink(link), + withLeaseErrors(link, 1, nil, xerrors.New("release error")), + withUpdateError(xerrors.New("update error"))) + // Both the release and update errors should be returned. ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) @@ -732,9 +695,6 @@ func TestRefreshToken(t *testing.T) { t.Run("ValidateFailure", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - const staticError = "static error" validated := false fake, config, link := setupOauth2Test(t, testConfig{ @@ -751,14 +711,10 @@ func TestRefreshToken(t *testing.T) { ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) link.OAuthExpiry = expired - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) _, err := config.RefreshToken(ctx, mDB, link) require.ErrorContains(t, err, "token failed to validate") @@ -769,9 +725,6 @@ func TestRefreshToken(t *testing.T) { t.Run("ValidateRetryGitHub", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - const staticError = "static error" validateCalls := 0 fake, config, link := setupOauth2Test(t, testConfig{ @@ -798,8 +751,10 @@ func TestRefreshToken(t *testing.T) { }, }) - // Since the token is not expired, no refresh lease will be acquired and - // it will only be validated. + // Since the token is not expired, no lock or refresh lease will be acquired + // and it will only be validated. + mDB := mockDB(t) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) @@ -809,9 +764,6 @@ func TestRefreshToken(t *testing.T) { t.Run("ValidateNoUpdate", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - validateCalls := 0 fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ @@ -829,8 +781,10 @@ func TestRefreshToken(t *testing.T) { }, }) - // Since the token is not expired, no refresh lease will be acquired and - // it will only be validated. + // Since the token is not expired, no lock or refresh lease will be acquired + // and it will only be validated. + mDB := mockDB(t) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) @@ -1121,9 +1075,6 @@ func TestRefreshToken(t *testing.T) { t.Run("SaveBeforeValidate_DBError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -1138,17 +1089,12 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - - // Allow getting the lease. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdateError(xerrors.New("db connection lost"))) - mDB.EXPECT(). - UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - Return(database.ExternalAuthLink{}, xerrors.New("db connection lost")) + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) @@ -1157,20 +1103,14 @@ func TestRefreshToken(t *testing.T) { "DB errors should not be treated as invalid token") }) - t.Run("WaitsIfLeasedOK", func(t *testing.T) { + t.Run("WaitsForConcurrentReplicaOK", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { return xerrors.New("should not be called") }), - oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { - return jwt.MapClaims{}, nil - }), }, ExternalAuthOpt: func(cfg *externalauth.Config) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() @@ -1180,24 +1120,24 @@ func TestRefreshToken(t *testing.T) { }, ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { link.OAuthExpiry = expired + // Simulate another replica already having a lease. + link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Hour), Valid: true} }, }) - ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - refreshed := database.ExternalAuthLink{ OAuthAccessToken: "winner-access-token", OAuthRefreshToken: "winner-refresh-token", } - // Simulate another replica already having a lease the first time, then - // resolve with an updated link the second time. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(database.ExternalAuthLink{}, sql.ErrNoRows).Times(1) - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(refreshed, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) + // On the second attempt, simulate the token having been refreshed by the + // other replica. That refreshed token should be returned instead of trying + // to initiate a new refresh. + mDB := mockDB(t, + withLink(link), + withLink(refreshed)) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) updated, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) @@ -1205,20 +1145,14 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, refreshed.OAuthRefreshToken, updated.OAuthRefreshToken) }) - t.Run("WaitsIfLeasedError", func(t *testing.T) { + t.Run("WaitsForConcurrentReplicaError", func(t *testing.T) { t.Parallel() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { return xerrors.New("should not be called") }), - oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { - return jwt.MapClaims{}, nil - }), }, ExternalAuthOpt: func(cfg *externalauth.Config) { cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() @@ -1228,27 +1162,62 @@ func TestRefreshToken(t *testing.T) { }, ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { link.OAuthExpiry = expired + // Simulate another replica already having a lease. + link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Hour), Valid: true} }, }) ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - refreshed := database.ExternalAuthLink{ + errored := database.ExternalAuthLink{ OauthRefreshFailureReason: "failed to refresh", } - // Simulate another replica already having a lease the first time, then - // resolve with the failure the second time. - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(database.ExternalAuthLink{}, sql.ErrNoRows).Times(1) - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(refreshed, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) + // On the second attempt, simulate the token having been failed to be + // refreshed by the other replica. That error should be returned instead of + // trying to initiate a new refresh. + mDB := mockDB(t, + withLink(link), + withLink(errored)) _, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) - require.ErrorContains(t, err, refreshed.OauthRefreshFailureReason) + require.ErrorContains(t, err, errored.OauthRefreshFailureReason) + }) + + t.Run("OverridesStaleLease", func(t *testing.T) { + t.Parallel() + + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + return nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) { + cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() + // Faster polling for faster tests. + cfg.RefreshLeaseInitialBackoff = time.Millisecond + cfg.RefreshLeaseMaxBackoff = time.Millisecond + }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + // Simulate another replica having a stale lease. + link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(-time.Hour), Valid: true} + }, + }) + + ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + + // Should be able to get a lease and update since the other replica's lease + // is stale. + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) + + _, err := config.RefreshToken(ctx, mDB, link) + require.NoError(t, err) }) // OptimisticLockPreventsStaleOverwrite verifies that the @@ -1386,31 +1355,6 @@ func TestRefreshTokenWithScopes(t *testing.T) { expired := dbtime.Now().Add(-time.Hour) - // mockDBPassthrough returns a mock store that uses the provided link for the - // initial fetch then echoes the UpdateExternalAuthLink params back as a - // populated ExternalAuthLink, letting the test read what RefreshToken decided - // to persist. - mockDBPassthrough := func(t *testing.T, link database.ExternalAuthLink) database.Store { - t.Helper() - ctrl := gomock.NewController(t) - mDB := dbmock.NewMockStore(ctrl) - mDB.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(link, nil).Times(1) - mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Any()). - Return(nil).Times(1) - mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). - DoAndReturn(func(_ context.Context, p database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { - return database.ExternalAuthLink{ - ProviderID: p.ProviderID, - UserID: p.UserID, - OAuthAccessToken: p.OAuthAccessToken, - OAuthRefreshToken: p.OAuthRefreshToken, - OAuthExpiry: p.OAuthExpiry, - }, nil - }).AnyTimes() - return mDB - } - t.Run("EchoesConfiguredScopesOnRefresh", func(t *testing.T) { t.Parallel() client, captured := fakeAS(t, @@ -1423,7 +1367,11 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthRefreshToken: "old-r", OAuthExpiry: expired, } - _, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) + _, err := cfg.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, "refresh_token", captured.Get("grant_type")) @@ -1444,7 +1392,11 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthRefreshToken: "old-r", OAuthExpiry: expired, } - _, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) + _, err := cfg.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, "refresh_token", captured.Get("grant_type")) @@ -1466,7 +1418,11 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthRefreshToken: "prior-r", OAuthExpiry: expired, } - link, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) + link, err := cfg.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, "prior-r", link.OAuthRefreshToken, "prior refresh_token must be preserved when AS omits a new one (RFC 6749 §6)") @@ -1484,7 +1440,11 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthRefreshToken: "prior-r", OAuthExpiry: expired, } - link, err := cfg.RefreshToken(ctx, mockDBPassthrough(t, link), link) + mDB := mockDB(t, + withLink(link), + withLease(link), + withUpdatePassthrough()) + link, err := cfg.RefreshToken(ctx, mDB, link) require.NoError(t, err) require.Equal(t, "rotated-r", link.OAuthRefreshToken, "rotated refresh_token from AS must be persisted") @@ -2249,6 +2209,100 @@ func setupOauth2Test(t *testing.T, settings testConfig) (*oidctest.FakeIDP, *ext return fake, config, link } +type mockDBOption func(*dbmock.MockStore) + +// withLink expects that a lock is acquired and the link is fetched; returning +// the provided link for that fetch. +func withLink(link database.ExternalAuthLink) mockDBOption { + return withLinkTimes(link, 1) +} + +// withLinkTimes is like withLink but lets you specify how many times. +func withLinkTimes(link database.ExternalAuthLink, times int) mockDBOption { + return func(mDB *dbmock.MockStore) { + mDB.EXPECT().InTx(gomock.Any(), gomock.Any()).Times(times).DoAndReturn( + func(f func(database.Store) error, opts *database.TxOptions) error { + return f(mDB) + }, + ) + mDB.EXPECT().TryAcquireLock(gomock.Any(), gomock.Any()). + Return(true, nil).Times(times) + mDB.EXPECT().GetExternalAuthLink(gomock.Any(), gomock.Any()). + Return(link, nil).Times(times) + } +} + +// withLease expects that a lease is acquired and released for the provided +// link without any errors. +func withLease(link database.ExternalAuthLink) mockDBOption { + return withLeaseErrors(link, 1, nil, nil) +} + +// withLeaseTimes is like withLease but lets you specify how many times. +func withLeaseTimes(link database.ExternalAuthLink, times int) mockDBOption { + return withLeaseErrors(link, times, nil, nil) +} + +// withLeaseErrors expects that a lease is acquired and released for the +// provided link, return the provided errors. If acquireErr is non-nil then a +// release is not expected and releaseErr goes unused. +func withLeaseErrors(link database.ExternalAuthLink, times int, acquireErr, releaseErr error) mockDBOption { + return func(mDB *dbmock.MockStore) { + // Acquiring the lease. + mDB.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), gomock.Cond(func(params database.SetExternalAuthLinkRefreshLeaseParams) bool { + return params.ProviderID == link.ProviderID && + params.UserID == link.UserID && + params.RefreshLeaseExpiresAt.Valid && + params.RefreshLeaseExpiresAt.Time.After(dbtime.Now()) + })).Return(acquireErr).Times(times) + if acquireErr == nil { + // Releasing the lease. + mDB.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), database.SetExternalAuthLinkRefreshLeaseParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + RefreshLeaseExpiresAt: sql.NullTime{}, + }).Return(releaseErr).Times(times) + } + } +} + +// withUpdateError expects an update to be attempted and returns an error for +// that update. +func withUpdateError(err error) mockDBOption { + return func(mDB *dbmock.MockStore) { + mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). + Return(database.ExternalAuthLink{}, err).Times(1) + } +} + +// withUpdatePassthrough echoes the UpdateExternalAuthLink params back as a +// populated ExternalAuthLink, letting the test read what RefreshToken decided +// to persist, if anything. +func withUpdatePassthrough() mockDBOption { + return func(mDB *dbmock.MockStore) { + mDB.EXPECT().UpdateExternalAuthLink(gomock.Any(), gomock.Any()). + DoAndReturn(func(_ context.Context, p database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { + return database.ExternalAuthLink{ + ProviderID: p.ProviderID, + UserID: p.UserID, + OAuthAccessToken: p.OAuthAccessToken, + OAuthRefreshToken: p.OAuthRefreshToken, + OAuthExpiry: p.OAuthExpiry, + }, nil + }).Times(1) + } +} + +func mockDB(t *testing.T, opts ...mockDBOption) *dbmock.MockStore { + t.Helper() + ctrl := gomock.NewController(t) + mDB := dbmock.NewMockStore(ctrl) + for _, opt := range opts { + opt(mDB) + } + return mDB +} + func TestApplyDefaultsToConfig_CaseInsensitive(t *testing.T) { t.Parallel() diff --git a/enterprise/dbcrypt/dbcrypt.go b/enterprise/dbcrypt/dbcrypt.go index b871204e244..585aa7218cc 100644 --- a/enterprise/dbcrypt/dbcrypt.go +++ b/enterprise/dbcrypt/dbcrypt.go @@ -242,20 +242,6 @@ func (db *dbCrypt) GetExternalAuthLinksByUserID(ctx context.Context, userID uuid return links, nil } -func (db *dbCrypt) AcquireExternalAuthLinkRefreshLease(ctx context.Context, params database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - link, err := db.Store.AcquireExternalAuthLinkRefreshLease(ctx, params) - if err != nil { - return database.ExternalAuthLink{}, err - } - if err := db.decryptField(&link.OAuthAccessToken, link.OAuthAccessTokenKeyID); err != nil { - return database.ExternalAuthLink{}, err - } - if err := db.decryptField(&link.OAuthRefreshToken, link.OAuthRefreshTokenKeyID); err != nil { - return database.ExternalAuthLink{}, err - } - return link, nil -} - func (db *dbCrypt) UpdateExternalAuthLink(ctx context.Context, params database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { if err := db.encryptField(¶ms.OAuthAccessToken, ¶ms.OAuthAccessTokenKeyID); err != nil { return database.ExternalAuthLink{}, err diff --git a/enterprise/dbcrypt/dbcrypt_internal_test.go b/enterprise/dbcrypt/dbcrypt_internal_test.go index 2721973fccd..63f66c29441 100644 --- a/enterprise/dbcrypt/dbcrypt_internal_test.go +++ b/enterprise/dbcrypt/dbcrypt_internal_test.go @@ -332,61 +332,6 @@ func TestExternalAuthLinks(t *testing.T) { }) }) - t.Run("AcquireExternalAuthLinkRefreshLease", func(t *testing.T) { - t.Run("OK", func(t *testing.T) { - t.Parallel() - db, crypt, ciphers := setup(t) - link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ - OAuthAccessToken: "access", - OAuthRefreshToken: "refresh", - }) - link, err := db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ - UserID: link.UserID, - ProviderID: link.ProviderID, - RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now(), Valid: true}, - }) - require.NoError(t, err) - requireEncryptedEquals(t, ciphers[0], link.OAuthAccessToken, "access") - requireEncryptedEquals(t, ciphers[0], link.OAuthRefreshToken, "refresh") - }) - t.Run("Decrypt", func(t *testing.T) { - t.Parallel() - _, crypt, ciphers := setup(t) - link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ - OAuthAccessToken: "access", - OAuthRefreshToken: "refresh", - }) - link, err := crypt.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ - UserID: link.UserID, - ProviderID: link.ProviderID, - RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now(), Valid: true}, - }) - require.NoError(t, err) - require.Equal(t, "access", link.OAuthAccessToken) - require.Equal(t, "refresh", link.OAuthRefreshToken) - require.Equal(t, ciphers[0].HexDigest(), link.OAuthAccessTokenKeyID.String) - require.Equal(t, ciphers[0].HexDigest(), link.OAuthRefreshTokenKeyID.String) - }) - t.Run("DecryptErr", func(t *testing.T) { - t.Parallel() - db, crypt, ciphers := setup(t) - link := dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{ - OAuthAccessToken: fakeBase64RandomData(t, 32), - OAuthRefreshToken: fakeBase64RandomData(t, 32), - OAuthAccessTokenKeyID: sql.NullString{String: ciphers[0].HexDigest(), Valid: true}, - OAuthRefreshTokenKeyID: sql.NullString{String: ciphers[0].HexDigest(), Valid: true}, - }) - link, err := crypt.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ - UserID: link.UserID, - ProviderID: link.ProviderID, - RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now(), Valid: true}, - }) - require.Error(t, err, "expected an error") - var derr *DecryptFailedError - require.ErrorAs(t, err, &derr, "expected a decrypt error") - }) - }) - t.Run("GetExternalAuthLinksByUserID", func(t *testing.T) { t.Parallel() From 521cbc98eb7894cf26392a96ba14cbb54b6dd4a9 Mon Sep 17 00:00:00 2001 From: Asher Date: Tue, 18 Aug 2026 12:08:22 -0800 Subject: [PATCH 07/12] Replace label with for condition --- coderd/externalauth/externalauth.go | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index eef0cf104bc..d1cfb21d3cc 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -297,7 +297,7 @@ func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.St r := retry.New(initial, maximum) lockID := database.GenLockID(fmt.Sprintf("external-auth-refresh:%s-%s", c.ID, externalAuthLink.UserID.String())) gotLease := false - for { + for !gotLease { // A refresh can take an arbitrary amount of time, so to avoid holding a // connection to the database during that time, we only lock temporarily to // mark the row as being refreshed if not already being refreshed. @@ -342,7 +342,7 @@ func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.St switch { case gotLease: // We now hold the lock. - goto refresh + break case err != nil: // Some kind of DB error. return externalAuthLink, err @@ -362,7 +362,6 @@ func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.St } } -refresh: defer func() { refreshErr = errors.Join(refreshErr, db.SetExternalAuthLinkRefreshLease(ctx, database.SetExternalAuthLinkRefreshLeaseParams{ ProviderID: externalAuthLink.ProviderID, From eddbf91d2f36d73cf415d08a3cfc5c538ffd34ab Mon Sep 17 00:00:00 2001 From: Asher Date: Tue, 18 Aug 2026 12:43:06 -0800 Subject: [PATCH 08/12] Add guard when releasing the lease --- coderd/database/querier.go | 2 ++ coderd/database/queries.sql.go | 17 +++++++++++++---- coderd/database/queries/externalauth.sql | 5 ++++- coderd/externalauth/externalauth.go | 5 +++++ coderd/externalauth/externalauth_test.go | 12 +++++++----- 5 files changed, 31 insertions(+), 10 deletions(-) diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 0563a66dc87..b3deca89d54 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -1333,6 +1333,8 @@ type sqlcQuerier interface { // refresh endpoint. Does not bump updated_at: context pinning is // background state and must not reorder chat lists. SetChatContextSnapshot(ctx context.Context, arg SetChatContextSnapshotParams) error + // If an old lease is set, the row will be only updated if it matches the + // current lease. SetExternalAuthLinkRefreshLease(ctx context.Context, arg SetExternalAuthLinkRefreshLeaseParams) error SoftDeleteChatMessageByID(ctx context.Context, id int64) error SoftDeleteChatMessagesAfterID(ctx context.Context, arg SoftDeleteChatMessagesAfterIDParams) error diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 36a10be0ebd..a5e6328ad8f 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -14117,16 +14117,25 @@ SET WHERE provider_id = $2 AND user_id = $3 + AND (refresh_lease_expires_at = $4 OR $4 IS NULL) ` type SetExternalAuthLinkRefreshLeaseParams struct { - RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` - ProviderID string `db:"provider_id" json:"provider_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` + ProviderID string `db:"provider_id" json:"provider_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + OldRefreshLeaseExpiresAt sql.NullTime `db:"old_refresh_lease_expires_at" json:"old_refresh_lease_expires_at"` } +// If an old lease is set, the row will be only updated if it matches the +// current lease. func (q *sqlQuerier) SetExternalAuthLinkRefreshLease(ctx context.Context, arg SetExternalAuthLinkRefreshLeaseParams) error { - _, err := q.db.ExecContext(ctx, setExternalAuthLinkRefreshLease, arg.RefreshLeaseExpiresAt, arg.ProviderID, arg.UserID) + _, err := q.db.ExecContext(ctx, setExternalAuthLinkRefreshLease, + arg.RefreshLeaseExpiresAt, + arg.ProviderID, + arg.UserID, + arg.OldRefreshLeaseExpiresAt, + ) return err } diff --git a/coderd/database/queries/externalauth.sql b/coderd/database/queries/externalauth.sql index 93d19a078c9..cf56cdc823c 100644 --- a/coderd/database/queries/externalauth.sql +++ b/coderd/database/queries/externalauth.sql @@ -50,10 +50,13 @@ WHERE RETURNING *; -- name: SetExternalAuthLinkRefreshLease :exec +-- If an old lease is set, the row will be only updated if it matches the +-- current lease. UPDATE external_auth_links SET refresh_lease_expires_at = @refresh_lease_expires_at WHERE provider_id = @provider_id - AND user_id = @user_id; + AND user_id = @user_id + AND (refresh_lease_expires_at = @old_refresh_lease_expires_at OR @old_refresh_lease_expires_at IS NULL); diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index d1cfb21d3cc..e703ca8220a 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -367,6 +367,11 @@ func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.St ProviderID: externalAuthLink.ProviderID, UserID: externalAuthLink.UserID, RefreshLeaseExpiresAt: sql.NullTime{}, + // The row will only update if we hold the current lease. This is + // somewhat redundant since if we lost the lease our context would be + // expired anyway, so it is not actually possible to get the sql.ErrNoRows + // that would result from this. + OldRefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, })) }() diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index f9b56bfb7bd..879c1d0e918 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -2257,11 +2257,13 @@ func withLeaseErrors(link database.ExternalAuthLink, times int, acquireErr, rele })).Return(acquireErr).Times(times) if acquireErr == nil { // Releasing the lease. - mDB.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), database.SetExternalAuthLinkRefreshLeaseParams{ - ProviderID: link.ProviderID, - UserID: link.UserID, - RefreshLeaseExpiresAt: sql.NullTime{}, - }).Return(releaseErr).Times(times) + mDB.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), gomock.Cond(func(params database.SetExternalAuthLinkRefreshLeaseParams) bool { + return params.ProviderID == link.ProviderID && + params.UserID == link.UserID && + !params.RefreshLeaseExpiresAt.Valid && + params.OldRefreshLeaseExpiresAt.Valid && + params.OldRefreshLeaseExpiresAt.Time.After(dbtime.Now()) + })).Return(releaseErr).Times(times) } } } From f925063642de1e7bfe90f8fea286fbc00a494c9d Mon Sep 17 00:00:00 2001 From: Asher Date: Tue, 18 Aug 2026 16:21:21 -0800 Subject: [PATCH 09/12] Compute lease date in DB Also switch to a SQL function instead of the lock pattern. --- coderd/database/dbauthz/dbauthz.go | 21 +- coderd/database/dbauthz/dbauthz_test.go | 19 +- coderd/database/dbmetrics/querymetrics.go | 24 +- coderd/database/dbmock/dbmock.go | 43 ++- coderd/database/dump.sql | 77 +++-- .../000571_external_auth_lock.down.sql | 1 + .../000571_external_auth_lock.up.sql | 30 ++ coderd/database/querier.go | 8 +- coderd/database/queries.sql.go | 63 ++-- coderd/database/queries/externalauth.sql | 14 +- coderd/externalauth/externalauth.go | 124 +++----- coderd/externalauth/externalauth_test.go | 273 +++++++++++------- enterprise/dbcrypt/dbcrypt.go | 14 + enterprise/dbcrypt/dbcrypt_internal_test.go | 55 ++++ 14 files changed, 497 insertions(+), 269 deletions(-) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 80c92ea6f14..f68a7a1b2f3 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -1731,6 +1731,13 @@ func scopedOrgRoleIdentifiers(names []string, orgID uuid.UUID) []rbac.RoleIdenti return out } +func (q *querier) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + fetch := func(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) + } + return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.AcquireExternalAuthLinkRefreshLease)(ctx, arg) +} + func (q *querier) AcquireLock(ctx context.Context, id int64) error { return q.db.AcquireLock(ctx, id) } @@ -7037,6 +7044,13 @@ func (q *querier) RegisterWorkspaceProxy(ctx context.Context, arg database.Regis return updateWithReturn(q.log, q.auth, fetch, q.db.RegisterWorkspaceProxy)(ctx, arg) } +func (q *querier) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { + fetch := func(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) + } + return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.ReleaseExternalAuthLinkRefreshLease)(ctx, arg) +} + func (q *querier) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { // This is a system function to clear user groups in group sync. if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil { @@ -7094,13 +7108,6 @@ func (q *querier) SetChatContextSnapshot(ctx context.Context, arg database.SetCh return q.db.SetChatContextSnapshot(ctx, arg) } -func (q *querier) SetExternalAuthLinkRefreshLease(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) error { - fetch := func(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { - return q.db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{UserID: arg.UserID, ProviderID: arg.ProviderID}) - } - return fetchAndExec(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.SetExternalAuthLinkRefreshLease)(ctx, arg) -} - func (q *querier) SoftDeleteChatMessageByID(ctx context.Context, id int64) error { msg, err := q.db.GetChatMessageByID(ctx, id) if err != nil { diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 9601ae98b86..b49126fd9f1 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -3202,13 +3202,22 @@ func (s *MethodTestSuite) TestUser() { dbm.EXPECT().UpdateExternalAuthLink(gomock.Any(), arg).Return(link, nil).AnyTimes() check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) })) - s.Run("SetExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + s.Run("AcquireExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{}) dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() - link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Minute), Valid: true} - arg := database.SetExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} - dbm.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(nil).AnyTimes() - check.Args(arg).Asserts(link, policy.ActionUpdatePersonal) + timeout := 10 * time.Second + arg := database.AcquireExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, TimeoutMs: timeout.Milliseconds()} + dbm.EXPECT().AcquireExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(link, nil).AnyTimes() + check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) + })) + s.Run("ReleaseExternalAuthLinkRefreshLease", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + link := testutil.Fake(s.T(), faker, database.ExternalAuthLink{ + RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now().Add(time.Minute), Valid: true}, + }) + dbm.EXPECT().GetExternalAuthLink(gomock.Any(), database.GetExternalAuthLinkParams{ProviderID: link.ProviderID, UserID: link.UserID}).Return(link, nil).AnyTimes() + arg := database.ReleaseExternalAuthLinkRefreshLeaseParams{ProviderID: link.ProviderID, UserID: link.UserID, RefreshLeaseExpiresAt: link.RefreshLeaseExpiresAt} + dbm.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), arg).Return(nil).AnyTimes() + check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns() })) s.Run("UpdateUserLink", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { link := testutil.Fake(s.T(), faker, database.UserLink{}) diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 9d5a45d4740..e295484cdad 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -105,6 +105,14 @@ func (m queryMetricsStore) DeleteOrganization(ctx context.Context, id uuid.UUID) return r0 } +func (m queryMetricsStore) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + start := time.Now() + r0, r1 := m.s.AcquireExternalAuthLinkRefreshLease(ctx, arg) + m.queryLatencies.WithLabelValues("AcquireExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "AcquireExternalAuthLinkRefreshLease").Inc() + return r0, r1 +} + func (m queryMetricsStore) AcquireLock(ctx context.Context, pgAdvisoryXactLock int64) error { start := time.Now() r0 := m.s.AcquireLock(ctx, pgAdvisoryXactLock) @@ -4953,6 +4961,14 @@ func (m queryMetricsStore) RegisterWorkspaceProxy(ctx context.Context, arg datab return r0, r1 } +func (m queryMetricsStore) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { + start := time.Now() + r0 := m.s.ReleaseExternalAuthLinkRefreshLease(ctx, arg) + m.queryLatencies.WithLabelValues("ReleaseExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ReleaseExternalAuthLinkRefreshLease").Inc() + return r0 +} + func (m queryMetricsStore) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { start := time.Now() r0, r1 := m.s.RemoveUserFromGroups(ctx, arg) @@ -5001,14 +5017,6 @@ func (m queryMetricsStore) SetChatContextSnapshot(ctx context.Context, arg datab return r0 } -func (m queryMetricsStore) SetExternalAuthLinkRefreshLease(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) error { - start := time.Now() - r0 := m.s.SetExternalAuthLinkRefreshLease(ctx, arg) - m.queryLatencies.WithLabelValues("SetExternalAuthLinkRefreshLease").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "SetExternalAuthLinkRefreshLease").Inc() - return r0 -} - func (m queryMetricsStore) SoftDeleteChatMessageByID(ctx context.Context, id int64) error { start := time.Now() r0 := m.s.SoftDeleteChatMessageByID(ctx, id) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 73ce4a58851..3c14163b0b6 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -45,6 +45,21 @@ func (m *MockStore) EXPECT() *MockStoreMockRecorder { return m.recorder } +// AcquireExternalAuthLinkRefreshLease mocks base method. +func (m *MockStore) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "AcquireExternalAuthLinkRefreshLease", ctx, arg) + ret0, _ := ret[0].(database.ExternalAuthLink) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// AcquireExternalAuthLinkRefreshLease indicates an expected call of AcquireExternalAuthLinkRefreshLease. +func (mr *MockStoreMockRecorder) AcquireExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AcquireExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).AcquireExternalAuthLinkRefreshLease), ctx, arg) +} + // AcquireLock mocks base method. func (m *MockStore) AcquireLock(ctx context.Context, pgAdvisoryXactLock int64) error { m.ctrl.T.Helper() @@ -9385,6 +9400,20 @@ func (mr *MockStoreMockRecorder) RegisterWorkspaceProxy(ctx, arg any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterWorkspaceProxy", reflect.TypeOf((*MockStore)(nil).RegisterWorkspaceProxy), ctx, arg) } +// ReleaseExternalAuthLinkRefreshLease mocks base method. +func (m *MockStore) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg database.ReleaseExternalAuthLinkRefreshLeaseParams) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReleaseExternalAuthLinkRefreshLease", ctx, arg) + ret0, _ := ret[0].(error) + return ret0 +} + +// ReleaseExternalAuthLinkRefreshLease indicates an expected call of ReleaseExternalAuthLinkRefreshLease. +func (mr *MockStoreMockRecorder) ReleaseExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReleaseExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).ReleaseExternalAuthLinkRefreshLease), ctx, arg) +} + // RemoveUserFromGroups mocks base method. func (m *MockStore) RemoveUserFromGroups(ctx context.Context, arg database.RemoveUserFromGroupsParams) ([]uuid.UUID, error) { m.ctrl.T.Helper() @@ -9473,20 +9502,6 @@ func (mr *MockStoreMockRecorder) SetChatContextSnapshot(ctx, arg any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetChatContextSnapshot", reflect.TypeOf((*MockStore)(nil).SetChatContextSnapshot), ctx, arg) } -// SetExternalAuthLinkRefreshLease mocks base method. -func (m *MockStore) SetExternalAuthLinkRefreshLease(ctx context.Context, arg database.SetExternalAuthLinkRefreshLeaseParams) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "SetExternalAuthLinkRefreshLease", ctx, arg) - ret0, _ := ret[0].(error) - return ret0 -} - -// SetExternalAuthLinkRefreshLease indicates an expected call of SetExternalAuthLinkRefreshLease. -func (mr *MockStoreMockRecorder) SetExternalAuthLinkRefreshLease(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetExternalAuthLinkRefreshLease", reflect.TypeOf((*MockStore)(nil).SetExternalAuthLinkRefreshLease), ctx, arg) -} - // SoftDeleteChatMessageByID mocks base method. func (m *MockStore) SoftDeleteChatMessageByID(ctx context.Context, id int64) error { m.ctrl.T.Helper() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index bd3aa4f9676..e364428137f 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -723,6 +723,60 @@ CREATE TYPE workspace_transition AS ENUM ( 'delete' ); +CREATE TABLE external_auth_links ( + provider_id text NOT NULL, + user_id uuid NOT NULL, + created_at timestamp with time zone NOT NULL, + updated_at timestamp with time zone NOT NULL, + oauth_access_token text NOT NULL, + oauth_refresh_token text NOT NULL, + oauth_expiry timestamp with time zone NOT NULL, + oauth_access_token_key_id text, + oauth_refresh_token_key_id text, + oauth_extra jsonb, + oauth_refresh_failure_reason text DEFAULT ''::text NOT NULL, + refresh_lease_expires_at timestamp with time zone +); + +COMMENT ON COLUMN external_auth_links.oauth_access_token_key_id IS 'The ID of the key used to encrypt the OAuth access token. If this is NULL, the access token is not encrypted'; + +COMMENT ON COLUMN external_auth_links.oauth_refresh_token_key_id IS 'The ID of the key used to encrypt the OAuth refresh token. If this is NULL, the refresh token is not encrypted'; + +COMMENT ON COLUMN external_auth_links.oauth_refresh_failure_reason IS 'This error means the refresh token is invalid. Cached so we can avoid calling the external provider again for the same error.'; + +COMMENT ON COLUMN external_auth_links.refresh_lease_expires_at IS 'Indicates a replica is refreshing the token; prevents concurrent refreshes.'; + +CREATE FUNCTION acquire_external_auth_link_refresh_lease(arg_provider_id text, arg_user_id uuid, timeout_ms bigint) RETURNS SETOF external_auth_links + LANGUAGE plpgsql + AS $$ +DECLARE r external_auth_links; +BEGIN + UPDATE external_auth_links + SET + refresh_lease_expires_at = NOW() + (timeout_ms || ' ms')::interval + WHERE + provider_id = arg_provider_id + AND user_id = arg_user_id + AND (refresh_lease_expires_at IS NULL OR refresh_lease_expires_at < NOW()) + RETURNING * INTO r; + -- Got the lease, return the one row. + IF FOUND THEN + RETURN NEXT r; + RETURN; + END IF; + -- Differentiate between unable to get the lease and the row being gone. + IF EXISTS (SELECT 1 FROM external_auth_links WHERE arg_provider_id = arg_provider_id AND user_id = user_id) THEN + RAISE EXCEPTION 'row is currently leased by another replica' + USING ERRCODE = 'check_violation', + CONSTRAINT = 'external_auth_link_active_lease'; + END IF; + -- Row is gone, return nothing. + RETURN; +END; +$$; + +COMMENT ON FUNCTION acquire_external_auth_link_refresh_lease(arg_provider_id text, arg_user_id uuid, timeout_ms bigint) IS 'Acquire a lease on the external auth link and return the row. If there is already an active lease, an exception is raised.'; + CREATE FUNCTION aggregate_usage_event() RETURNS trigger LANGUAGE plpgsql AS $$ @@ -2324,29 +2378,6 @@ COMMENT ON COLUMN dbcrypt_keys.revoked_at IS 'The time at which the key was revo COMMENT ON COLUMN dbcrypt_keys.test IS 'A column used to test the encryption.'; -CREATE TABLE external_auth_links ( - provider_id text NOT NULL, - user_id uuid NOT NULL, - created_at timestamp with time zone NOT NULL, - updated_at timestamp with time zone NOT NULL, - oauth_access_token text NOT NULL, - oauth_refresh_token text NOT NULL, - oauth_expiry timestamp with time zone NOT NULL, - oauth_access_token_key_id text, - oauth_refresh_token_key_id text, - oauth_extra jsonb, - oauth_refresh_failure_reason text DEFAULT ''::text NOT NULL, - refresh_lease_expires_at timestamp with time zone -); - -COMMENT ON COLUMN external_auth_links.oauth_access_token_key_id IS 'The ID of the key used to encrypt the OAuth access token. If this is NULL, the access token is not encrypted'; - -COMMENT ON COLUMN external_auth_links.oauth_refresh_token_key_id IS 'The ID of the key used to encrypt the OAuth refresh token. If this is NULL, the refresh token is not encrypted'; - -COMMENT ON COLUMN external_auth_links.oauth_refresh_failure_reason IS 'This error means the refresh token is invalid. Cached so we can avoid calling the external provider again for the same error.'; - -COMMENT ON COLUMN external_auth_links.refresh_lease_expires_at IS 'Indicates a replica is refreshing the token; prevents concurrent refreshes.'; - CREATE TABLE files ( hash character varying(64) NOT NULL, created_at timestamp with time zone NOT NULL, diff --git a/coderd/database/migrations/000571_external_auth_lock.down.sql b/coderd/database/migrations/000571_external_auth_lock.down.sql index c51dc71bbc7..82f935fabd6 100644 --- a/coderd/database/migrations/000571_external_auth_lock.down.sql +++ b/coderd/database/migrations/000571_external_auth_lock.down.sql @@ -1 +1,2 @@ ALTER TABLE external_auth_links DROP COLUMN IF EXISTS refresh_lease_expires_at; +DROP FUNCTION IF EXISTS acquire_external_auth_link_refresh_lease; diff --git a/coderd/database/migrations/000571_external_auth_lock.up.sql b/coderd/database/migrations/000571_external_auth_lock.up.sql index 6853d9d8791..08b0db4599d 100644 --- a/coderd/database/migrations/000571_external_auth_lock.up.sql +++ b/coderd/database/migrations/000571_external_auth_lock.up.sql @@ -1,2 +1,32 @@ ALTER TABLE external_auth_links ADD COLUMN IF NOT EXISTS refresh_lease_expires_at timestamp WITH time zone DEFAULT NULL; COMMENT ON COLUMN external_auth_links.refresh_lease_expires_at IS 'Indicates a replica is refreshing the token; prevents concurrent refreshes.'; + +CREATE OR REPLACE FUNCTION acquire_external_auth_link_refresh_lease(arg_provider_id text, arg_user_id uuid, timeout_ms bigint) +RETURNS SETOF external_auth_links AS $$ +DECLARE r external_auth_links; +BEGIN + UPDATE external_auth_links + SET + refresh_lease_expires_at = NOW() + (timeout_ms || ' ms')::interval + WHERE + provider_id = arg_provider_id + AND user_id = arg_user_id + AND (refresh_lease_expires_at IS NULL OR refresh_lease_expires_at < NOW()) + RETURNING * INTO r; + -- Got the lease, return the one row. + IF FOUND THEN + RETURN NEXT r; + RETURN; + END IF; + -- Differentiate between unable to get the lease and the row being gone. + IF EXISTS (SELECT 1 FROM external_auth_links WHERE arg_provider_id = arg_provider_id AND user_id = user_id) THEN + RAISE EXCEPTION 'row is currently leased by another replica' + USING ERRCODE = 'check_violation', + CONSTRAINT = 'external_auth_link_active_lease'; + END IF; + -- Row is gone, return nothing. + RETURN; +END; +$$ LANGUAGE plpgsql; + +COMMENT ON FUNCTION acquire_external_auth_link_refresh_lease IS 'Acquire a lease on the external auth link and return the row. If there is already an active lease, an exception is raised.'; diff --git a/coderd/database/querier.go b/coderd/database/querier.go index b3deca89d54..9fb78cd2bfa 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -13,6 +13,9 @@ import ( ) type sqlcQuerier interface { + // Set the lease to expire according to the provided timeout. If there is + // already a lease, an exception is raised. + AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg AcquireExternalAuthLinkRefreshLeaseParams) (ExternalAuthLink, error) // Blocks until the lock is acquired. // // This must be called from within a transaction. The lock will be automatically @@ -1315,6 +1318,8 @@ type sqlcQuerier interface { PopNextQueuedMessage(ctx context.Context, chatID uuid.UUID) (ChatQueuedMessage, error) ReduceWorkspaceAgentShareLevelToAuthenticatedByTemplate(ctx context.Context, templateID uuid.UUID) error RegisterWorkspaceProxy(ctx context.Context, arg RegisterWorkspaceProxyParams) (WorkspaceProxy, error) + // The lease is only removed if it is the current lease. + ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg ReleaseExternalAuthLinkRefreshLeaseParams) error RemoveUserFromGroups(ctx context.Context, arg RemoveUserFromGroupsParams) ([]uuid.UUID, error) // Mutates only created_at on the target row; ids are unchanged so // consumers can keep tracking queued messages by id. @@ -1333,9 +1338,6 @@ type sqlcQuerier interface { // refresh endpoint. Does not bump updated_at: context pinning is // background state and must not reorder chat lists. SetChatContextSnapshot(ctx context.Context, arg SetChatContextSnapshotParams) error - // If an old lease is set, the row will be only updated if it matches the - // current lease. - SetExternalAuthLinkRefreshLease(ctx context.Context, arg SetExternalAuthLinkRefreshLeaseParams) error SoftDeleteChatMessageByID(ctx context.Context, id int64) error SoftDeleteChatMessagesAfterID(ctx context.Context, arg SoftDeleteChatMessagesAfterIDParams) error SoftDeleteContextFileMessages(ctx context.Context, chatID uuid.UUID) error diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index a5e6328ad8f..afd937e2bca 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -13956,6 +13956,38 @@ func (q *sqlQuerier) RevokeDBCryptKey(ctx context.Context, activeKeyDigest strin return err } +const acquireExternalAuthLinkRefreshLease = `-- name: AcquireExternalAuthLinkRefreshLease :one +SELECT provider_id, user_id, created_at, updated_at, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, oauth_extra, oauth_refresh_failure_reason, refresh_lease_expires_at from acquire_external_auth_link_refresh_lease($1, $2, $3) +` + +type AcquireExternalAuthLinkRefreshLeaseParams struct { + ProviderID string `db:"provider_id" json:"provider_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + TimeoutMs int64 `db:"timeout_ms" json:"timeout_ms"` +} + +// Set the lease to expire according to the provided timeout. If there is +// already a lease, an exception is raised. +func (q *sqlQuerier) AcquireExternalAuthLinkRefreshLease(ctx context.Context, arg AcquireExternalAuthLinkRefreshLeaseParams) (ExternalAuthLink, error) { + row := q.db.QueryRowContext(ctx, acquireExternalAuthLinkRefreshLease, arg.ProviderID, arg.UserID, arg.TimeoutMs) + var i ExternalAuthLink + err := row.Scan( + &i.ProviderID, + &i.UserID, + &i.CreatedAt, + &i.UpdatedAt, + &i.OAuthAccessToken, + &i.OAuthRefreshToken, + &i.OAuthExpiry, + &i.OAuthAccessTokenKeyID, + &i.OAuthRefreshTokenKeyID, + &i.OAuthExtra, + &i.OauthRefreshFailureReason, + &i.RefreshLeaseExpiresAt, + ) + return i, err +} + const deleteExternalAuthLink = `-- name: DeleteExternalAuthLink :exec DELETE FROM external_auth_links WHERE provider_id = $1 AND user_id = $2 ` @@ -14109,33 +14141,26 @@ func (q *sqlQuerier) InsertExternalAuthLink(ctx context.Context, arg InsertExter return i, err } -const setExternalAuthLinkRefreshLease = `-- name: SetExternalAuthLinkRefreshLease :exec +const releaseExternalAuthLinkRefreshLease = `-- name: ReleaseExternalAuthLinkRefreshLease :exec UPDATE external_auth_links SET - refresh_lease_expires_at = $1 + refresh_lease_expires_at = NULL WHERE - provider_id = $2 - AND user_id = $3 - AND (refresh_lease_expires_at = $4 OR $4 IS NULL) + provider_id = $1 + AND user_id = $2 + AND refresh_lease_expires_at = $3 ` -type SetExternalAuthLinkRefreshLeaseParams struct { - RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` - ProviderID string `db:"provider_id" json:"provider_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` - OldRefreshLeaseExpiresAt sql.NullTime `db:"old_refresh_lease_expires_at" json:"old_refresh_lease_expires_at"` +type ReleaseExternalAuthLinkRefreshLeaseParams struct { + ProviderID string `db:"provider_id" json:"provider_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + RefreshLeaseExpiresAt sql.NullTime `db:"refresh_lease_expires_at" json:"refresh_lease_expires_at"` } -// If an old lease is set, the row will be only updated if it matches the -// current lease. -func (q *sqlQuerier) SetExternalAuthLinkRefreshLease(ctx context.Context, arg SetExternalAuthLinkRefreshLeaseParams) error { - _, err := q.db.ExecContext(ctx, setExternalAuthLinkRefreshLease, - arg.RefreshLeaseExpiresAt, - arg.ProviderID, - arg.UserID, - arg.OldRefreshLeaseExpiresAt, - ) +// The lease is only removed if it is the current lease. +func (q *sqlQuerier) ReleaseExternalAuthLinkRefreshLease(ctx context.Context, arg ReleaseExternalAuthLinkRefreshLeaseParams) error { + _, err := q.db.ExecContext(ctx, releaseExternalAuthLinkRefreshLease, arg.ProviderID, arg.UserID, arg.RefreshLeaseExpiresAt) return err } diff --git a/coderd/database/queries/externalauth.sql b/coderd/database/queries/externalauth.sql index cf56cdc823c..aeaf714420f 100644 --- a/coderd/database/queries/externalauth.sql +++ b/coderd/database/queries/externalauth.sql @@ -49,14 +49,18 @@ WHERE AND (refresh_lease_expires_at = $3 OR $3 IS NULL) RETURNING *; --- name: SetExternalAuthLinkRefreshLease :exec --- If an old lease is set, the row will be only updated if it matches the --- current lease. +-- name: AcquireExternalAuthLinkRefreshLease :one +-- Set the lease to expire according to the provided timeout. If there is +-- already a lease, an exception is raised. +SELECT * from acquire_external_auth_link_refresh_lease(@provider_id, @user_id, @timeout_ms); + +-- name: ReleaseExternalAuthLinkRefreshLease :exec +-- The lease is only removed if it is the current lease. UPDATE external_auth_links SET - refresh_lease_expires_at = @refresh_lease_expires_at + refresh_lease_expires_at = NULL WHERE provider_id = @provider_id AND user_id = @user_id - AND (refresh_lease_expires_at = @old_refresh_lease_expires_at OR @old_refresh_lease_expires_at IS NULL); + AND refresh_lease_expires_at = @refresh_lease_expires_at; diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index e703ca8220a..a59b22f53da 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -64,6 +64,10 @@ const ( // defaultRefreshLeaseMaxBackoff is the maximum wait between polls to check // whether another replica has finished refreshing. defaultRefreshLeaseMaxBackoff = 500 * time.Millisecond + + // externalAuthLinkActiveLeaseConstraint indicates the lease could not be + // acquired because something else has an active lease. + externalAuthLinkActiveLeaseConstraint database.CheckConstraint = "external_auth_link_active_lease" ) // SingleflightGroup exposes a subset of singleflight.Group for easier testing. @@ -281,11 +285,9 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu // concurrency protection between multiple instances by using a lease column on // the link's row. func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, timeout time.Duration) (newLink database.ExternalAuthLink, refreshErr error) { - lease := dbtime.Now().Add(timeout) - // There may be other replicas also wanting to refresh; try to get a lease // on the row. This also ensures we have the latest link. - var dblink database.ExternalAuthLink + var leasedLink database.ExternalAuthLink initial := defaultRefreshLeaseInitialBackoff if c.RefreshLeaseInitialBackoff > 0 { initial = c.RefreshLeaseInitialBackoff @@ -295,88 +297,51 @@ func (c *Config) refreshAndValidateWithLease(ctx context.Context, db database.St maximum = c.RefreshLeaseMaxBackoff } r := retry.New(initial, maximum) - lockID := database.GenLockID(fmt.Sprintf("external-auth-refresh:%s-%s", c.ID, externalAuthLink.UserID.String())) - gotLease := false - for !gotLease { - // A refresh can take an arbitrary amount of time, so to avoid holding a - // connection to the database during that time, we only lock temporarily to - // mark the row as being refreshed if not already being refreshed. - err := db.InTx(func(tx database.Store) error { - ok, err := tx.TryAcquireLock(ctx, lockID) - if err != nil { - return xerrors.Errorf("try acquire external auth lock: %w", err) - } - // Link is being refreshed by something else. - if !ok { - return nil - } - // Fetch the latest link so we have up-to-date lease information. - dblink, err = tx.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ + // Make sure to release the lease if we manage to get one before returning. + defer func() { + if leasedLink.RefreshLeaseExpiresAt.Valid { + refreshErr = errors.Join(refreshErr, db.ReleaseExternalAuthLinkRefreshLease(ctx, database.ReleaseExternalAuthLinkRefreshLeaseParams{ ProviderID: externalAuthLink.ProviderID, UserID: externalAuthLink.UserID, - }) - if err != nil { - return err - } - // Link is being refreshed by something else. - if dblink.RefreshLeaseExpiresAt.Valid && - dblink.RefreshLeaseExpiresAt.Time.After(dbtime.Now()) { - return nil - } - // Link was refreshed while we waited. - if dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken { - return nil - } - // Acquire the lease. - err = tx.SetExternalAuthLinkRefreshLease(ctx, database.SetExternalAuthLinkRefreshLeaseParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, - }) - if err != nil { - return err - } - gotLease = true - return nil - }, nil) + // The row will only update if we hold the current lease. This is + // somewhat redundant since if we lost the lease our context would be + // expired anyway, so it is not actually possible to get the sql.ErrNoRows + // that would result from this. + RefreshLeaseExpiresAt: leasedLink.RefreshLeaseExpiresAt, + })) + } + }() + for !leasedLink.RefreshLeaseExpiresAt.Valid { + // Acquiring the lease also returns the link, so we can get the expiry date + // that the database sets and check to see if it was refreshed in the + // meantime. + var err error + leasedLink, err = db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + ProviderID: externalAuthLink.ProviderID, + UserID: externalAuthLink.UserID, + TimeoutMs: timeout.Milliseconds(), + }) switch { - case gotLease: - // We now hold the lock. - break - case err != nil: - // Some kind of DB error. - return externalAuthLink, err - // If we got a different token, it means something else refreshed either - // while we were waiting or before we got a hold of the lease but after we - // initially fetched the link. - case dblink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken: - if dblink.OauthRefreshFailureReason != "" { - return externalAuthLink, refreshError(dblink, dblink.OauthRefreshFailureReason) - } - return dblink, nil - default: - // Something still holds the lock; keep waiting. + // Something still holds the lock; keep waiting. + case database.IsCheckViolation(err, externalAuthLinkActiveLeaseConstraint): if !r.Wait(ctx) { return externalAuthLink, ctx.Err() } + // Some kind of DB error or the row does not exist. + case err != nil: + return externalAuthLink, err + // Something else refreshed either while we were waiting or before we got a + // hold of the lease but after we initially fetched the link. + case leasedLink.OAuthRefreshToken != externalAuthLink.OAuthRefreshToken: + if leasedLink.OauthRefreshFailureReason != "" { + return externalAuthLink, refreshError(leasedLink, leasedLink.OauthRefreshFailureReason) + } + return leasedLink, nil } } - defer func() { - refreshErr = errors.Join(refreshErr, db.SetExternalAuthLinkRefreshLease(ctx, database.SetExternalAuthLinkRefreshLeaseParams{ - ProviderID: externalAuthLink.ProviderID, - UserID: externalAuthLink.UserID, - RefreshLeaseExpiresAt: sql.NullTime{}, - // The row will only update if we hold the current lease. This is - // somewhat redundant since if we lost the lease our context would be - // expired anyway, so it is not actually possible to get the sql.ErrNoRows - // that would result from this. - OldRefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, - })) - }() - // Otherwise the token has still not been updated; refresh it now. - newLink, refreshErr = c.refreshAndValidateToken(ctx, db, dblink, lease) + newLink, refreshErr = c.refreshAndValidateToken(ctx, db, leasedLink) return newLink, refreshErr } @@ -391,8 +356,9 @@ func refreshError(link database.ExternalAuthLink, reason string) error { } // refreshAndValidateToken does the actual token refresh, persists the result to -// the database, then validates the token. -func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink, lease time.Time) (database.ExternalAuthLink, error) { +// the database, then validates the token. The provided link must be up to date +// with the currently held lease. +func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, externalAuthLink database.ExternalAuthLink) (database.ExternalAuthLink, error) { existingToken := externalAuthLink.OAuthToken() // This is additional defensive programming. Because TokenSource is an @@ -449,7 +415,7 @@ func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, // somewhat redundant since if we lost the lease our context would be // expired anyway, so it is not actually possible to get the // sql.ErrNoRows that would result from this. - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + RefreshLeaseExpiresAt: externalAuthLink.RefreshLeaseExpiresAt, }) if updateErr != nil { // This error should be rare. @@ -506,7 +472,7 @@ func (c *Config) refreshAndValidateToken(ctx context.Context, db database.Store, // redundant since if we lost the lease our context would be expired anyway, // so it is not actually possible to get the sql.ErrNoRows that would result // from this. - RefreshLeaseExpiresAt: sql.NullTime{Time: lease, Valid: true}, + RefreshLeaseExpiresAt: externalAuthLink.RefreshLeaseExpiresAt, }) if err != nil { return updatedAuthLink, xerrors.Errorf("persist refreshed token: %w", err) diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index 879c1d0e918..cb0aee258c5 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -20,6 +20,7 @@ import ( "github.com/coreos/go-oidc/v3/oidc" "github.com/golang-jwt/jwt/v4" "github.com/google/uuid" + "github.com/lib/pq" "github.com/prometheus/client_golang/prometheus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -34,6 +35,7 @@ import ( "github.com/coder/coder/v2/coderd" "github.com/coder/coder/v2/coderd/coderdtest/oidctest" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbmock" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/database/dbtime" @@ -135,9 +137,7 @@ func TestRefreshToken(t *testing.T) { OAuthExpiry: expired, } - mDB := mockDB(t, - withLink(link), - withLease(link)) + mDB := mockDB(t, withLease(link)) ctx := testutil.Context(t, testutil.WaitLong) _, err := config.RefreshToken(ctx, mDB, link) @@ -163,7 +163,6 @@ func TestRefreshToken(t *testing.T) { }) mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) @@ -221,8 +220,10 @@ func TestRefreshToken(t *testing.T) { // attempts, the bad refresh token attempt, and then finally the last // attempt with no refresh token set. mDB := mockDB(t, - withLinkTimes(link, 4), - withLeaseTimes(link, 4)) + withLease(link), + withLease(link), + withLease(link), + withLease(link)) ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) @@ -284,7 +285,6 @@ func TestRefreshToken(t *testing.T) { // Once more, this time with the zeroed-out refresh token due to the bad // refresh error. - withLink(link)(mDB) withLease(link)(mDB) // When the refresh token is empty, no api calls should be made @@ -435,7 +435,6 @@ func TestRefreshToken(t *testing.T) { // Only one call should try to acquire and release a lease and only one call // should update the link. mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) @@ -494,13 +493,11 @@ func TestRefreshToken(t *testing.T) { winnerLink.OAuthRefreshToken = "winner-refresh-token" winnerLink.OAuthAccessToken = "winner-access-token" - // Simulate that another caller updated the link. This will also - // short-circuit acquiring the lease. - mDB := mockDB(t, withLink(winnerLink)) - + // Simulate that another caller updated the link. // UpdateExternalAuthLinkRefreshToken should NOT be called because trying to // get the lease detected the nearly-concurrent refresh. It should instead // return the winning token. + mDB := mockDB(t, withLease(winnerLink)) result, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err, "loser should succeed using the winner's token") require.Equal(t, winnerLink.OAuthAccessToken, result.OAuthAccessToken) @@ -623,9 +620,7 @@ func TestRefreshToken(t *testing.T) { }, }) - mDB := mockDB(t, - withLink(link), - withLeaseErrors(link, 1, xerrors.New("acquire error"), nil)) + mDB := mockDB(t, withLeaseErrors(link, xerrors.New("acquire error"), nil)) ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) _, err := config.RefreshToken(ctx, mDB, link) @@ -649,8 +644,7 @@ func TestRefreshToken(t *testing.T) { }) mDB := mockDB(t, - withLink(link), - withLeaseErrors(link, 1, nil, xerrors.New("release error")), + withLeaseErrors(link, nil, xerrors.New("release error")), withUpdatePassthrough()) // Although the refresh was successful, an error is still returned due to @@ -679,8 +673,7 @@ func TestRefreshToken(t *testing.T) { }) mDB := mockDB(t, - withLink(link), - withLeaseErrors(link, 1, nil, xerrors.New("release error")), + withLeaseErrors(link, nil, xerrors.New("release error")), withUpdateError(xerrors.New("update error"))) // Both the release and update errors should be returned. @@ -712,7 +705,6 @@ func TestRefreshToken(t *testing.T) { link.OAuthExpiry = expired mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) @@ -1090,7 +1082,6 @@ func TestRefreshToken(t *testing.T) { }) mDB := mockDB(t, - withLink(link), withLease(link), withUpdateError(xerrors.New("db connection lost"))) @@ -1120,29 +1111,30 @@ func TestRefreshToken(t *testing.T) { }, ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { link.OAuthExpiry = expired - // Simulate another replica already having a lease. - link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Hour), Valid: true} }, }) - refreshed := database.ExternalAuthLink{ - OAuthAccessToken: "winner-access-token", - OAuthRefreshToken: "winner-refresh-token", - } + winnerLink := link + winnerLink.OAuthAccessToken = "winner-access-token" + winnerLink.OAuthRefreshToken = "winner-refresh-token" - // On the second attempt, simulate the token having been refreshed by the - // other replica. That refreshed token should be returned instead of trying - // to initiate a new refresh. + // Simulate another replica already having a lease. On the second attempt, + // simulate the token having been refreshed by the other replica. That + // refreshed token should be returned instead of trying to initiate a new + // refresh. mDB := mockDB(t, - withLink(link), - withLink(refreshed)) + withLeaseErrors(link, &pq.Error{ + Code: pq.ErrorCode("23514"), // check_violation + Constraint: "external_auth_link_active_lease", + }, nil), + withLease(winnerLink)) ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) updated, err := config.RefreshToken(ctx, mDB, link) require.NoError(t, err) - require.Equal(t, refreshed.OAuthAccessToken, updated.OAuthAccessToken) - require.Equal(t, refreshed.OAuthRefreshToken, updated.OAuthRefreshToken) + require.Equal(t, winnerLink.OAuthAccessToken, updated.OAuthAccessToken) + require.Equal(t, winnerLink.OAuthRefreshToken, updated.OAuthRefreshToken) }) t.Run("WaitsForConcurrentReplicaError", func(t *testing.T) { @@ -1162,32 +1154,96 @@ func TestRefreshToken(t *testing.T) { }, ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { link.OAuthExpiry = expired - // Simulate another replica already having a lease. - link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(time.Hour), Valid: true} }, }) ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - errored := database.ExternalAuthLink{ - OauthRefreshFailureReason: "failed to refresh", - } + winnerLink := link + winnerLink.OAuthRefreshToken = "" + winnerLink.OauthRefreshFailureReason = "failed to refresh" - // On the second attempt, simulate the token having been failed to be - // refreshed by the other replica. That error should be returned instead of - // trying to initiate a new refresh. + // Simulate another replica already having a lease. On the second attempt, + // simulate the token having been failed to be refreshed by the other + // replica. That error should be returned instead of trying to initiate a + // new refresh. mDB := mockDB(t, - withLink(link), - withLink(errored)) + withLeaseErrors(link, &pq.Error{ + Code: pq.ErrorCode("23514"), // check_violation + Constraint: "external_auth_link_active_lease", + }, nil), + withLease(winnerLink)) _, err := config.RefreshToken(ctx, mDB, link) require.Error(t, err) - require.ErrorContains(t, err, errored.OauthRefreshFailureReason) + require.ErrorContains(t, err, winnerLink.OauthRefreshFailureReason) }) - t.Run("OverridesStaleLease", func(t *testing.T) { + t.Run("AcquireLeaseAtomicity", func(t *testing.T) { + t.Parallel() + if testing.Short() { + t.SkipNow() + } + + db, _ := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitLong) + + one := dbtestutil.StartTx(t, db, nil) + two := dbtestutil.StartTx(t, db, nil) + + user := dbgen.User(t, db, database.User{}) + link := dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{ + UserID: user.ID, + }) + + // The winner acquires the lease inside an open transaction, holding the + // row lock without committing. + acquired, err := one.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + TimeoutMs: time.Hour.Milliseconds(), + }) + require.NoError(t, err) + require.True(t, acquired.RefreshLeaseExpiresAt.Valid) + + // The loser's UPDATE targets the same row and must block on the winner's + // uncommitted row lock rather than return anything. + loserErr := make(chan error, 1) + go func() { + _, err := two.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + TimeoutMs: time.Hour.Milliseconds(), + }) + loserErr <- err + }() + + select { + case err := <-loserErr: + t.Fatalf("loser returned %v before the winner committed; expected it to block on the row lock", err) + case <-time.After(testutil.IntervalMedium): + } + + // Commit the winner. The loser unblocks, re-checks its predicate against + // the committed row, and errors with the active lease check violation. + require.NoError(t, one.Done()) + err = testutil.RequireReceive(ctx, t, loserErr) + require.True(t, database.IsCheckViolation(err, "external_auth_link_active_lease")) + + // The stored lease is the winner's, untouched by the loser. + final, err := db.GetExternalAuthLink(ctx, database.GetExternalAuthLinkParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + }) + require.NoError(t, err) + require.Equal(t, final.RefreshLeaseExpiresAt.Valid, acquired.RefreshLeaseExpiresAt.Valid) + }) + + t.Run("LeaseDeletedLink", func(t *testing.T) { t.Parallel() + db, _ := dbtestutil.NewDB(t) + fake, config, link := setupOauth2Test(t, testConfig{ FakeIDPOpts: []oidctest.FakeIDPOpt{ oidctest.WithRefresh(func(_ string) error { @@ -1202,21 +1258,48 @@ func TestRefreshToken(t *testing.T) { }, ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { link.OAuthExpiry = expired - // Simulate another replica having a stale lease. - link.RefreshLeaseExpiresAt = sql.NullTime{Time: dbtime.Now().Add(-time.Hour), Valid: true} }, }) ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) + _, err := config.RefreshToken(ctx, db, link) + require.ErrorIs(t, err, sql.ErrNoRows) + }) - // Should be able to get a lease and update since the other replica's lease - // is stale. - mDB := mockDB(t, - withLink(link), - withLease(link), - withUpdatePassthrough()) + t.Run("OverridesStaleLease", func(t *testing.T) { + t.Parallel() - _, err := config.RefreshToken(ctx, mDB, link) + db, _ := dbtestutil.NewDB(t) + + fake, config, link := setupOauth2Test(t, testConfig{ + DB: db, + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + return nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) { + cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() + // Faster polling for faster tests. + cfg.RefreshLeaseInitialBackoff = time.Millisecond + cfg.RefreshLeaseMaxBackoff = time.Millisecond + }, + ExternalAuthLinkOpts: func(link *database.ExternalAuthLink) { + link.OAuthExpiry = expired + }, + }) + + // Simulate another replica having a stale lease. + ctx := testutil.Context(t, testutil.WaitLong) + _, err := db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + TimeoutMs: -time.Hour.Milliseconds(), + }) + require.NoError(t, err) + + ctx = oidc.ClientContext(ctx, fake.HTTPClient(nil)) + _, err = config.RefreshToken(ctx, db, link) require.NoError(t, err) }) @@ -1280,7 +1363,7 @@ func TestRefreshToken(t *testing.T) { OauthRefreshFailureReason: "simulated failure from stale caller B", OAuthRefreshToken: "", UpdatedAt: dbtime.Now(), - // This shou* prevent the write because it does not match. + // This should prevent the write because it does not match. RefreshLeaseExpiresAt: sql.NullTime{Time: dbtime.Now().Add(time.Hour), Valid: true}, // Preserve then. OAuthAccessToken: link.OAuthAccessToken, @@ -1368,7 +1451,6 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthExpiry: expired, } mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) _, err := cfg.RefreshToken(ctx, mDB, link) @@ -1393,7 +1475,6 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthExpiry: expired, } mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) _, err := cfg.RefreshToken(ctx, mDB, link) @@ -1419,7 +1500,6 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthExpiry: expired, } mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) link, err := cfg.RefreshToken(ctx, mDB, link) @@ -1441,7 +1521,6 @@ func TestRefreshTokenWithScopes(t *testing.T) { OAuthExpiry: expired, } mDB := mockDB(t, - withLink(link), withLease(link), withUpdatePassthrough()) link, err := cfg.RefreshToken(ctx, mDB, link) @@ -2211,59 +2290,41 @@ func setupOauth2Test(t *testing.T, settings testConfig) (*oidctest.FakeIDP, *ext type mockDBOption func(*dbmock.MockStore) -// withLink expects that a lock is acquired and the link is fetched; returning -// the provided link for that fetch. -func withLink(link database.ExternalAuthLink) mockDBOption { - return withLinkTimes(link, 1) -} - -// withLinkTimes is like withLink but lets you specify how many times. -func withLinkTimes(link database.ExternalAuthLink, times int) mockDBOption { - return func(mDB *dbmock.MockStore) { - mDB.EXPECT().InTx(gomock.Any(), gomock.Any()).Times(times).DoAndReturn( - func(f func(database.Store) error, opts *database.TxOptions) error { - return f(mDB) - }, - ) - mDB.EXPECT().TryAcquireLock(gomock.Any(), gomock.Any()). - Return(true, nil).Times(times) - mDB.EXPECT().GetExternalAuthLink(gomock.Any(), gomock.Any()). - Return(link, nil).Times(times) - } -} - -// withLease expects that a lease is acquired and released for the provided -// link without any errors. +// withLease wraps withLeaseErrors with nil errors. func withLease(link database.ExternalAuthLink) mockDBOption { - return withLeaseErrors(link, 1, nil, nil) + return withLeaseErrors(link, nil, nil) } -// withLeaseTimes is like withLease but lets you specify how many times. -func withLeaseTimes(link database.ExternalAuthLink, times int) mockDBOption { - return withLeaseErrors(link, times, nil, nil) -} - -// withLeaseErrors expects that a lease is acquired and released for the -// provided link, return the provided errors. If acquireErr is non-nil then a -// release is not expected and releaseErr goes unused. -func withLeaseErrors(link database.ExternalAuthLink, times int, acquireErr, releaseErr error) mockDBOption { +// withLeaseErrors expects that a lease is acquired and that same lease is then +// released. If acquireErr is non-nil then a release is not expected and +// releaseErr goes unused. +func withLeaseErrors(link database.ExternalAuthLink, acquireErr, releaseErr error) mockDBOption { + var lease sql.NullTime return func(mDB *dbmock.MockStore) { - // Acquiring the lease. - mDB.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), gomock.Cond(func(params database.SetExternalAuthLinkRefreshLeaseParams) bool { - return params.ProviderID == link.ProviderID && - params.UserID == link.UserID && - params.RefreshLeaseExpiresAt.Valid && - params.RefreshLeaseExpiresAt.Time.After(dbtime.Now()) - })).Return(acquireErr).Times(times) + mDB.EXPECT().AcquireExternalAuthLinkRefreshLease( + gomock.Any(), + gomock.Cond(func(params database.AcquireExternalAuthLinkRefreshLeaseParams) bool { + return params.ProviderID == link.ProviderID && + params.UserID == link.UserID && + params.TimeoutMs > 0 + })).DoAndReturn(func(_ context.Context, params database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + if acquireErr == nil { + // Return the same link but now with a lease attached. + lease = sql.NullTime{Valid: true, Time: dbtime.Now().Add(time.Duration(params.TimeoutMs) * time.Millisecond)} + leasedLink := link + leasedLink.RefreshLeaseExpiresAt = lease + return leasedLink, nil + } + return database.ExternalAuthLink{}, acquireErr + }).Times(1) if acquireErr == nil { - // Releasing the lease. - mDB.EXPECT().SetExternalAuthLinkRefreshLease(gomock.Any(), gomock.Cond(func(params database.SetExternalAuthLinkRefreshLeaseParams) bool { + mDB.EXPECT().ReleaseExternalAuthLinkRefreshLease(gomock.Any(), gomock.Cond(func(params database.ReleaseExternalAuthLinkRefreshLeaseParams) bool { return params.ProviderID == link.ProviderID && params.UserID == link.UserID && - !params.RefreshLeaseExpiresAt.Valid && - params.OldRefreshLeaseExpiresAt.Valid && - params.OldRefreshLeaseExpiresAt.Time.After(dbtime.Now()) - })).Return(releaseErr).Times(times) + params.RefreshLeaseExpiresAt.Valid && + // Should have passed the same lease back in. + params.RefreshLeaseExpiresAt.Time.Equal(lease.Time) + })).Return(releaseErr).Times(1) } } } diff --git a/enterprise/dbcrypt/dbcrypt.go b/enterprise/dbcrypt/dbcrypt.go index 585aa7218cc..b871204e244 100644 --- a/enterprise/dbcrypt/dbcrypt.go +++ b/enterprise/dbcrypt/dbcrypt.go @@ -242,6 +242,20 @@ func (db *dbCrypt) GetExternalAuthLinksByUserID(ctx context.Context, userID uuid return links, nil } +func (db *dbCrypt) AcquireExternalAuthLinkRefreshLease(ctx context.Context, params database.AcquireExternalAuthLinkRefreshLeaseParams) (database.ExternalAuthLink, error) { + link, err := db.Store.AcquireExternalAuthLinkRefreshLease(ctx, params) + if err != nil { + return database.ExternalAuthLink{}, err + } + if err := db.decryptField(&link.OAuthAccessToken, link.OAuthAccessTokenKeyID); err != nil { + return database.ExternalAuthLink{}, err + } + if err := db.decryptField(&link.OAuthRefreshToken, link.OAuthRefreshTokenKeyID); err != nil { + return database.ExternalAuthLink{}, err + } + return link, nil +} + func (db *dbCrypt) UpdateExternalAuthLink(ctx context.Context, params database.UpdateExternalAuthLinkParams) (database.ExternalAuthLink, error) { if err := db.encryptField(¶ms.OAuthAccessToken, ¶ms.OAuthAccessTokenKeyID); err != nil { return database.ExternalAuthLink{}, err diff --git a/enterprise/dbcrypt/dbcrypt_internal_test.go b/enterprise/dbcrypt/dbcrypt_internal_test.go index 63f66c29441..93b95e23188 100644 --- a/enterprise/dbcrypt/dbcrypt_internal_test.go +++ b/enterprise/dbcrypt/dbcrypt_internal_test.go @@ -332,6 +332,61 @@ func TestExternalAuthLinks(t *testing.T) { }) }) + t.Run("AcquireExternalAuthLinkRefreshLease", func(t *testing.T) { + t.Run("OK", func(t *testing.T) { + t.Parallel() + db, crypt, ciphers := setup(t) + link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ + OAuthAccessToken: "access", + OAuthRefreshToken: "refresh", + }) + link, err := db.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + UserID: link.UserID, + ProviderID: link.ProviderID, + TimeoutMs: 10, + }) + require.NoError(t, err) + requireEncryptedEquals(t, ciphers[0], link.OAuthAccessToken, "access") + requireEncryptedEquals(t, ciphers[0], link.OAuthRefreshToken, "refresh") + }) + t.Run("Decrypt", func(t *testing.T) { + t.Parallel() + _, crypt, ciphers := setup(t) + link := dbgen.ExternalAuthLink(t, crypt, database.ExternalAuthLink{ + OAuthAccessToken: "access", + OAuthRefreshToken: "refresh", + }) + link, err := crypt.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + UserID: link.UserID, + ProviderID: link.ProviderID, + TimeoutMs: 10, + }) + require.NoError(t, err) + require.Equal(t, "access", link.OAuthAccessToken) + require.Equal(t, "refresh", link.OAuthRefreshToken) + require.Equal(t, ciphers[0].HexDigest(), link.OAuthAccessTokenKeyID.String) + require.Equal(t, ciphers[0].HexDigest(), link.OAuthRefreshTokenKeyID.String) + }) + t.Run("DecryptErr", func(t *testing.T) { + t.Parallel() + db, crypt, ciphers := setup(t) + link := dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{ + OAuthAccessToken: fakeBase64RandomData(t, 32), + OAuthRefreshToken: fakeBase64RandomData(t, 32), + OAuthAccessTokenKeyID: sql.NullString{String: ciphers[0].HexDigest(), Valid: true}, + OAuthRefreshTokenKeyID: sql.NullString{String: ciphers[0].HexDigest(), Valid: true}, + }) + link, err := crypt.AcquireExternalAuthLinkRefreshLease(ctx, database.AcquireExternalAuthLinkRefreshLeaseParams{ + UserID: link.UserID, + ProviderID: link.ProviderID, + TimeoutMs: 10, + }) + require.Error(t, err, "expected an error") + var derr *DecryptFailedError + require.ErrorAs(t, err, &derr, "expected a decrypt error") + }) + }) + t.Run("GetExternalAuthLinksByUserID", func(t *testing.T) { t.Parallel() From 26c37add9364225db1079f3a2628c0592d4b5459 Mon Sep 17 00:00:00 2001 From: Asher Date: Tue, 18 Aug 2026 16:23:48 -0800 Subject: [PATCH 10/12] Bump migration number New migrations have since come in on main. --- ...rnal_auth_lock.down.sql => 000574_external_auth_lock.down.sql} | 0 ...external_auth_lock.up.sql => 000574_external_auth_lock.up.sql} | 0 2 files changed, 0 insertions(+), 0 deletions(-) rename coderd/database/migrations/{000571_external_auth_lock.down.sql => 000574_external_auth_lock.down.sql} (100%) rename coderd/database/migrations/{000571_external_auth_lock.up.sql => 000574_external_auth_lock.up.sql} (100%) diff --git a/coderd/database/migrations/000571_external_auth_lock.down.sql b/coderd/database/migrations/000574_external_auth_lock.down.sql similarity index 100% rename from coderd/database/migrations/000571_external_auth_lock.down.sql rename to coderd/database/migrations/000574_external_auth_lock.down.sql diff --git a/coderd/database/migrations/000571_external_auth_lock.up.sql b/coderd/database/migrations/000574_external_auth_lock.up.sql similarity index 100% rename from coderd/database/migrations/000571_external_auth_lock.up.sql rename to coderd/database/migrations/000574_external_auth_lock.up.sql From d533f0a928e83164ae52b4dc2a706d1e7d089f23 Mon Sep 17 00:00:00 2001 From: Asher Date: Wed, 19 Aug 2026 12:04:19 -0800 Subject: [PATCH 11/12] Fix select fallback condition --- coderd/database/dump.sql | 2 +- .../000574_external_auth_lock.up.sql | 2 +- coderd/externalauth/externalauth_test.go | 20 ++++++++++++++++--- 3 files changed, 19 insertions(+), 5 deletions(-) diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index 9c03e488dba..999b98bc509 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -767,7 +767,7 @@ BEGIN RETURN; END IF; -- Differentiate between unable to get the lease and the row being gone. - IF EXISTS (SELECT 1 FROM external_auth_links WHERE arg_provider_id = arg_provider_id AND user_id = user_id) THEN + IF EXISTS (SELECT 1 FROM external_auth_links WHERE provider_id = arg_provider_id AND user_id = arg_user_id) THEN RAISE EXCEPTION 'row is currently leased by another replica' USING ERRCODE = 'check_violation', CONSTRAINT = 'external_auth_link_active_lease'; diff --git a/coderd/database/migrations/000574_external_auth_lock.up.sql b/coderd/database/migrations/000574_external_auth_lock.up.sql index 08b0db4599d..56ee05ae576 100644 --- a/coderd/database/migrations/000574_external_auth_lock.up.sql +++ b/coderd/database/migrations/000574_external_auth_lock.up.sql @@ -19,7 +19,7 @@ BEGIN RETURN; END IF; -- Differentiate between unable to get the lease and the row being gone. - IF EXISTS (SELECT 1 FROM external_auth_links WHERE arg_provider_id = arg_provider_id AND user_id = user_id) THEN + IF EXISTS (SELECT 1 FROM external_auth_links WHERE provider_id = arg_provider_id AND user_id = arg_user_id) THEN RAISE EXCEPTION 'row is currently leased by another replica' USING ERRCODE = 'check_violation', CONSTRAINT = 'external_auth_link_active_lease'; diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index 9eb00592f03..1001668dce3 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -1308,7 +1308,7 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, final.RefreshLeaseExpiresAt.Valid, acquired.RefreshLeaseExpiresAt.Valid) }) - t.Run("LeaseDeletedLink", func(t *testing.T) { + t.Run("LeaseMissingLink", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) @@ -1330,8 +1330,22 @@ func TestRefreshToken(t *testing.T) { }, }) - ctx := oidc.ClientContext(testutil.Context(t, testutil.WaitLong), fake.HTTPClient(nil)) - _, err := config.RefreshToken(ctx, db, link) + // Insert another link to ensure the link acquisition function's select + // fallback matches on the right provider/user. + ctx := testutil.Context(t, testutil.WaitLong) + _, err := db.InsertExternalAuthLink(ctx, database.InsertExternalAuthLinkParams{ + ProviderID: "decoy-provider", + UserID: uuid.New(), + CreatedAt: dbtime.Now(), + UpdatedAt: dbtime.Now(), + OAuthAccessToken: "x", + OAuthRefreshToken: "x", + OAuthExpiry: dbtime.Now().Add(time.Hour), + }) + require.NoError(t, err) + + ctx = oidc.ClientContext(ctx, fake.HTTPClient(nil)) + _, err = config.RefreshToken(ctx, db, link) require.ErrorIs(t, err, sql.ErrNoRows) }) From 5a1de17b80fb540a9f5437b107f6835911a8c516 Mon Sep 17 00:00:00 2001 From: Asher Date: Wed, 19 Aug 2026 12:31:57 -0800 Subject: [PATCH 12/12] Bump migration number again --- ...rnal_auth_lock.down.sql => 000577_external_auth_lock.down.sql} | 0 ...external_auth_lock.up.sql => 000577_external_auth_lock.up.sql} | 0 2 files changed, 0 insertions(+), 0 deletions(-) rename coderd/database/migrations/{000574_external_auth_lock.down.sql => 000577_external_auth_lock.down.sql} (100%) rename coderd/database/migrations/{000574_external_auth_lock.up.sql => 000577_external_auth_lock.up.sql} (100%) diff --git a/coderd/database/migrations/000574_external_auth_lock.down.sql b/coderd/database/migrations/000577_external_auth_lock.down.sql similarity index 100% rename from coderd/database/migrations/000574_external_auth_lock.down.sql rename to coderd/database/migrations/000577_external_auth_lock.down.sql diff --git a/coderd/database/migrations/000574_external_auth_lock.up.sql b/coderd/database/migrations/000577_external_auth_lock.up.sql similarity index 100% rename from coderd/database/migrations/000574_external_auth_lock.up.sql rename to coderd/database/migrations/000577_external_auth_lock.up.sql