fix: follow-up fixes for conditional VCS requests (#27711)

Follow-ups from #27627 

- Memoizes `Config.Git()` with a mutex so the provider's ETag response
cache survives across calls. Only successful construction is cached;
errors are retried.
- Moves the HTTP client onto `Config.HTTPClient`, wired through
`ConvertConfig`, so `Git()` no longer takes a per-call argument that
would be silently ignored after memoization.
- `newGitHub` and `newGitLab` now return `(Provider, error)`,
eliminating the typed-nil-interface class in `gitprovider.New` rather
than the single instance.
- Gates the 304 branch on a `haveCached` flag instead of a nil body
check.
- Only caches bodies that decode successfully, preventing poisoned
entries.
- Keys the response cache on the full token digest rather than a
truncated prefix.
- Tests added: `TestConfigGitMemoizesProvider`,
`TestConfigGitRetriesOnConstructorError`,
`TestGitLabConstructorErrorReturnsNilInterface`,
`TestResponseCacheStore`,
`TestConditionalRequestReuse/MalformedResponseNotCached`;
`TestConvertYAML/CustomScopesAndEndpoint` now asserts
`Config.HTTPClient` wiring.

Follow-ups tracked in #28139, #28140, #28141, #28142.

> 🤖 Generated by Coder Agents on behalf of @johnstcn.
This commit is contained in:
Cian Johnston
2026-08-18 09:00:20 +01:00
committed by GitHub
parent b674d40d39
commit 6079c514ee
11 changed files with 273 additions and 52 deletions
+35 -7
View File
@@ -114,6 +114,10 @@ type Config struct {
// (e.g., "https://api.github.com" for GitHub). Derived from
// defaults when not explicitly configured.
APIBaseURL string
// If nil, http.DefaultClient is used. The value is read once at
// the first successful Git() call; later assignments have no
// effect because the provider is memoized.
HTTPClient *http.Client
// AppInstallURL is for GitHub App's (and hopefully others eventually)
// to provide a link to install the app. There's installation
// of the application, and user authentication. It's possible
@@ -159,18 +163,40 @@ type Config struct {
// RefreshGroup deduplicates concurrent requests.
RefreshGroup SingleflightGroup
gitProviderMu sync.Mutex
// gitProvider memoizes the provider so the GitHub ETag response
// cache survives across Git calls.
gitProvider gitprovider.Provider
}
// Git returns a Provider for this config if the provider type is a
// supported git hosting provider. Returns (nil, nil) for non-git
// providers (e.g. Slack, JFrog). Returns a non-nil error if provider
// construction fails.
func (c *Config) Git(client *http.Client) (gitprovider.Provider, error) {
// Git returns a Provider for this config. It returns (nil, nil) when
// this config's type has no provider implementation, which covers both
// non-git types (e.g. Slack, JFrog) and git types that are not
// implemented yet (bitbucket-*, azure-devops*, gitea). Callers cannot
// distinguish the two cases from the return values. Returns a non-nil
// error if provider construction fails.
//
// The provider is built on the first successful call and cached for
// the lifetime of the Config, so its in-memory response cache
// survives across calls. The provider uses c.HTTPClient for API
// requests; if c.HTTPClient is nil, http.DefaultClient is used.
func (c *Config) Git() (gitprovider.Provider, error) {
norm := strings.ToLower(c.Type)
if !codersdk.EnhancedExternalAuthProvider(norm).Git() {
return nil, nil //nolint:nilnil // nil provider means non-git type, not an error
}
return gitprovider.New(norm, c.APIBaseURL, client)
c.gitProviderMu.Lock()
defer c.gitProviderMu.Unlock()
if c.gitProvider != nil {
return c.gitProvider, nil
}
p, err := gitprovider.New(norm, c.APIBaseURL, c.HTTPClient)
if err != nil {
return nil, err
}
c.gitProvider = p
return c.gitProvider, nil
}
// GenerateTokenExtra generates the extra token data to store in the database.
@@ -912,7 +938,8 @@ func (c *DeviceAuth) formatDeviceCodeURL() (string, error) {
// ConvertConfig converts the SDK configuration entry format
// to the parsed and ready-to-consume in coderd provider type.
func ConvertConfig(ctx context.Context, logger slog.Logger, instrument *promoauth.Factory, entries []codersdk.ExternalAuthConfig, accessURL *url.URL) ([]*Config, error) {
// If httpClient is nil, http.DefaultClient is used.
func ConvertConfig(ctx context.Context, logger slog.Logger, instrument *promoauth.Factory, entries []codersdk.ExternalAuthConfig, accessURL *url.URL, httpClient *http.Client) ([]*Config, error) {
ids := map[string]struct{}{}
configs := []*Config{}
for _, entry := range entries {
@@ -1013,6 +1040,7 @@ func ConvertConfig(ctx context.Context, logger slog.Logger, instrument *promoaut
ClientSecret: entry.ClientSecret,
Regex: regex,
APIBaseURL: entry.APIBaseURL,
HTTPClient: httpClient,
Type: entry.Type,
NoRefresh: entry.NoRefresh,
ValidateURL: entry.ValidateURL,
+79 -7
View File
@@ -37,11 +37,80 @@ import (
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/database/dbtime"
"github.com/coder/coder/v2/coderd/externalauth"
"github.com/coder/coder/v2/coderd/externalauth/gitprovider"
"github.com/coder/coder/v2/coderd/promoauth"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
)
func TestConfigGitMemoizesProvider(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
const etag = `"config-git-memo-etag"`
var conditionalRequests atomic.Int64
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if inm := r.Header.Get("If-None-Match"); inm != "" {
conditionalRequests.Add(1)
assert.Equal(t, etag, inm)
w.Header().Set("ETag", etag)
w.WriteHeader(http.StatusNotModified)
return
}
w.Header().Set("Content-Type", "application/json")
w.Header().Set("ETag", etag)
_, _ = w.Write([]byte(`[]`))
}))
defer srv.Close()
cfg := &externalauth.Config{
Type: string(codersdk.EnhancedExternalAuthProviderGitHub),
APIBaseURL: srv.URL + "/api/v3",
HTTPClient: srv.Client(),
}
gp1, err := cfg.Git()
require.NoError(t, err)
require.NotNil(t, gp1)
branch := gitprovider.BranchRef{Owner: "owner", Repo: "repo", Branch: "feat"}
// Cold poll: populates the provider's ETag cache.
_, err = gp1.ResolveBranchPullRequest(ctx, "test-token", branch)
require.NoError(t, err)
// Re-resolve the provider, as the worker does on every poll.
gp2, err := cfg.Git()
require.NoError(t, err)
require.NotNil(t, gp2)
assert.Same(t, gp1, gp2, "Git must return the same provider instance so its ETag cache survives across calls")
_, err = gp2.ResolveBranchPullRequest(ctx, "test-token", branch)
require.NoError(t, err)
assert.Equal(t, int64(1), conditionalRequests.Load(), "second poll should have revalidated with If-None-Match using the cache from the first poll")
}
func TestConfigGitRetriesOnConstructorError(t *testing.T) {
t.Parallel()
cfg := &externalauth.Config{
Type: string(codersdk.EnhancedExternalAuthProviderGitLab),
APIBaseURL: "://invalid",
}
_, err1 := cfg.Git()
require.Error(t, err1)
_, err2 := cfg.Git()
require.Error(t, err2)
// A memoized error would be the same instance; a retried
// construction produces a fresh error each call.
require.NotErrorIs(t, err2, err1, "construction errors must be retried, not memoized")
}
func TestRefreshToken(t *testing.T) {
t.Parallel()
expired := time.Now().Add(time.Hour * -1)
@@ -1084,7 +1153,7 @@ func TestRefreshTokenWithScopes(t *testing.T) {
AuthURL: "https://login.microsoftonline.com/tenant/oauth2/authorize",
TokenURL: "https://login.microsoftonline.com/tenant/oauth2/token",
Scopes: scopes,
}}, &url.URL{Scheme: "https", Host: "coder.example.com"})
}}, &url.URL{Scheme: "https", Host: "coder.example.com"}, nil)
require.NoError(t, err)
return configs[0]
}
@@ -1209,7 +1278,7 @@ func TestValidateToken(t *testing.T) {
ClientID: "id",
ClientSecret: "secret",
ValidateURL: validateURL,
}}, &url.URL{})
}}, &url.URL{}, nil)
require.NoError(t, err)
return configs[0], logs
}
@@ -1614,7 +1683,7 @@ func TestExchangeWithClientSecret(t *testing.T) {
Type: codersdk.EnhancedExternalAuthProviderJFrog.String(),
ClientID: "id",
ClientSecret: "secret",
}}, &url.URL{})
}}, &url.URL{}, nil)
require.NoError(t, err)
config := configs[0]
@@ -1740,7 +1809,7 @@ func TestConvertYAML(t *testing.T) {
}} {
t.Run(tc.Name, func(t *testing.T) {
t.Parallel()
output, err := externalauth.ConvertConfig(context.Background(), testutil.Logger(t), instrument, tc.Input, &url.URL{})
output, err := externalauth.ConvertConfig(context.Background(), testutil.Logger(t), instrument, tc.Input, &url.URL{}, nil)
if tc.Error != "" {
require.Error(t, err)
require.Contains(t, err.Error(), tc.Error)
@@ -1752,6 +1821,7 @@ func TestConvertYAML(t *testing.T) {
t.Run("CustomScopesAndEndpoint", func(t *testing.T) {
t.Parallel()
client := new(http.Client)
config, err := externalauth.ConvertConfig(context.Background(), testutil.Logger(t), instrument, []codersdk.ExternalAuthConfig{{
Type: string(codersdk.EnhancedExternalAuthProviderGitLab),
ClientID: "id",
@@ -1760,9 +1830,10 @@ func TestConvertYAML(t *testing.T) {
TokenURL: "https://token.com",
RedirectURL: "https://redirect.com",
Scopes: []string{"read"},
}}, &url.URL{Scheme: "https", Host: "default.com"})
}}, &url.URL{Scheme: "https", Host: "default.com"}, client)
require.NoError(t, err)
require.Equal(t, "https://auth.com?client_id=id&redirect_uri=https%3A%2F%2Fredirect.com%2Fexternal-auth%2Fgitlab%2Fcallback&response_type=code&scope=read", config[0].AuthCodeURL(""))
assert.Same(t, client, config[0].HTTPClient, "ConvertConfig must wire the provided client onto every Config")
})
t.Run("RevokeTimeoutSet", func(t *testing.T) {
@@ -1771,7 +1842,7 @@ func TestConvertYAML(t *testing.T) {
Type: string(codersdk.EnhancedExternalAuthProviderGitLab),
ClientID: "id",
ClientSecret: "secret",
}}, &url.URL{})
}}, &url.URL{}, nil)
require.NoError(t, err)
require.Equal(t, 10*time.Second, configs[0].RevokeTimeout)
})
@@ -1784,7 +1855,7 @@ func TestConvertYAML(t *testing.T) {
ClientSecret: "secret",
AuthURL: "https://gitlab.corp.com/oauth/authorize",
TokenURL: "https://gitlab.corp.com/oauth/token",
}}, &url.URL{})
}}, &url.URL{}, nil)
require.NoError(t, err)
require.Len(t, configs, 1)
require.Equal(t, "https://gitlab.corp.com/api/v4", configs[0].APIBaseURL)
@@ -1966,6 +2037,7 @@ func TestApplyDefaultsToConfig_CaseInsensitive(t *testing.T) {
ClientSecret: "test-secret",
}},
accessURL,
nil,
)
require.NoError(t, err)
require.Len(t, configs, 1)
+14 -9
View File
@@ -1,6 +1,7 @@
package gitprovider
import (
"bytes"
"container/list"
"crypto/sha256"
"encoding/hex"
@@ -13,7 +14,10 @@ import (
const (
// defaultResponseCacheEntries is the maximum number of cached
// responses retained. Once exceeded, the least-recently-used
// entry is evicted.
// entry is evicted. Sized above the gitsync worker's steady-state
// working set: defaultBatchSize (50) rows re-acquired every
// DiffStatusTTL (120s) over defaultInterval (10s) ticks, at 2-3
// cache keys per row. See coderd/x/gitsync.
defaultResponseCacheEntries = 2048
// maxCachedBodyBytes is the largest response body that will be
@@ -60,7 +64,8 @@ func newResponseCache(maxSize int) *responseCache {
}
// load returns the cached ETag and body for key, if present, and
// marks the entry as most-recently-used.
// marks the entry as most-recently-used. The returned body aliases
// the cache's copy and must not be mutated.
func (c *responseCache) load(key string) (etag string, body []byte, ok bool) {
c.mu.Lock()
defer c.mu.Unlock()
@@ -82,6 +87,7 @@ func (c *responseCache) store(key, etag string, body []byte) {
return
}
stored := bytes.Clone(body)
c.mu.Lock()
defer c.mu.Unlock()
@@ -89,15 +95,14 @@ func (c *responseCache) store(key, etag string, body []byte) {
c.ll.MoveToFront(elem)
cr := elem.Value.(*cachedResponse)
cr.etag = etag
cr.body = body
// Replace the body slice entirely; never write into the
// existing slice in place. A concurrent reader may hold a
// reference to the old slice while json.Unmarshal is reading
// it.
cr.body = stored
return
}
// Copy the body so we never retain a slice that the caller may
// later reuse or mutate.
stored := make([]byte, len(body))
copy(stored, body)
elem := c.ll.PushFront(&cachedResponse{key: key, etag: etag, body: stored})
c.entries[key] = elem
@@ -123,5 +128,5 @@ func (c *responseCache) evictOldest() {
// another, without keeping raw credentials in memory.
func responseCacheKey(requestURL, token string) string {
sum := sha256.Sum256([]byte(token))
return requestURL + "\x00" + hex.EncodeToString(sum[:8])
return requestURL + "\x00" + hex.EncodeToString(sum[:])
}
+14 -14
View File
@@ -37,7 +37,7 @@ type githubProvider struct {
repositorySSHPathPattern *regexp.Regexp
}
func newGitHub(apiBaseURL string, httpClient *http.Client, clock quartz.Clock) *githubProvider {
func newGitHub(apiBaseURL string, httpClient *http.Client, clock quartz.Clock) (Provider, error) {
if apiBaseURL == "" {
apiBaseURL = defaultGitHubAPIBaseURL
}
@@ -72,7 +72,7 @@ func newGitHub(apiBaseURL string, httpClient *http.Client, clock quartz.Clock) *
repositorySSHPathPattern: regexp.MustCompile(
`^(?:ssh://)?git@` + escapedHost + `[:/]([A-Za-z0-9_.-]+)/([A-Za-z0-9_.-]+?)(?:\.git)?/?$`,
),
}
}, nil
}
// deriveWebBaseURL converts a GitHub API base URL to the
@@ -412,12 +412,13 @@ func (g *githubProvider) decodeJSON(
// changed, which is cheaper than a full body and does not count
// against the primary REST rate limit.
cacheKey := responseCacheKey(requestURL, token)
var cachedBody []byte
if g.cache != nil {
if etag, body, ok := g.cache.load(cacheKey); ok {
req.Header.Set("If-None-Match", etag)
cachedBody = body
}
var (
cachedBody []byte
haveCached bool
)
if etag, body, ok := g.cache.load(cacheKey); ok {
req.Header.Set("If-None-Match", etag)
cachedBody, haveCached = body, true
}
resp, err := g.httpClient.Do(req)
@@ -427,7 +428,7 @@ func (g *githubProvider) decodeJSON(
defer resp.Body.Close()
// Nothing changed since the cached response: reuse the stored body.
if resp.StatusCode == http.StatusNotModified && cachedBody != nil {
if resp.StatusCode == http.StatusNotModified && haveCached {
if err := json.Unmarshal(cachedBody, dest); err != nil {
return xerrors.Errorf("decode cached github response: %w", err)
}
@@ -457,14 +458,13 @@ func (g *githubProvider) decodeJSON(
return xerrors.Errorf("read github response: %w", err)
}
// Cache the validator so the next poll can be made conditional.
if g.cache != nil {
g.cache.store(cacheKey, resp.Header.Get("ETag"), body)
}
if err := json.Unmarshal(body, dest); err != nil {
return xerrors.Errorf("decode github response: %w", err)
}
// Only cache bodies we could successfully decode, so a malformed
// response does not poison the cache.
g.cache.store(cacheKey, resp.Header.Get("ETag"), body)
return nil
}
@@ -1061,4 +1061,40 @@ func TestConditionalRequestReuse(t *testing.T) {
assert.Equal(t, 0, conditionalRequests,
"a different token must not send If-None-Match from another token's cache")
})
t.Run("MalformedResponseNotCached", func(t *testing.T) {
t.Parallel()
const etag = `"poison-etag"`
var conditionalRequests int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("If-None-Match") != "" {
conditionalRequests++
w.Header().Set("ETag", etag)
w.WriteHeader(http.StatusNotModified)
return
}
w.Header().Set("Content-Type", "application/json")
w.Header().Set("ETag", etag)
// A body that cannot be decoded as a pull request list.
_, _ = w.Write([]byte(`<html>proxy error</html>`))
}))
defer srv.Close()
gp, err := gitprovider.New("github", srv.URL+"/api/v3", srv.Client())
require.NoError(t, err)
require.NotNil(t, gp)
branch := gitprovider.BranchRef{Owner: "owner", Repo: "repo", Branch: "feat"}
_, err = gp.ResolveBranchPullRequest(context.Background(), "test-token", branch)
require.Error(t, err)
_, err = gp.ResolveBranchPullRequest(context.Background(), "test-token", branch)
require.Error(t, err)
assert.Equal(t, 0, conditionalRequests,
"a body that failed to decode must not be cached with its ETag")
})
}
+1 -1
View File
@@ -25,7 +25,7 @@ type gitlabProvider struct {
clock quartz.Clock
}
func newGitLab(baseURL string, httpClient *http.Client, clock quartz.Clock) (*gitlabProvider, error) {
func newGitLab(baseURL string, httpClient *http.Client, clock quartz.Clock) (Provider, error) {
if baseURL == "" {
baseURL = "https://gitlab.com"
}
@@ -262,6 +262,18 @@ func TestGitLabResolveBranchPullRequest(t *testing.T) {
})
}
func TestGitLabConstructorErrorReturnsNilInterface(t *testing.T) {
t.Parallel()
// A construction error must return a true nil Provider, not a
// non-nil interface boxing a nil *gitlabProvider. The gp == nil
// check is what callers write, and require.Nil alone would pass
// for a non-nil interface holding a nil pointer.
gp, err := gitprovider.New("gitlab", "://invalid/", nil)
require.Error(t, err)
require.True(t, gp == nil, "provider must be nil when construction fails, got %T", gp)
}
func TestGitLabRateLimit(t *testing.T) {
t.Parallel()
@@ -192,7 +192,7 @@ func New(providerType string, apiBaseURL string, httpClient *http.Client, opts .
switch providerType {
case "github":
return newGitHub(apiBaseURL, httpClient, o.clock), nil
return newGitHub(apiBaseURL, httpClient, o.clock)
case "gitlab":
return newGitLab(apiBaseURL, httpClient, o.clock)
default:
@@ -207,9 +207,6 @@ func New(providerType string, apiBaseURL string, httpClient *http.Client, opts .
// resetHeader (unix timestamp). Returns zero if no recognizable header
// is present.
func parseRetryAfter(h http.Header, resetHeader string, clk quartz.Clock) time.Duration {
if clk == nil {
clk = quartz.NewReal()
}
// Retry-After header: seconds until retry.
if ra := h.Get("Retry-After"); ra != "" {
if secs, err := strconv.Atoi(ra); err == nil {
@@ -236,9 +233,6 @@ func checkRateLimitError(resp *http.Response, clk quartz.Clock, resetHeader stri
if resp.StatusCode != http.StatusForbidden && resp.StatusCode != http.StatusTooManyRequests {
return nil
}
if clk == nil {
clk = quartz.NewReal()
}
retryAfter := parseRetryAfter(resp.Header, resetHeader, clk)
if retryAfter <= 0 {
return nil
@@ -7,6 +7,7 @@ import (
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/coder/quartz"
)
@@ -112,13 +113,85 @@ func TestParseRetryAfter(t *testing.T) {
d := parseRetryAfter(h, "X-Ratelimit-Reset", clk)
assert.Equal(t, 60*time.Second, d)
})
}
t.Run("NilClock", func(t *testing.T) {
func TestResponseCacheStore(t *testing.T) {
t.Parallel()
// Stores the same key twice using a buffer the caller mutates
// between stores, and verifies load returns the bodies passed at
// store time rather than the mutated buffer.
t.Run("UpdateReplacesBody", func(t *testing.T) {
t.Parallel()
h := http.Header{}
h.Set("Retry-After", "1")
d := parseRetryAfter(h, "X-Ratelimit-Reset", nil)
assert.Equal(t, time.Second, d)
cache := newResponseCache(4)
const key = "k"
buf := []byte(`{"v":1}`)
cache.store(key, `"etag-1"`, buf)
// Mutate the caller's buffer: the cache must hold its own copy.
for i := range buf {
buf[i] = 'X'
}
etag, body, ok := cache.load(key)
require.True(t, ok)
assert.Equal(t, `"etag-1"`, etag)
assert.Equal(t, `{"v":1}`, string(body))
// Reuse the same buffer for a second store of the same key.
buf = append(buf[:0], `{"v":2}`...)
cache.store(key, `"etag-2"`, buf)
for i := range buf {
buf[i] = 'Y'
}
etag, body, ok = cache.load(key)
require.True(t, ok)
assert.Equal(t, `"etag-2"`, etag)
assert.Equal(t, `{"v":2}`, string(body))
})
// Fills the cache past maxSize and verifies the
// least-recently-used entry is evicted, not merely the
// oldest-inserted one.
t.Run("EvictsLeastRecentlyUsed", func(t *testing.T) {
t.Parallel()
cache := newResponseCache(2)
cache.store("a", `"etag-a"`, []byte(`{"k":"a"}`))
cache.store("b", `"etag-b"`, []byte(`{"k":"b"}`))
// Access "a" so "b" becomes the least-recently-used entry.
_, _, ok := cache.load("a")
require.True(t, ok)
// This third store exceeds maxSize and must evict "b",
// not the older "a".
cache.store("c", `"etag-c"`, []byte(`{"k":"c"}`))
_, _, ok = cache.load("b")
assert.False(t, ok, "least-recently-used entry must be evicted")
etag, body, ok := cache.load("a")
require.True(t, ok, "recently-accessed entry must survive eviction")
assert.Equal(t, `"etag-a"`, etag)
assert.Equal(t, `{"k":"a"}`, string(body))
etag, body, ok = cache.load("c")
require.True(t, ok)
assert.Equal(t, `"etag-c"`, etag)
assert.Equal(t, `{"k":"c"}`, string(body))
})
// Bodies larger than maxCachedBodyBytes are not stored, so a
// single oversized response cannot unbound the cache.
t.Run("OversizedBodyNotStored", func(t *testing.T) {
t.Parallel()
cache := newResponseCache(4)
cache.store("big", `"etag-big"`, make([]byte, maxCachedBodyBytes+1))
_, _, ok := cache.load("big")
assert.False(t, ok, "bodies exceeding maxCachedBodyBytes must not be cached")
})
}