refactor(aibridge): clean up keypool and provider error handling (#25609)

## Description

Cleans up how key pool errors are represented and how they get turned into HTTP responses. Consolidates two error types into a single type with a kind tag, and gives the response helpers in both providers consistent names.

## Changes

- Replaced the keypool sentinel and transient error struct with one error type that carries a kind and a retry-after duration.
- Updated `KeyFailoverConfig.BuildKeyPoolResponse` to take the typed key pool error, so each provider can shape the exhaustion response in its own format.
- Removed the per-provider `MarkKey` callback from `KeyFailoverConfig` since providers can rely on the shared `MarkKeyOnStatus` helper.
- Renamed the response-error helpers so OpenAI and Anthropic use the same naming.

Related to: https://linear.app/codercom/issue/AIGOV-334/aibridge-follow-ups-from-key-failover-prs

> [!NOTE]
> Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira
This commit is contained in:
Susana Ferreira
2026-05-25 18:58:29 +01:00
committed by GitHub
parent 5d178ada9f
commit 22109a54ad
20 changed files with 428 additions and 522 deletions
+14 -17
View File
@@ -2,11 +2,10 @@ package keypool
import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/aibridge/utils"
)
@@ -16,6 +15,9 @@ type KeyFailoverConfig struct {
// Pool is the key pool to walk. Nil disables key failover.
Pool *Pool
ProviderName string
Logger slog.Logger
// IsBYOK returns true when the request already carries
// user-supplied auth. BYOK requests skip key failover.
IsBYOK func(*http.Request) bool
@@ -24,14 +26,9 @@ type KeyFailoverConfig struct {
// in the format the provider expects.
InjectAuthKey func(*http.Header, string)
// MarkKey marks the key based on the upstream response.
// Returns true when the response is a key-specific error,
// causing the walker to advance and retry with the next key.
MarkKey func(ctx context.Context, key *Key, resp *http.Response) bool
// BuildExhaustedResponse returns the response sent to the
// client when the walker has no more keys to try.
BuildExhaustedResponse func(err error) *http.Response
// BuildKeyPoolResponse renders the response sent to the client
// when the walker has no more keys to try.
BuildKeyPoolResponse func(*Error) *http.Response
}
// keyFailoverTransport retries inner across the key pool on
@@ -74,12 +71,12 @@ func (t *keyFailoverTransport) RoundTrip(req *http.Request) (*http.Response, err
// Fresh walker per request, independent of other inflight requests.
walker := t.config.Pool.Walker()
for {
key, err := walker.Next()
if err != nil {
resp := t.config.BuildExhaustedResponse(err)
key, keyPoolErr := walker.Next()
if keyPoolErr != nil {
resp := t.config.BuildKeyPoolResponse(keyPoolErr)
if resp == nil {
// Fallback if BuildExhaustedResponse returns nil.
body := []byte(fmt.Sprintf(`{"error":"key pool exhausted: %s"}`, err))
// Fallback if BuildKeyPoolResponse returns nil.
body := []byte(`{"error":"key pool unavailable"}`)
resp = utils.NewJSONErrorResponse(http.StatusBadGateway, 0, body)
}
return resp, nil
@@ -97,8 +94,8 @@ func (t *keyFailoverTransport) RoundTrip(req *http.Request) (*http.Response, err
// Transport-level error, not a key issue.
return resp, rtErr
}
// MarkKey returns true on key-specific failures (e.g. 401/403/429).
if t.config.MarkKey(req.Context(), key, resp) {
// MarkKeyOnStatus returns true on key-specific failures (e.g. 401/403/429).
if MarkKeyOnStatus(req.Context(), key, resp, t.config.Logger, t.config.ProviderName) {
// Drain and retry with the next key.
_, _ = io.Copy(io.Discard, resp.Body)
_ = resp.Body.Close()
+2 -2
View File
@@ -91,8 +91,8 @@ func TestMarkKeyOnStatus(t *testing.T) {
clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0"}, clk)
require.NoError(t, err)
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
resp := &http.Response{
StatusCode: tc.statusCode,
+43 -31
View File
@@ -10,6 +10,8 @@ import (
"github.com/coder/quartz"
)
// Configuration validation type errors. These surface when the
// pool is built from invalid input.
var (
// ErrNoKeys is returned when the input is empty.
ErrNoKeys = xerrors.New("no keys provided")
@@ -18,20 +20,35 @@ var (
ErrDuplicateKey = xerrors.New("duplicate key")
)
// ErrPermanentKeyPool is returned when every key in the
// pool has been permanently marked unavailable.
var ErrPermanentKeyPool = xerrors.New("all keys permanently unavailable")
// ErrorKind classifies a runtime key-pool failure.
type ErrorKind int
// TransientKeyPoolError is returned when no key is currently
// available but at least one will recover. RetryAfter is the
// soonest remaining cooldown across the pool, or 0 if a key
// just became valid mid-walk.
type TransientKeyPoolError struct {
const (
// ErrorKindRateLimited means no key is currently available
// but at least one key will recover after a cooldown.
ErrorKindRateLimited ErrorKind = iota
// ErrorKindPermanent means every key is permanently marked
// and no key can satisfy the request.
ErrorKindPermanent
)
// Error is returned when no key is available for the
// current attempt. RetryAfter is the soonest remaining
// cooldown across the pool.
type Error struct {
Kind ErrorKind
RetryAfter time.Duration
}
func (e *TransientKeyPoolError) Error() string {
return fmt.Sprintf("all keys exhausted (retry after %s)", e.RetryAfter)
func (e *Error) Error() string {
switch e.Kind {
case ErrorKindPermanent:
return "all configured keys failed authentication"
case ErrorKindRateLimited:
return fmt.Sprintf("all configured keys are rate-limited (retry after %s)", e.RetryAfter)
default:
return "key pool error"
}
}
// KeyState represents the current state of a key in the pool.
@@ -176,20 +193,21 @@ func (k *Key) MarkPermanent() bool {
return true
}
// keyPoolError returns ErrPermanentKeyPool if every key
// is permanently unavailable, or *TransientKeyPoolError if
// at least one key is temporarily unavailable. When multiple
// keys are temporary, the smallest remaining cooldown is used
// as the retry-after.
func (p *Pool) keyPoolError() error {
// keyPoolError returns an Error summarizing why no
// key is currently available. When at least one key is
// temporary, the smallest remaining cooldown is used as the
// retry-after.
func (p *Pool) keyPoolError() *Error {
var retryAfter time.Duration
var hasCooldown bool
for i := range p.keys {
state, cooldown := p.keys[i].stateAndCooldown()
switch state {
// Recoverable now: signal transient with zero retry-after.
// Recoverable now: a key's cooldown expired between the walker's
// check and this scan. Return Retry-After: 0 to indicate that
// an immediate retry will succeed.
case KeyStateValid:
return &TransientKeyPoolError{}
return &Error{Kind: ErrorKindRateLimited}
// Recoverable later: track soonest remaining cooldown.
case KeyStateTemporary:
if !hasCooldown || cooldown < retryAfter {
@@ -201,9 +219,9 @@ func (p *Pool) keyPoolError() error {
}
}
if hasCooldown {
return &TransientKeyPoolError{RetryAfter: retryAfter}
return &Error{Kind: ErrorKindRateLimited, RetryAfter: retryAfter}
}
return ErrPermanentKeyPool
return &Error{Kind: ErrorKindPermanent}
}
// PoolState returns a snapshot of each key's state in the pool's
@@ -236,16 +254,10 @@ func (p *Pool) Walker() *Walker {
// Next returns a Key handle for the next available key without
// modifying the pool state.
//
// Returns *TransientKeyPoolError or ErrPermanentKeyPool
// when no more keys are available.
func (w *Walker) Next() (*Key, error) {
pool := w.pool
if pool == nil {
return nil, ErrPermanentKeyPool
}
for i := w.pos; i < len(pool.keys); i++ {
key := &pool.keys[i]
// Returns *Error when no more keys are available.
func (w *Walker) Next() (*Key, *Error) {
for i := w.pos; i < len(w.pool.keys); i++ {
key := &w.pool.keys[i]
if key.State() != KeyStateValid {
continue
}
@@ -255,5 +267,5 @@ func (w *Walker) Next() (*Key, error) {
}
// No keys available.
return nil, pool.keyPoolError()
return nil, w.pool.keyPoolError()
}
+94 -103
View File
@@ -1,7 +1,6 @@
package keypool_test
import (
"errors"
"sync"
"testing"
"time"
@@ -43,16 +42,15 @@ func TestNewKeyPool(t *testing.T) {
// Verify all keys are returned in order and valid.
walker := pool.Walker()
for _, expected := range tc.expectedKeys {
key, err := walker.Next()
require.NoError(t, err)
key, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
assert.Equal(t, expected, key.Value())
assert.Equal(t, keypool.KeyStateValid, key.State())
}
// No more keys available.
_, err = walker.Next()
var transient *keypool.TransientKeyPoolError
require.ErrorAs(t, err, &transient, "expected transient exhaustion: walker returned all valid keys, none marked permanent")
_, keyPoolErr := walker.Next()
require.Equal(t, &keypool.Error{Kind: keypool.ErrorKindRateLimited}, keyPoolErr, "expected rate-limited exhaustion: walker returned all valid keys, none marked permanent")
})
}
}
@@ -69,8 +67,8 @@ func TestState(t *testing.T) {
// Fresh key is valid.
name: "fresh_key_is_valid",
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
return key
},
expectedState: keypool.KeyStateValid,
@@ -79,8 +77,8 @@ func TestState(t *testing.T) {
// Active cooldown makes the key temporary.
name: "active_cooldown_is_temporary",
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
return key
},
@@ -90,8 +88,8 @@ func TestState(t *testing.T) {
// Expired cooldown returns the key to valid.
name: "expired_cooldown_is_valid",
setup: func(t *testing.T, pool *keypool.Pool, clk *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(30 * time.Second)
clk.Advance(35 * time.Second)
return key
@@ -102,8 +100,8 @@ func TestState(t *testing.T) {
// Permanent key is permanent.
name: "permanent_key",
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkPermanent()
return key
},
@@ -113,8 +111,8 @@ func TestState(t *testing.T) {
// Permanent takes precedence over active cooldown.
name: "permanent_with_cooldown_is_permanent",
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
key.MarkPermanent()
return key
@@ -152,8 +150,8 @@ func TestMarkTemporary(t *testing.T) {
name: "valid_to_temporary",
cooldown: 60 * time.Second,
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
return key
},
expectedState: keypool.KeyStateTemporary,
@@ -165,8 +163,8 @@ func TestMarkTemporary(t *testing.T) {
name: "temporary_to_temporary_extends_cooldown",
cooldown: 60 * time.Second,
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(10 * time.Second)
return key
},
@@ -179,8 +177,8 @@ func TestMarkTemporary(t *testing.T) {
name: "temporary_to_temporary_keeps_longer_cooldown",
cooldown: 10 * time.Second,
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
return key
},
@@ -192,8 +190,8 @@ func TestMarkTemporary(t *testing.T) {
name: "permanent_to_temporary_is_no_op",
cooldown: 60 * time.Second,
setup: func(t *testing.T, pool *keypool.Pool, _ *quartz.Mock) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkPermanent()
return key
},
@@ -231,8 +229,8 @@ func TestMarkPermanent(t *testing.T) {
// valid -> permanent: key becomes permanently unavailable.
name: "valid_to_permanent",
setup: func(t *testing.T, pool *keypool.Pool) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
return key
},
expectedState: keypool.KeyStatePermanent,
@@ -243,8 +241,8 @@ func TestMarkPermanent(t *testing.T) {
// to auth failure.
name: "temporary_to_permanent",
setup: func(t *testing.T, pool *keypool.Pool) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
return key
},
@@ -255,8 +253,8 @@ func TestMarkPermanent(t *testing.T) {
// permanent -> permanent: no-op, already permanent.
name: "permanent_to_permanent",
setup: func(t *testing.T, pool *keypool.Pool) *keypool.Key {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkPermanent()
return key
},
@@ -290,7 +288,7 @@ func TestWalkerNext(t *testing.T) {
setup func(t *testing.T, pool *keypool.Pool)
advance time.Duration
expectedValid []string
expectedErr error
expectedErr *keypool.Error
}{
{
// Given: key-0: valid, key-1: valid, key-2: valid.
@@ -299,7 +297,7 @@ func TestWalkerNext(t *testing.T) {
keys: []string{"key-0", "key-1", "key-2"},
setup: func(_ *testing.T, _ *keypool.Pool) {},
expectedValid: []string{"key-0", "key-1", "key-2"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary, key-1: valid, key-2: valid.
@@ -307,12 +305,12 @@ func TestWalkerNext(t *testing.T) {
name: "skips_temporary_keys",
keys: []string{"key-0", "key-1", "key-2"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
},
expectedValid: []string{"key-1", "key-2"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: permanent, key-1: permanent, key-2: valid.
@@ -321,15 +319,15 @@ func TestWalkerNext(t *testing.T) {
keys: []string{"key-0", "key-1", "key-2"},
setup: func(t *testing.T, pool *keypool.Pool) {
walker := pool.Walker()
key0, err := walker.Next()
require.NoError(t, err)
key0, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key0.MarkPermanent()
key1, err := walker.Next()
require.NoError(t, err)
key1, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key1.MarkPermanent()
},
expectedValid: []string{"key-2"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary (30s), key-1: valid.
@@ -338,13 +336,13 @@ func TestWalkerNext(t *testing.T) {
name: "expired_temporary_is_available",
keys: []string{"key-0", "key-1"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(30 * time.Second)
},
advance: 35 * time.Second,
expectedValid: []string{"key-0", "key-1"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary (zero, default 60s), key-1: valid.
@@ -353,13 +351,13 @@ func TestWalkerNext(t *testing.T) {
name: "default_cooldown_not_expired",
keys: []string{"key-0", "key-1"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(0)
},
advance: 50 * time.Second,
expectedValid: []string{"key-1"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary (zero, default 60s), key-1: valid.
@@ -368,13 +366,13 @@ func TestWalkerNext(t *testing.T) {
name: "default_cooldown_expired",
keys: []string{"key-0", "key-1"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(0)
},
advance: 65 * time.Second,
expectedValid: []string{"key-0", "key-1"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary (negative, default 60s), key-1: valid.
@@ -383,13 +381,13 @@ func TestWalkerNext(t *testing.T) {
name: "negative_cooldown_uses_default",
keys: []string{"key-0", "key-1"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(-10 * time.Second)
},
advance: 65 * time.Second,
expectedValid: []string{"key-0", "key-1"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary (60s), then marked again with shorter cooldown (10s).
@@ -398,14 +396,14 @@ func TestWalkerNext(t *testing.T) {
name: "shorter_cooldown_preserves_longer_not_expired",
keys: []string{"key-0"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
key.MarkTemporary(10 * time.Second)
},
advance: 15 * time.Second,
expectedValid: []string{},
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 45 * time.Second},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited, RetryAfter: 45 * time.Second},
},
{
// Given: key-0: temporary (60s), then marked again with shorter cooldown (10s).
@@ -414,14 +412,14 @@ func TestWalkerNext(t *testing.T) {
name: "shorter_cooldown_preserves_longer_expired",
keys: []string{"key-0"},
setup: func(t *testing.T, pool *keypool.Pool) {
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
key.MarkTemporary(60 * time.Second)
key.MarkTemporary(10 * time.Second)
},
advance: 65 * time.Second,
expectedValid: []string{"key-0"},
expectedErr: &keypool.TransientKeyPoolError{},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited},
},
{
// Given: key-0: temporary (60s), key-1: temporary (10s), key-2: temporary (30s).
@@ -431,18 +429,18 @@ func TestWalkerNext(t *testing.T) {
keys: []string{"key-0", "key-1", "key-2"},
setup: func(t *testing.T, pool *keypool.Pool) {
walker := pool.Walker()
key0, err := walker.Next()
require.NoError(t, err)
key0, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key0.MarkTemporary(60 * time.Second)
key1, err := walker.Next()
require.NoError(t, err)
key1, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key1.MarkTemporary(10 * time.Second)
key2, err := walker.Next()
require.NoError(t, err)
key2, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key2.MarkTemporary(30 * time.Second)
},
expectedValid: []string{},
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 10 * time.Second},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited, RetryAfter: 10 * time.Second},
},
{
// Given: key-0: temporary, key-1: temporary.
@@ -451,15 +449,15 @@ func TestWalkerNext(t *testing.T) {
keys: []string{"key-0", "key-1"},
setup: func(t *testing.T, pool *keypool.Pool) {
walker := pool.Walker()
key0, err := walker.Next()
require.NoError(t, err)
key0, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key0.MarkTemporary(60 * time.Second)
key1, err := walker.Next()
require.NoError(t, err)
key1, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key1.MarkTemporary(60 * time.Second)
},
expectedValid: []string{},
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 60 * time.Second},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited, RetryAfter: 60 * time.Second},
},
{
// Given: key-0: permanent, key-1: permanent.
@@ -468,15 +466,15 @@ func TestWalkerNext(t *testing.T) {
keys: []string{"key-0", "key-1"},
setup: func(t *testing.T, pool *keypool.Pool) {
walker := pool.Walker()
key0, err := walker.Next()
require.NoError(t, err)
key0, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key0.MarkPermanent()
key1, err := walker.Next()
require.NoError(t, err)
key1, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key1.MarkPermanent()
},
expectedValid: []string{},
expectedErr: keypool.ErrPermanentKeyPool,
expectedErr: &keypool.Error{Kind: keypool.ErrorKindPermanent},
},
{
// Given: key-0: permanent, key-1: temporary, key-2: permanent.
@@ -485,18 +483,18 @@ func TestWalkerNext(t *testing.T) {
keys: []string{"key-0", "key-1", "key-2"},
setup: func(t *testing.T, pool *keypool.Pool) {
walker := pool.Walker()
key0, err := walker.Next()
require.NoError(t, err)
key0, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key0.MarkPermanent()
key1, err := walker.Next()
require.NoError(t, err)
key1, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key1.MarkTemporary(60 * time.Second)
key2, err := walker.Next()
require.NoError(t, err)
key2, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
key2.MarkPermanent()
},
expectedValid: []string{},
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 60 * time.Second},
expectedErr: &keypool.Error{Kind: keypool.ErrorKindRateLimited, RetryAfter: 60 * time.Second},
},
}
@@ -516,21 +514,14 @@ func TestWalkerNext(t *testing.T) {
walker := pool.Walker()
for _, expectedKey := range tc.expectedValid {
key, err := walker.Next()
require.NoError(t, err)
key, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
assert.Equal(t, expectedKey, key.Value())
}
// After all expected keys, the walker should be exhausted.
_, err = walker.Next()
var wantTransient *keypool.TransientKeyPoolError
if errors.As(tc.expectedErr, &wantTransient) {
var got *keypool.TransientKeyPoolError
require.ErrorAs(t, err, &got)
assert.Equal(t, wantTransient.RetryAfter, got.RetryAfter)
} else {
require.ErrorIs(t, err, tc.expectedErr)
}
_, keyPoolErr := walker.Next()
require.Equal(t, tc.expectedErr, keyPoolErr)
})
}
}
@@ -595,8 +586,8 @@ func TestKeyConcurrent(t *testing.T) {
clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0"}, clk)
require.NoError(t, err)
key, err := pool.Walker().Next()
require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
const numGoroutines = 10
var wg sync.WaitGroup
@@ -628,29 +619,29 @@ func TestWalkerIndependence(t *testing.T) {
walker := pool.Walker()
// First attempt: get key-0.
key, err := walker.Next()
require.NoError(t, err)
key, keyPoolErr := walker.Next()
require.Nil(t, keyPoolErr)
assert.Equal(t, "key-0", key.Value())
// Simulate 429: mark key-0 temporary.
key.MarkTemporary(60 * time.Second)
// Second attempt: walker advances to key-1.
key, err = walker.Next()
require.NoError(t, err)
key, keyPoolErr = walker.Next()
require.Nil(t, keyPoolErr)
assert.Equal(t, "key-1", key.Value())
// Simulate 401: mark key-1 permanent.
key.MarkPermanent()
// Third attempt: walker advances to key-2.
key, err = walker.Next()
require.NoError(t, err)
key, keyPoolErr = walker.Next()
require.Nil(t, keyPoolErr)
assert.Equal(t, "key-2", key.Value())
// A new walker should skip key-0 (temporary) and key-1
// (permanent), and return key-2.
key2, err := pool.Walker().Next()
require.NoError(t, err)
key2, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr)
assert.Equal(t, "key-2", key2.Value())
}