diff --git a/cli/create_test.go b/cli/create_test.go index 73778be1d63d6..b8fb1b4e64de6 100644 --- a/cli/create_test.go +++ b/cli/create_test.go @@ -11,6 +11,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/sync/singleflight" "github.com/coder/coder/v2/cli" "github.com/coder/coder/v2/cli/clitest" @@ -2117,6 +2118,7 @@ func TestCreateWithGitAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, IncludeProvisionerDaemon: true, }) diff --git a/coderd/aitasks_test.go b/coderd/aitasks_test.go index a5425eba62b13..9ee6df9e1ad40 100644 --- a/coderd/aitasks_test.go +++ b/coderd/aitasks_test.go @@ -17,6 +17,7 @@ import ( "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" agentapisdk "github.com/coder/agentapi-sdk-go" @@ -1524,6 +1525,7 @@ func TestCreateTaskExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1582,6 +1584,7 @@ func TestCreateTaskExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1633,6 +1636,7 @@ func TestCreateTaskExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1665,6 +1669,7 @@ func TestCreateTaskExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) diff --git a/coderd/coderdtest/oidctest/idp.go b/coderd/coderdtest/oidctest/idp.go index a7f608c632cfd..4bf6d0287dac1 100644 --- a/coderd/coderdtest/oidctest/idp.go +++ b/coderd/coderdtest/oidctest/idp.go @@ -33,6 +33,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/oauth2" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" "cdr.dev/slog/v3" @@ -1641,6 +1642,7 @@ func (f *FakeIDP) ExternalAuthConfig(t testing.TB, id string, custom *ExternalAu Scopes: []string{}, CodeURL: f.locked.Provider().DeviceCodeURL, }, + RefreshGroup: new(singleflight.Group), } if !custom.UseDeviceAuth { diff --git a/coderd/externalauth/externalauth.go b/coderd/externalauth/externalauth.go index 6633ab1936e6c..ed88a4843dd25 100644 --- a/coderd/externalauth/externalauth.go +++ b/coderd/externalauth/externalauth.go @@ -19,6 +19,7 @@ import ( "github.com/sqlc-dev/pqtype" "golang.org/x/oauth2" xgithub "golang.org/x/oauth2/github" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd/database" @@ -53,6 +54,12 @@ const ( defaultRefreshRetryTimeout = 10 * time.Second ) +// SingleflightGroup exposes a subset of singleflight.Group for easier testing. +// singleflight.Group should be used instead of implementing this in production. +type SingleflightGroup interface { + DoChan(key string, fn func() (any, error)) <-chan singleflight.Result +} + // Config is used for authentication for Git operations. type Config struct { promoauth.InstrumentedOAuth2Config @@ -142,6 +149,9 @@ type Config struct { // defaultRefreshRetryTimeout. A negative value disables transient-failure // retries entirely, so exactly one refresh attempt is made. RefreshRetryTimeout time.Duration + + // RefreshGroup deduplicates concurrent requests. + RefreshGroup SingleflightGroup } // Git returns a Provider for this config if the provider type is a @@ -190,6 +200,37 @@ 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) { + // 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) { + // 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 + // seconds for updating the database and validating the link. + timeout := 10 * time.Second + if c.RefreshRetryTimeout > 0 { + timeout += c.RefreshRetryTimeout + } + rctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), timeout) + defer cancel() + return c.innerRefreshToken(rctx, db, externalAuthLink) + }) + select { + case results := <-ch: + if newlink, ok := results.Val.(database.ExternalAuthLink); ok { + return newlink, results.Err + } else if results.Err == nil { + return externalAuthLink, xerrors.Errorf("got invalid type from token refresh: %T", results.Val) + } + return externalAuthLink, results.Err + case <-ctx.Done(): + return externalAuthLink, ctx.Err() + } +} + +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 && @@ -237,21 +278,17 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu // // The error message is saved for debugging purposes. if isFailedRefresh(existingToken, err) { - // Before caching the failure, re-read the external auth link - // from the database. A 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. + // 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 { - // Another caller won the refresh race and stored a new - // refresh token. Return their updated link instead of - // caching a failure. return currentLink, nil } @@ -322,15 +359,9 @@ func (c *Config) RefreshToken(ctx context.Context, db database.Store, externalAu // 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. - // Use a detached context for the DB write only. The IDP already - // consumed the old refresh token, so if the caller's request - // context is canceled mid-save, the new token would be lost. - persistCtx, persistCancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second) - defer persistCancel() - originalAccessToken := externalAuthLink.OAuthAccessToken if token.AccessToken != originalAccessToken { - updatedAuthLink, err := db.UpdateExternalAuthLink(persistCtx, database.UpdateExternalAuthLinkParams{ + updatedAuthLink, err := db.UpdateExternalAuthLink(ctx, database.UpdateExternalAuthLinkParams{ ProviderID: c.ID, UserID: externalAuthLink.UserID, UpdatedAt: dbtime.Now(), @@ -924,6 +955,7 @@ func ConvertConfig(instrument *promoauth.Factory, entries []codersdk.ExternalAut MCPToolAllowRegex: mcpToolAllow, MCPToolDenyRegex: mcpToolDeny, CodeChallengeMethodsSupported: slice.StringEnums[promoauth.Oauth2PKCEChallengeMethod](entry.CodeChallengeMethodsSupported), + RefreshGroup: new(singleflight.Group), } if entry.DeviceFlow { diff --git a/coderd/externalauth/externalauth_test.go b/coderd/externalauth/externalauth_test.go index 9951e06dc96e1..8d91f6df9e2fb 100644 --- a/coderd/externalauth/externalauth_test.go +++ b/coderd/externalauth/externalauth_test.go @@ -1,6 +1,7 @@ package externalauth_test import ( + "bytes" "context" "encoding/json" "fmt" @@ -8,7 +9,9 @@ import ( "net/http" "net/http/httptest" "net/url" + "runtime/debug" "strings" + "sync" "sync/atomic" "testing" "time" @@ -21,6 +24,8 @@ import ( "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" "golang.org/x/oauth2" + "golang.org/x/sync/errgroup" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd" @@ -109,6 +114,7 @@ func TestRefreshToken(t *testing.T) { return nil, xerrors.New("failure") }, }, + RefreshGroup: new(singleflight.Group), } _, err := config.RefreshToken(context.Background(), nil, database.ExternalAuthLink{ @@ -336,13 +342,95 @@ func TestRefreshToken(t *testing.T) { "permanent failures should not be retried") }) - // ConcurrentRefreshRace tests that when multiple concurrent requests - // race to refresh the same token, the loser does not poison the - // database with a cached "bad_refresh_token" failure. This - // reproduces the issue described in coder/coder#17069 where - // providers with single-use refresh tokens (e.g., GitHub Apps) - // reject the second refresh attempt, and the resulting error was - // incorrectly cached. + // ConcurrentRefreshGroup tests that when requests try to refresh a token + // while another request is pending, they wait on the first caller and share + // the result instead of all attempting to perform the refresh. + 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{ + AccessToken: "winner-access-token", + RefreshToken: "winner-refresh-token", + Expiry: time.Now().Add(time.Hour), + } + + var refreshCalls atomic.Int64 + config := &externalauth.Config{ + InstrumentedOAuth2Config: &testutil.OAuth2Config{ + // The first call to refresh will succeed and all others will fail. The + // first will wait for all callers to join the group before returning. + TokenSourceFunc: func() (*oauth2.Token, error) { + if refreshCalls.Add(1) == 1 { + // Wait for all the other calls to be subscribed, to prevent + // the test from flaking. + subscribed := 1 + for { + <-ch + subscribed++ + if subscribed >= parallelRequests { + return refreshedToken, nil + } + } + } + return nil, xerrors.New("bad_refresh_token") + }, + }, + RefreshGroup: &group{ + notify: ch, + }, + } + + link := database.ExternalAuthLink{OAuthExpiry: expired} + refreshedLink := database.ExternalAuthLink{ + OAuthAccessToken: refreshedToken.AccessToken, + OAuthRefreshToken: refreshedToken.RefreshToken, + OAuthExpiry: refreshedToken.Expiry, + } + + // 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) + + // When we fire off all requests in parallel... + ctx := testutil.Context(t, testutil.WaitLong) + var eg errgroup.Group + results := make([]database.ExternalAuthLink, parallelRequests) + for i := range parallelRequests { + eg.Go(func() error { + result, err := config.RefreshToken(ctx, mDB, link) + results[i] = result + return err + }) + } + + // No call should error. + err := eg.Wait() + require.NoError(t, err) + + // All calls should have picked up the winning token. + for i := range parallelRequests { + require.Equal(t, refreshedLink, results[i]) + } + + // Only one refresh call should have actually been made. + require.Equal(t, int64(1), refreshCalls.Load()) + }) + + // 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. + // + // 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. t.Run("ConcurrentRefreshRace", func(t *testing.T) { t.Parallel() @@ -387,6 +475,106 @@ func TestRefreshToken(t *testing.T) { require.Equal(t, "winner-refresh-token", result.OAuthRefreshToken) }) + // ConcurrentContextCancel tests that if one request is canceled, it does not + // cancel other requests waiting on it. + t.Run("ConcurrentContextCanceled", func(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + parallelRequests := 5 + ch := make(chan string) + + var refreshCalls atomic.Int64 + ctx := testutil.Context(t, testutil.WaitLong) + cancelOnRefresh, cancel := context.WithCancel(ctx) + defer cancel() + + // Use to know when the first call has started the group, so we know which + // context we can cancel. + listening := make(chan struct{}) + + fake, config, link := setupOauth2Test(t, testConfig{ + FakeIDPOpts: []oidctest.FakeIDPOpt{ + oidctest.WithRefresh(func(_ string) error { + if refreshCalls.Add(1) == 1 { + close(listening) + // Wait for all the other calls to be subscribed, to prevent + // the test from flaking. + subscribed := 1 + for { + <-ch + subscribed++ + if subscribed >= parallelRequests { + // Cancel the parent context after refresh succeeds + // but before the DB save and validation. + cancel() + return nil + } + } + } + // Should never reach here. + return xerrors.New("bad_refresh_token") + }), + oidctest.WithDynamicUserInfo(func(_ string) (jwt.MapClaims, error) { + return jwt.MapClaims{}, nil + }), + }, + ExternalAuthOpt: func(cfg *externalauth.Config) { + cfg.Type = codersdk.EnhancedExternalAuthProviderGitHub.String() + cfg.RefreshGroup = &group{notify: ch} + }, + DB: db, + }) + + oldAccessToken := link.OAuthAccessToken + oldRefreshToken := link.OAuthRefreshToken + link.OAuthExpiry = expired + + var wg sync.WaitGroup + // Start the first call with the cancelable context. + wg.Add(1) + go func() { + defer wg.Done() + ctx := oidc.ClientContext(cancelOnRefresh, fake.HTTPClient(nil)) + _, err := config.RefreshToken(ctx, db, link) + assert.ErrorIs(t, err, context.Canceled) + }() + + // Wait for it to start the group, to make sure the callback above is + // canceling the right context (if we fire them all at once, any one of them + // could start the group). + <-listening + + // Now we can fire off the remaining requests. + for range parallelRequests - 1 { + wg.Add(1) + go func() { + defer wg.Done() + ctx := oidc.ClientContext(ctx, fake.HTTPClient(nil)) + result, err := config.RefreshToken(ctx, db, link) + assert.NoError(t, err) + assert.NotEqual(t, oldAccessToken, result.OAuthAccessToken) + assert.NotEqual(t, oldRefreshToken, result.OAuthRefreshToken) + }() + } + + wg.Wait() + + // DB link should have been updated. + dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + }) + require.NoError(t, err) + require.NotEqual(t, oldAccessToken, dbLink.OAuthAccessToken, + "DB should have the new access token despite context cancellation") + require.NotEqual(t, oldRefreshToken, dbLink.OAuthRefreshToken, + "DB should have the new refresh token despite context cancellation") + + // Only one refresh call should have actually been made. + require.Equal(t, int64(1), refreshCalls.Load()) + }) + // ValidateFailure tests if the token is no longer valid with a 401 response. t.Run("ValidateFailure", func(t *testing.T) { t.Parallel() @@ -666,18 +854,21 @@ func TestRefreshToken(t *testing.T) { link.OAuthExpiry = expired _, err := config.RefreshToken(ctx, db, link) - require.NoError(t, err) + require.ErrorIs(t, err, context.Canceled) require.Equal(t, int64(1), refreshCalls.Load()) - dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ - ProviderID: link.ProviderID, - UserID: link.UserID, - }) - require.NoError(t, err) - require.NotEqual(t, oldAccessToken, dbLink.OAuthAccessToken, - "DB should have the new access token despite context cancellation") - require.NotEqual(t, oldRefreshToken, dbLink.OAuthRefreshToken, - "DB should have the new refresh token despite context cancellation") + require.Eventually(t, func() bool { + dbLink, err := db.GetExternalAuthLink(context.Background(), database.GetExternalAuthLinkParams{ + ProviderID: link.ProviderID, + UserID: link.UserID, + }) + if err != nil { + return false + } + return err == nil && + dbLink.OAuthAccessToken != oldAccessToken && + dbLink.OAuthRefreshToken != oldRefreshToken + }, testutil.WaitShort, testutil.IntervalFast, "never saw refresh token db updated") }) // SaveBeforeValidate_RateLimited tests the full path: refresh @@ -1009,6 +1200,7 @@ func TestValidateToken(t *testing.T) { ID: "test-validate", Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), ValidateURL: validateURL, + RefreshGroup: new(singleflight.Group), } } @@ -1613,6 +1805,7 @@ func setupOauth2Test(t *testing.T, settings testConfig) (*oidctest.FakeIDP, *ext RevokeURL: fake.WellknownConfig().RevokeURL, RevokeTimeout: 1 * time.Second, CodeChallengeMethodsSupported: []promoauth.Oauth2PKCEChallengeMethod{promoauth.PKCEChallengeMethodSha256}, + RefreshGroup: new(singleflight.Group), } settings.ExternalAuthOpt(config) @@ -1689,3 +1882,166 @@ type roundTripper func(req *http.Request) (*http.Response, error) func (r roundTripper) RoundTrip(req *http.Request) (*http.Response, error) { return r(req) } + +var _ externalauth.SingleflightGroup = (*group)(nil) + +// The following has been copied from x/sync/singleflight but has been modified +// to notify when callers join the group so the tests can be deterministic. + +// Copyright 2013 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// errGoexit indicates runtime.Goexit was called in +// the user-given function. +var errGoexit = xerrors.New("runtime.Goexit was called") + +// A panicError is an arbitrary value recovered from a panic +// with the stack trace during the execution of the given function. +type panicError struct { + value any + stack []byte +} + +// Error implements error interface. +func (p *panicError) Error() string { + return fmt.Sprintf("%v\n\n%s", p.value, p.stack) +} + +func (p *panicError) Unwrap() error { + err, ok := p.value.(error) + if !ok { + return nil + } + + return err +} + +func newPanicError(v any) error { + stack := debug.Stack() + + // The first line of the stack trace is of the form "goroutine N [status]:" + // but by the time the panic reaches Do the goroutine may no longer exist + // and its status will have changed. Trim out the misleading line. + if line := bytes.IndexByte(stack, '\n'); line >= 0 { + stack = stack[line+1:] + } + return &panicError{value: v, stack: stack} +} + +// call is an in-flight or completed singleflight.Do call +type call struct { + wg sync.WaitGroup + + // These fields are written once before the WaitGroup is done + // and are only read after the WaitGroup is done. + val any + err error + + // These fields are read and written with the singleflight + // mutex held before the WaitGroup is done, and are read but + // not written after the WaitGroup is done. + dups int + chans []chan<- singleflight.Result +} + +// group represents a class of work and forms a namespace in +// which units of work can be executed with duplicate suppression. +type group struct { + mu sync.Mutex // protects m + m map[string]*call // lazily initialized + notify chan string +} + +// DoChan is like Do but returns a channel that will receive the +// results when they are ready. +// +// The returned channel will not be closed. +func (g *group) DoChan(key string, fn func() (any, error)) <-chan singleflight.Result { + ch := make(chan singleflight.Result, 1) + g.mu.Lock() + if g.m == nil { + g.m = make(map[string]*call) + } + if c, ok := g.m[key]; ok { + c.dups++ + c.chans = append(c.chans, ch) + g.notify <- key + g.mu.Unlock() + return ch + } + c := &call{chans: []chan<- singleflight.Result{ch}} + c.wg.Add(1) + g.m[key] = c + g.mu.Unlock() + + go g.doCall(c, key, fn) + + return ch +} + +// doCall handles the single call for a key. +func (g *group) doCall(c *call, key string, fn func() (any, error)) { + normalReturn := false + recovered := false + + // use double-defer to distinguish panic from runtime.Goexit, + // more details see https://golang.org/cl/134395 + defer func() { + // the given function invoked runtime.Goexit + if !normalReturn && !recovered { + c.err = errGoexit + } + + g.mu.Lock() + defer g.mu.Unlock() + c.wg.Done() + if g.m[key] == c { + delete(g.m, key) + } + + //nolint:errorlint // Avoid changing the original code. + if e, ok := c.err.(*panicError); ok { + // In order to prevent the waiting channels from being blocked forever, + // needs to ensure that this panic cannot be recovered. + //nolint:revive // Avoid changing the original code. + if len(c.chans) > 0 { + go panic(e) + select {} // Keep this goroutine around so that it will appear in the crash dump. + } else { + panic(e) + } + } else if c.err == errGoexit { //nolint:revive // Avoid changing the original code. + // Already in the process of goexit, no need to call again + } else { + // Normal return + for _, ch := range c.chans { + ch <- singleflight.Result{Val: c.val, Err: c.err, Shared: c.dups > 0} + } + } + }() + + func() { + defer func() { + if !normalReturn { + // Ideally, we would wait to take a stack trace until we've determined + // whether this is a panic or a runtime.Goexit. + // + // Unfortunately, the only way we can distinguish the two is to see + // whether the recover stopped the goroutine from terminating, and by + // the time we know that, the part of the stack trace relevant to the + // panic has been discarded. + if r := recover(); r != nil { + c.err = newPanicError(r) + } + } + }() + + c.val, c.err = fn() + normalReturn = true + }() + + if !normalReturn { + recovered = true + } +} diff --git a/coderd/externalauth_test.go b/coderd/externalauth_test.go index 4aa327313b10f..e30a81f861264 100644 --- a/coderd/externalauth_test.go +++ b/coderd/externalauth_test.go @@ -17,6 +17,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/oauth2" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd/coderdtest" @@ -519,6 +520,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) user := coderdtest.CreateFirstUser(t, client) @@ -549,6 +551,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) resp := coderdtest.RequestExternalAuthCallback(t, "github", client) @@ -563,6 +566,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) _ = coderdtest.CreateFirstUser(t, client) @@ -586,6 +590,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) maliciousHost := "https://malicious.com" @@ -619,6 +624,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) user := coderdtest.CreateFirstUser(t, client) @@ -676,10 +682,11 @@ func TestExternalAuthCallback(t *testing.T) { Expiry: dbtime.Now().Add(-time.Hour), }, }, - ID: "github", - Regex: regexp.MustCompile(`github\.com`), - Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), - NoRefresh: true, + ID: "github", + Regex: regexp.MustCompile(`github\.com`), + Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + NoRefresh: true, + RefreshGroup: new(singleflight.Group), }}, }) user := coderdtest.CreateFirstUser(t, client) @@ -726,6 +733,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) user := coderdtest.CreateFirstUser(t, client) @@ -791,6 +799,7 @@ func TestExternalAuthCallback(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) user := coderdtest.CreateFirstUser(t, client) diff --git a/coderd/promoauth/oauth2_test.go b/coderd/promoauth/oauth2_test.go index a2cb6f9bc4069..f2cd9dd83e7fa 100644 --- a/coderd/promoauth/oauth2_test.go +++ b/coderd/promoauth/oauth2_test.go @@ -13,6 +13,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/oauth2" + "golang.org/x/sync/singleflight" "github.com/coder/coder/v2/coderd/coderdtest/oidctest" "github.com/coder/coder/v2/coderd/coderdtest/promhelp" @@ -50,6 +51,7 @@ func TestInstrument(t *testing.T) { InstrumentedOAuth2Config: factory.New(id, idp.OIDCConfig(t, []string{})), ID: "test", ValidateURL: must[*url.URL](t)(idp.IssuerURL().Parse("/oauth2/userinfo")).String(), + RefreshGroup: new(singleflight.Group), } // 0 Requests before we start diff --git a/coderd/provisionerdserver/provisionerdserver_test.go b/coderd/provisionerdserver/provisionerdserver_test.go index 8f0112decf732..80ad75a493dc7 100644 --- a/coderd/provisionerdserver/provisionerdserver_test.go +++ b/coderd/provisionerdserver/provisionerdserver_test.go @@ -23,6 +23,7 @@ import ( "github.com/stretchr/testify/require" "go.opentelemetry.io/otel/trace" "golang.org/x/oauth2" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" "google.golang.org/protobuf/types/known/timestamppb" "storj.io/drpc" @@ -381,6 +382,7 @@ func TestAcquireJob(t *testing.T) { externalAuthConfigs: []*externalauth.Config{{ ID: gitAuthProvider.Id, InstrumentedOAuth2Config: &testutil.OAuth2Config{}, + RefreshGroup: new(singleflight.Group), }}, }) ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort) diff --git a/coderd/templateversions_test.go b/coderd/templateversions_test.go index c3d2153f3421e..9f5e494464fc4 100644 --- a/coderd/templateversions_test.go +++ b/coderd/templateversions_test.go @@ -15,6 +15,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/sync/errgroup" + "golang.org/x/sync/singleflight" "github.com/coder/coder/v2/coderd/audit" "github.com/coder/coder/v2/coderd/coderdtest" @@ -1009,6 +1010,7 @@ func TestTemplateVersionsExternalAuth(t *testing.T) { ID: "github", Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) user := coderdtest.CreateFirstUser(t, client) diff --git a/coderd/workspaceagents_test.go b/coderd/workspaceagents_test.go index 65c7eb1dbdf6d..f41921c7bb5f2 100644 --- a/coderd/workspaceagents_test.go +++ b/coderd/workspaceagents_test.go @@ -27,6 +27,7 @@ import ( "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" "golang.org/x/oauth2" + "golang.org/x/sync/singleflight" "golang.org/x/xerrors" "google.golang.org/protobuf/types/known/timestamppb" "tailscale.com/tailcfg" @@ -3736,6 +3737,7 @@ func TestWorkspaceAgentsExternalAuthExpiresAt(t *testing.T) { Regex: regexp.MustCompile(`.*`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), // ValidateURL intentionally omitted: tokens are always valid. + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, ownerClient) diff --git a/coderd/workspaces_test.go b/coderd/workspaces_test.go index 2c4627d3662ff..5140e7e91fea1 100644 --- a/coderd/workspaces_test.go +++ b/coderd/workspaces_test.go @@ -19,6 +19,7 @@ import ( "github.com/prometheus/client_golang/prometheus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/sync/singleflight" "cdr.dev/slog/v3" "github.com/coder/coder/v2/agent/agenttest" @@ -1501,6 +1502,7 @@ func TestCreateWorkspaceExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1553,6 +1555,7 @@ func TestCreateWorkspaceExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1601,6 +1604,7 @@ func TestCreateWorkspaceExternalAuth(t *testing.T) { Regex: regexp.MustCompile(`github\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1640,6 +1644,7 @@ func TestCreateWorkspaceExternalAuth(t *testing.T) { Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), DisplayName: "GitHub", ValidateURL: validateSrv.URL, + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) @@ -1680,6 +1685,7 @@ func TestCreateWorkspaceExternalAuth(t *testing.T) { ID: "fallback-provider", Regex: regexp.MustCompile(`fallback\.example\.com`), Type: codersdk.EnhancedExternalAuthProviderGitHub.String(), + RefreshGroup: new(singleflight.Group), }}, }) first := coderdtest.CreateFirstUser(t, client) diff --git a/enterprise/aibridged_integration_test.go b/enterprise/aibridged_integration_test.go index 3b66f9e36d1df..03f71996e7edf 100644 --- a/enterprise/aibridged_integration_test.go +++ b/enterprise/aibridged_integration_test.go @@ -18,6 +18,7 @@ import ( "go.opentelemetry.io/otel/attribute" sdktrace "go.opentelemetry.io/otel/sdk/trace" "go.opentelemetry.io/otel/sdk/trace/tracetest" + "golang.org/x/sync/singleflight" "github.com/coder/coder/v2/aibridge" "github.com/coder/coder/v2/aibridge/aibridgetest" @@ -162,6 +163,7 @@ func TestIntegration(t *testing.T) { Type: "mock", DisplayName: "Mock", MCPURL: mockMCPServer.URL, + RefreshGroup: new(singleflight.Group), }, }, },