From 48607126544bff44885f433418b2880f078f91bb Mon Sep 17 00:00:00 2001 From: sav-labs Date: Wed, 13 May 2026 17:55:08 +0200 Subject: [PATCH 1/3] feat(coderd/database): add UpdateUserLinkRefreshToken with optimistic lock Adds a new query for updating user_links token columns that only writes when the refresh token currently stored in the database still matches the one the caller read before contacting the IdP. Callers should treat sql.ErrNoRows as 'another caller refreshed first' and re-read the row. This is a prerequisite for deduplicating concurrent OIDC refresh requests in ValidateAPIKey across coderd replicas. --- coderd/database/dbauthz/dbauthz.go | 10 ++++ coderd/database/dbauthz/dbauthz_test.go | 15 ++++++ coderd/database/dbmetrics/querymetrics.go | 8 +++ coderd/database/dbmock/dbmock.go | 15 ++++++ coderd/database/querier.go | 7 +++ coderd/database/queries.sql.go | 64 +++++++++++++++++++++++ coderd/database/queries/user_links.sql | 24 +++++++++ 7 files changed, 143 insertions(+) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 83c391068bd..e45acb0d466 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -7137,6 +7137,16 @@ func (q *querier) UpdateUserLink(ctx context.Context, arg database.UpdateUserLin return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.UpdateUserLink)(ctx, arg) } +func (q *querier) UpdateUserLinkRefreshToken(ctx context.Context, arg database.UpdateUserLinkRefreshTokenParams) (database.UserLink, error) { + fetch := func(ctx context.Context, arg database.UpdateUserLinkRefreshTokenParams) (database.UserLink, error) { + return q.db.GetUserLinkByUserIDLoginType(ctx, database.GetUserLinkByUserIDLoginTypeParams{ + UserID: arg.UserID, + LoginType: arg.LoginType, + }) + } + return fetchAndQuery(q.log, q.auth, policy.ActionUpdatePersonal, fetch, q.db.UpdateUserLinkRefreshToken)(ctx, arg) +} + func (q *querier) UpdateUserLoginType(ctx context.Context, arg database.UpdateUserLoginTypeParams) (database.User, error) { if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceSystem); err != nil { return database.User{}, err diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 11fdb3e47ae..f46cacedb0d 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -2976,6 +2976,21 @@ func (s *MethodTestSuite) TestUser() { dbm.EXPECT().UpdateUserLink(gomock.Any(), arg).Return(link, nil).AnyTimes() check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) })) + s.Run("UpdateUserLinkRefreshToken", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + link := testutil.Fake(s.T(), faker, database.UserLink{}) + arg := database.UpdateUserLinkRefreshTokenParams{ + OAuthAccessToken: link.OAuthAccessToken, + OAuthRefreshToken: link.OAuthRefreshToken, + OAuthExpiry: link.OAuthExpiry, + UserID: link.UserID, + LoginType: link.LoginType, + Claims: database.UserLinkClaims{}, + OldOauthRefreshToken: link.OAuthRefreshToken, + } + dbm.EXPECT().GetUserLinkByUserIDLoginType(gomock.Any(), database.GetUserLinkByUserIDLoginTypeParams{UserID: link.UserID, LoginType: link.LoginType}).Return(link, nil).AnyTimes() + dbm.EXPECT().UpdateUserLinkRefreshToken(gomock.Any(), arg).Return(link, nil).AnyTimes() + check.Args(arg).Asserts(link, policy.ActionUpdatePersonal).Returns(link) + })) s.Run("UpdateUserRoles", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { u := testutil.Fake(s.T(), faker, database.User{RBACRoles: []string{codersdk.RoleTemplateAdmin}}) o := u diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 12191ac07ef..05b858b9f33 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -5097,6 +5097,14 @@ func (m queryMetricsStore) UpdateUserLink(ctx context.Context, arg database.Upda return r0, r1 } +func (m queryMetricsStore) UpdateUserLinkRefreshToken(ctx context.Context, arg database.UpdateUserLinkRefreshTokenParams) (database.UserLink, error) { + start := time.Now() + r0, r1 := m.s.UpdateUserLinkRefreshToken(ctx, arg) + m.queryLatencies.WithLabelValues("UpdateUserLinkRefreshToken").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateUserLinkRefreshToken").Inc() + return r0, r1 +} + func (m queryMetricsStore) UpdateUserLoginType(ctx context.Context, arg database.UpdateUserLoginTypeParams) (database.User, error) { start := time.Now() r0, r1 := m.s.UpdateUserLoginType(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 1eab730283e..5a105ac9a54 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -9603,6 +9603,21 @@ func (mr *MockStoreMockRecorder) UpdateUserLink(ctx, arg any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateUserLink", reflect.TypeOf((*MockStore)(nil).UpdateUserLink), ctx, arg) } +// UpdateUserLinkRefreshToken mocks base method. +func (m *MockStore) UpdateUserLinkRefreshToken(ctx context.Context, arg database.UpdateUserLinkRefreshTokenParams) (database.UserLink, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpdateUserLinkRefreshToken", ctx, arg) + ret0, _ := ret[0].(database.UserLink) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpdateUserLinkRefreshToken indicates an expected call of UpdateUserLinkRefreshToken. +func (mr *MockStoreMockRecorder) UpdateUserLinkRefreshToken(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateUserLinkRefreshToken", reflect.TypeOf((*MockStore)(nil).UpdateUserLinkRefreshToken), ctx, arg) +} + // UpdateUserLoginType mocks base method. func (m *MockStore) UpdateUserLoginType(ctx context.Context, arg database.UpdateUserLoginTypeParams) (database.User, error) { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index cb6b243f5ce..873a1d0aba8 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -1207,6 +1207,13 @@ type sqlcQuerier interface { UpdateUserHashedPassword(ctx context.Context, arg UpdateUserHashedPasswordParams) error UpdateUserLastSeenAt(ctx context.Context, arg UpdateUserLastSeenAtParams) (User, error) UpdateUserLink(ctx context.Context, arg UpdateUserLinkParams) (UserLink, 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 (across replicas, where + // in-process deduplication via singleflight cannot reach) from overwriting a + // valid token stored by the winner. Callers should treat sql.ErrNoRows as + // "another caller refreshed first" and re-read the row rather than erroring. + UpdateUserLinkRefreshToken(ctx context.Context, arg UpdateUserLinkRefreshTokenParams) (UserLink, error) UpdateUserLoginType(ctx context.Context, arg UpdateUserLoginTypeParams) (User, error) UpdateUserNotificationPreferences(ctx context.Context, arg UpdateUserNotificationPreferencesParams) (int64, error) UpdateUserProfile(ctx context.Context, arg UpdateUserProfileParams) (User, error) diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 8befc53e369..15f028f216f 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -24991,6 +24991,70 @@ func (q *sqlQuerier) UpdateUserLink(ctx context.Context, arg UpdateUserLinkParam return i, err } +const updateUserLinkRefreshToken = `-- name: UpdateUserLinkRefreshToken :one +UPDATE + user_links +SET + oauth_access_token = $1, + oauth_access_token_key_id = $2, + oauth_refresh_token = $3, + oauth_refresh_token_key_id = $4, + oauth_expiry = $5, + claims = $6 +WHERE + user_id = $7 +AND + login_type = $8 +AND + oauth_refresh_token = $9 +RETURNING user_id, login_type, linked_id, oauth_access_token, oauth_refresh_token, oauth_expiry, oauth_access_token_key_id, oauth_refresh_token_key_id, claims +` + +type UpdateUserLinkRefreshTokenParams struct { + 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"` + Claims UserLinkClaims `db:"claims" json:"claims"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + LoginType LoginType `db:"login_type" json:"login_type"` + OldOauthRefreshToken string `db:"old_oauth_refresh_token" json:"old_oauth_refresh_token"` +} + +// 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 (across replicas, where +// in-process deduplication via singleflight cannot reach) from overwriting a +// valid token stored by the winner. Callers should treat sql.ErrNoRows as +// "another caller refreshed first" and re-read the row rather than erroring. +func (q *sqlQuerier) UpdateUserLinkRefreshToken(ctx context.Context, arg UpdateUserLinkRefreshTokenParams) (UserLink, error) { + row := q.db.QueryRowContext(ctx, updateUserLinkRefreshToken, + arg.OAuthAccessToken, + arg.OAuthAccessTokenKeyID, + arg.OAuthRefreshToken, + arg.OAuthRefreshTokenKeyID, + arg.OAuthExpiry, + arg.Claims, + arg.UserID, + arg.LoginType, + arg.OldOauthRefreshToken, + ) + var i UserLink + err := row.Scan( + &i.UserID, + &i.LoginType, + &i.LinkedID, + &i.OAuthAccessToken, + &i.OAuthRefreshToken, + &i.OAuthExpiry, + &i.OAuthAccessTokenKeyID, + &i.OAuthRefreshTokenKeyID, + &i.Claims, + ) + return i, err +} + const createUserSecret = `-- name: CreateUserSecret :one INSERT INTO user_secrets ( id, diff --git a/coderd/database/queries/user_links.sql b/coderd/database/queries/user_links.sql index b352e808401..0f1536d8735 100644 --- a/coderd/database/queries/user_links.sql +++ b/coderd/database/queries/user_links.sql @@ -50,6 +50,30 @@ SET WHERE user_id = $7 AND login_type = $8 RETURNING *; +-- name: UpdateUserLinkRefreshToken :one +-- 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 (across replicas, where +-- in-process deduplication via singleflight cannot reach) from overwriting a +-- valid token stored by the winner. Callers should treat sql.ErrNoRows as +-- "another caller refreshed first" and re-read the row rather than erroring. +UPDATE + user_links +SET + oauth_access_token = @oauth_access_token, + oauth_access_token_key_id = @oauth_access_token_key_id, + oauth_refresh_token = @oauth_refresh_token, + oauth_refresh_token_key_id = @oauth_refresh_token_key_id, + oauth_expiry = @oauth_expiry, + claims = @claims +WHERE + user_id = @user_id +AND + login_type = @login_type +AND + oauth_refresh_token = @old_oauth_refresh_token +RETURNING *; + -- name: OIDCClaimFields :many -- OIDCClaimFields returns a list of distinct keys in the the merged_claims fields. -- This query is used to generate the list of available sync fields for idp sync settings. From 807fc6aaf3e12781128a5312643dab6427db79ea Mon Sep 17 00:00:00 2001 From: sav-labs Date: Wed, 13 May 2026 17:55:17 +0200 Subject: [PATCH 2/3] fix(coderd/httpmw): dedupe concurrent OIDC token refresh in ValidateAPIKey When the OAuth access token expires and a user's dashboard refresh fires multiple parallel API requests, each goroutine independently called the IdP with the same single-use refresh token. The first won; the rest got invalid_grant and could overwrite the valid token with stale state on the losing path. Wraps the refresh in a singleflight group keyed by user ID and login type so concurrent goroutines in the same coderd process share one IdP call. Persists the new tokens via UpdateUserLinkRefreshToken, whose optimistic-lock predicate converts cross-replica races into sql.ErrNoRows, which the helper resolves by re-reading the row written by the winner. Closes #25275. --- coderd/httpmw/apikey.go | 189 ++++++++++++++++++++++++++++++++-------- 1 file changed, 151 insertions(+), 38 deletions(-) diff --git a/coderd/httpmw/apikey.go b/coderd/httpmw/apikey.go index 40a87647f36..bdb3aae3c05 100644 --- a/coderd/httpmw/apikey.go +++ b/coderd/httpmw/apikey.go @@ -16,6 +16,7 @@ import ( "golang.org/x/net/idna" "golang.org/x/oauth2" "golang.org/x/xerrors" + "tailscale.com/util/singleflight" "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/apikey" @@ -118,6 +119,16 @@ func UserAuthorization(ctx context.Context) rbac.Subject { type OAuth2Configs struct { Github promoauth.OAuth2Config OIDC promoauth.OAuth2Config + + // RefreshGroup deduplicates concurrent OAuth refresh attempts inside a + // single coderd process. Multiple parallel HTTP requests from the same + // user that hit ValidateAPIKey after the OAuth access token has expired + // would otherwise each independently call the IdP with the same + // single-use refresh token; only one would win and the rest would fail + // with invalid_grant. The singleflight key is ":" so + // refreshes for different users do not block each other. The zero value + // is ready to use. + RefreshGroup singleflight.Group[string, database.UserLink] } func (c *OAuth2Configs) IsZero() bool { @@ -356,53 +367,30 @@ func ValidateAPIKey(ctx context.Context, cfg ValidateAPIKeyConfig, r *http.Reque } } - // We have a refresh token, so let's try it. - token, err := oauthConfig.TokenSource(r.Context(), &oauth2.Token{ - AccessToken: link.OAuthAccessToken, - RefreshToken: link.OAuthRefreshToken, - Expiry: link.OAuthExpiry, - }).Token() - // Hard error: we actively tried to refresh and the - // provider rejected it — surface even on optional-auth - // routes. - if err != nil { - return nil, &ValidateAPIKeyError{ - Code: http.StatusUnauthorized, - Response: codersdk.Response{ - Message: fmt.Sprintf( - "Could not refresh expired %s token. Try re-authenticating to resolve this issue.", - friendlyName), - Detail: err.Error(), - }, - Hard: true, + // Deduplicate concurrent refreshes for the same user. The + // closure either refreshes against the IdP or returns the + // link that another caller already refreshed. + refreshed, sfErr, _ := cfg.OAuth2Configs.RefreshGroup.Do( + key.UserID.String()+":"+string(key.LoginType), + func() (database.UserLink, error) { + return refreshOAuthLink(ctx, cfg.DB, oauthConfig, friendlyName, link) + }, + ) + if sfErr != nil { + var vErr *ValidateAPIKeyError + if errors.As(sfErr, &vErr) { + return nil, vErr } - } - link.OAuthAccessToken = token.AccessToken - link.OAuthRefreshToken = token.RefreshToken - link.OAuthExpiry = token.Expiry - //nolint:gocritic // system needs to update user link - link, err = cfg.DB.UpdateUserLink(dbauthz.AsSystemRestricted(ctx), database.UpdateUserLinkParams{ - UserID: link.UserID, - LoginType: link.LoginType, - OAuthAccessToken: link.OAuthAccessToken, - OAuthAccessTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthRefreshToken: link.OAuthRefreshToken, - OAuthRefreshTokenKeyID: sql.NullString{}, // dbcrypt will update as required - OAuthExpiry: link.OAuthExpiry, - // Refresh should keep the same debug context because we use - // the original claims for the group/role sync. - Claims: link.Claims, - }) - if err != nil { return nil, &ValidateAPIKeyError{ Code: http.StatusInternalServerError, Response: codersdk.Response{ Message: internalErrorMessage, - Detail: fmt.Sprintf("update user_link: %s.", err.Error()), + Detail: fmt.Sprintf("refresh user_link: %s.", sfErr.Error()), }, Hard: true, } } + link = refreshed } } @@ -1024,6 +1012,131 @@ func RedirectToLogin(rw http.ResponseWriter, r *http.Request, dashboardURL *url. http.Redirect(rw, r, u.String(), http.StatusSeeOther) } +// refreshOAuthLink performs a single OIDC/GitHub access-token refresh for +// the given user link. It is intended to be invoked from inside +// OAuth2Configs.RefreshGroup.Do so concurrent callers for the same user +// share the result. The function: +// +// 1. Opens a transaction and takes a postgres advisory lock keyed on +// (user_id, login_type). The lock serializes refreshes for one user +// across coderd replicas. Only one replica at a time enters the +// critical section, while peers block on AcquireLock and observe the +// winner's fresh token on re-read. +// 2. Re-reads the user_link row inside the lock so the goroutine that +// wins the singleflight (or the advisory lock from another replica) +// sees the latest tokens. +// 3. Skips the IdP call if the freshly read link is no longer expired. +// 4. Otherwise refreshes against the OAuth2 provider and persists the +// new tokens via UpdateUserLinkRefreshToken. The optimistic-lock +// predicate in that query is defense-in-depth against unexpected +// concurrent writers; under normal operation the advisory lock makes +// the update single-writer. +// +// Errors that should surface as user-visible HTTP responses are wrapped in +// *ValidateAPIKeyError so the caller can return them verbatim. +func refreshOAuthLink( + ctx context.Context, + db database.Store, + oauthConfig promoauth.OAuth2Config, + friendlyName string, + original database.UserLink, +) (database.UserLink, error) { + var result database.UserLink + //nolint:gocritic // system needs to refresh user_link + txErr := db.InTx(func(tx database.Store) error { + lockID := database.GenLockID(fmt.Sprintf("oauth-refresh:%s:%s", original.UserID, original.LoginType)) + if err := tx.AcquireLock(dbauthz.AsSystemRestricted(ctx), lockID); err != nil { + return &ValidateAPIKeyError{ + Code: http.StatusInternalServerError, + Response: codersdk.Response{ + Message: internalErrorMessage, + Detail: fmt.Sprintf("acquire oauth-refresh advisory lock: %s", err.Error()), + }, + Hard: true, + } + } + + link, err := tx.GetUserLinkByUserIDLoginType(dbauthz.AsSystemRestricted(ctx), database.GetUserLinkByUserIDLoginTypeParams{ + UserID: original.UserID, + LoginType: original.LoginType, + }) + if err != nil { + return &ValidateAPIKeyError{ + Code: http.StatusInternalServerError, + Response: codersdk.Response{ + Message: "A database error occurred", + Detail: fmt.Sprintf("re-read user_link for refresh: %s", err.Error()), + }, + Hard: true, + } + } + // Another goroutine or replica already refreshed; reuse its result. + if !link.OAuthExpiry.IsZero() && link.OAuthExpiry.After(dbtime.Now()) { + result = link + return nil + } + + token, err := oauthConfig.TokenSource(ctx, &oauth2.Token{ + AccessToken: link.OAuthAccessToken, + RefreshToken: link.OAuthRefreshToken, + Expiry: link.OAuthExpiry, + }).Token() + if err != nil { + return &ValidateAPIKeyError{ + Code: http.StatusUnauthorized, + Response: codersdk.Response{ + Message: fmt.Sprintf( + "Could not refresh expired %s token. Try re-authenticating to resolve this issue.", + friendlyName), + Detail: err.Error(), + }, + Hard: true, + } + } + + updated, err := tx.UpdateUserLinkRefreshToken(dbauthz.AsSystemRestricted(ctx), database.UpdateUserLinkRefreshTokenParams{ + UserID: link.UserID, + LoginType: link.LoginType, + 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, + // Refresh should keep the same debug context because we use + // the original claims for the group/role sync. + Claims: link.Claims, + OldOauthRefreshToken: link.OAuthRefreshToken, + }) + if err != nil { + return &ValidateAPIKeyError{ + Code: http.StatusInternalServerError, + Response: codersdk.Response{ + Message: internalErrorMessage, + Detail: fmt.Sprintf("update user_link: %s.", err.Error()), + }, + Hard: true, + } + } + result = updated + return nil + }, nil) + if txErr != nil { + var vErr *ValidateAPIKeyError + if errors.As(txErr, &vErr) { + return database.UserLink{}, vErr + } + return database.UserLink{}, &ValidateAPIKeyError{ + Code: http.StatusInternalServerError, + Response: codersdk.Response{ + Message: internalErrorMessage, + Detail: fmt.Sprintf("refresh user_link transaction: %s", txErr.Error()), + }, + Hard: true, + } + } + return result, nil +} + // CustomRedirectToLogin redirects the user to the login page with the `message` and // `redirect` query parameters set, with a provided code func CustomRedirectToLogin(rw http.ResponseWriter, r *http.Request, redirect string, message string, code int) { From b1bd3e87e16d4185533bc36a321c3bc206e475a0 Mon Sep 17 00:00:00 2001 From: sav-labs Date: Wed, 13 May 2026 17:55:25 +0200 Subject: [PATCH 3/3] test(coderd/httpmw): cover concurrent OIDC refresh deduplication Adds a regression test that spawns N concurrent ValidateAPIKey calls for one user with an expired OAuth access token and asserts the IdP mock observed exactly one refresh attempt. --- coderd/httpmw/apikey_test.go | 79 ++++++++++++++++++++++++++++++++++++ 1 file changed, 79 insertions(+) diff --git a/coderd/httpmw/apikey_test.go b/coderd/httpmw/apikey_test.go index d060330427b..de8ccd42fc4 100644 --- a/coderd/httpmw/apikey_test.go +++ b/coderd/httpmw/apikey_test.go @@ -692,6 +692,85 @@ func TestAPIKey(t *testing.T) { require.Equal(t, gotLink.OAuthRefreshToken, "moo") }) + t.Run("OAuthRefreshSingleflight", func(t *testing.T) { + t.Parallel() + // Verifies that concurrent ValidateAPIKey calls for the same user + // with an expired OAuth access token only trigger a single + // refresh against the IdP. Regression test for coder/coder#25275. + var ( + db, _ = dbtestutil.NewDB(t) + user = dbgen.User(t, db, database.User{}) + _ = dbgen.UserLink(t, db, database.UserLink{ + UserID: user.ID, + LoginType: database.LoginTypeGithub, + OAuthRefreshToken: "old-refresh", + OAuthAccessToken: "old-access", + OAuthExpiry: dbtime.Now().AddDate(0, 0, -1), + }) + ) + // Each request needs its own session token, but they share one user/link. + const concurrency = 5 + tokens := make([]string, concurrency) + for i := range tokens { + _, tok := dbgen.APIKey(t, db, database.APIKey{ + UserID: user.ID, + LastUsed: dbtime.Now(), + ExpiresAt: dbtime.Now().AddDate(0, 0, 1), + LoginType: database.LoginTypeGithub, + }) + tokens[i] = tok + } + + var refreshCalls atomic.Int64 + // Slow the token source slightly so concurrent goroutines actually + // queue on the singleflight group rather than serializing on the + // scheduler. + oauthCfg := &testutil.OAuth2Config{ + TokenSourceFunc: testutil.OAuth2TokenSource(func() (*oauth2.Token, error) { + refreshCalls.Add(1) + time.Sleep(50 * time.Millisecond) + return &oauth2.Token{ + AccessToken: "new-access", + RefreshToken: "new-refresh", + Expiry: dbtestutil.NowInDefaultTimezone().AddDate(0, 0, 1), + }, nil + }), + } + + mw := httpmw.ExtractAPIKeyMW(httpmw.ExtractAPIKeyConfig{ + DB: db, + OAuth2Configs: &httpmw.OAuth2Configs{Github: oauthCfg}, + RedirectToLogin: false, + }) + + start := make(chan struct{}) + results := make(chan int, concurrency) + for i := 0; i < concurrency; i++ { + tok := tokens[i] + go func() { + <-start + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set(codersdk.SessionTokenHeader, tok) + rw := httptest.NewRecorder() + mw(successHandler).ServeHTTP(rw, r) + results <- rw.Result().StatusCode + }() + } + close(start) + for i := 0; i < concurrency; i++ { + require.Equal(t, http.StatusOK, <-results) + } + require.EqualValues(t, 1, refreshCalls.Load(), + "expected exactly one refresh request to the IdP for %d concurrent callers", concurrency) + + gotLink, err := db.GetUserLinkByUserIDLoginType(context.Background(), database.GetUserLinkByUserIDLoginTypeParams{ + UserID: user.ID, + LoginType: database.LoginTypeGithub, + }) + require.NoError(t, err) + require.Equal(t, "new-refresh", gotLink.OAuthRefreshToken) + }) + t.Run("OAuthExpiredNoRefresh", func(t *testing.T) { t.Parallel() var (