mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add automatic key failover for AI Bridge Anthropic (#24836)
## Description Adds automatic key failover for centralized Anthropic provider. When a key pool is configured, each upstream call walks the pool and tries keys in order until one succeeds or the pool is exhausted. Keys are marked **temporary** on 429 (with cooldown from `Retry-After`) and **permanent** on 401/403. Errors that aren't key-specific don't trigger failover. Each agentic-loop iteration gets its own fresh walker, so a tool-call continuation can fail over independently of the initial request. BYOK is unchanged: BYOK requests run as a single attempt with no failover. ## Changes - `config.Anthropic` carries a `KeyPool`. `Key` remains for BYOK X-Api-Key set per interception. - Blocking interceptor: walks the pool, marks keys on key-specific failures, returns on first success or non-failover error. - Streaming interceptor: per-iteration walker. Pre-stream failures fail over to the next key; mid-stream errors are relayed as SSE events. - New `keypool` error types: `TransientExhaustionError` (carries soonest cooldown) and `ErrPermanentExhaustion`. Replace the prior `ErrAllKeysExhausted`. - Error responses now consistently include the outer `"type": "error"` field. ## Related Issues Related to: https://github.com/coder/internal/issues/1446 Related to: https://linear.app/codercom/issue/AIGOV-197/aibridge-automatic-key-failover-for-bridged-and-passthrough-routes ## Follow-up PRs - Bedrock multi-key support. - Refactor provider vs interceptor config separation. - Record the actually-used key in the interception credential hint after failover. > [!NOTE] > Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira
This commit is contained in:
@@ -0,0 +1,37 @@
|
||||
package keypool
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ParseRetryAfter extracts the cooldown duration from response
|
||||
// headers. It prefers the OpenAI-specific "retry-after-ms"
|
||||
// header (milliseconds) over the standard "Retry-After" header
|
||||
// (seconds). Returns zero if neither header is present or
|
||||
// parseable. The HTTP-date form of "Retry-After" is not parsed.
|
||||
func ParseRetryAfter(resp *http.Response) time.Duration {
|
||||
if resp == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
// OpenAI convention: millisecond precision.
|
||||
if val := resp.Header.Get("retry-after-ms"); val != "" {
|
||||
ms, err := strconv.ParseFloat(strings.TrimSpace(val), 64)
|
||||
if err == nil && ms > 0 {
|
||||
return time.Duration(ms * float64(time.Millisecond))
|
||||
}
|
||||
}
|
||||
|
||||
// Standard header: seconds.
|
||||
if val := resp.Header.Get("Retry-After"); val != "" {
|
||||
seconds, err := strconv.Atoi(strings.TrimSpace(val))
|
||||
if err == nil && seconds > 0 {
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
}
|
||||
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package keypool_test
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/coder/coder/v2/aibridge/keypool"
|
||||
)
|
||||
|
||||
func TestParseRetryAfter(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
headers map[string]string
|
||||
nilResponse bool
|
||||
expected time.Duration
|
||||
}{
|
||||
// nil response.
|
||||
{
|
||||
name: "nil_response",
|
||||
nilResponse: true,
|
||||
expected: 0,
|
||||
},
|
||||
// No headers set.
|
||||
{
|
||||
name: "no_headers",
|
||||
headers: nil,
|
||||
expected: 0,
|
||||
},
|
||||
// retry-after-ms (OpenAI, preferred).
|
||||
{
|
||||
name: "openai_retry_after_ms",
|
||||
headers: map[string]string{"retry-after-ms": "2500"},
|
||||
expected: 2500 * time.Millisecond,
|
||||
},
|
||||
{
|
||||
name: "whitespace_trimmed_ms",
|
||||
headers: map[string]string{"retry-after-ms": " 1500 "},
|
||||
expected: 1500 * time.Millisecond,
|
||||
},
|
||||
{
|
||||
name: "negative_ms_returns_zero",
|
||||
headers: map[string]string{"retry-after-ms": "-100"},
|
||||
expected: 0,
|
||||
},
|
||||
// Retry-After (standard, seconds).
|
||||
{
|
||||
name: "standard_retry_after_seconds",
|
||||
headers: map[string]string{"Retry-After": "60"},
|
||||
expected: 60 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "whitespace_trimmed_seconds",
|
||||
headers: map[string]string{"Retry-After": " 30 "},
|
||||
expected: 30 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "zero_seconds_returns_zero",
|
||||
headers: map[string]string{"Retry-After": "0"},
|
||||
expected: 0,
|
||||
},
|
||||
{
|
||||
name: "negative_seconds_returns_zero",
|
||||
headers: map[string]string{"Retry-After": "-5"},
|
||||
expected: 0,
|
||||
},
|
||||
// Both headers set: precedence and fallback.
|
||||
{
|
||||
name: "prefers_retry_after_ms_over_standard",
|
||||
headers: map[string]string{
|
||||
"retry-after-ms": "1500",
|
||||
"Retry-After": "30",
|
||||
},
|
||||
expected: 1500 * time.Millisecond,
|
||||
},
|
||||
{
|
||||
name: "falls_back_to_standard_when_ms_invalid",
|
||||
headers: map[string]string{"retry-after-ms": "invalid", "Retry-After": "10"},
|
||||
expected: 10 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "zero_ms_falls_back_to_standard",
|
||||
headers: map[string]string{"retry-after-ms": "0", "Retry-After": "5"},
|
||||
expected: 5 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "zero_ms_and_zero_seconds_return_zero",
|
||||
headers: map[string]string{"retry-after-ms": "0", "Retry-After": "0"},
|
||||
expected: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
var resp *http.Response
|
||||
if !tc.nilResponse {
|
||||
resp = &http.Response{Header: make(http.Header)}
|
||||
for key, val := range tc.headers {
|
||||
resp.Header.Set(key, val)
|
||||
}
|
||||
}
|
||||
assert.Equal(t, tc.expected, keypool.ParseRetryAfter(resp))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package keypool
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/aibridge/utils"
|
||||
)
|
||||
|
||||
// MarkKeyOnStatus marks key based on a key-specific HTTP
|
||||
// status code from resp (429 for temporary, 401 or 403 for
|
||||
// permanent). Returns true if the status was a key-specific
|
||||
// failover trigger so callers can retry with the next key.
|
||||
func MarkKeyOnStatus(
|
||||
ctx context.Context,
|
||||
key *Key,
|
||||
resp *http.Response,
|
||||
logger slog.Logger,
|
||||
providerName string,
|
||||
) bool {
|
||||
if resp == nil {
|
||||
return false
|
||||
}
|
||||
statusCode := resp.StatusCode
|
||||
switch statusCode {
|
||||
case http.StatusTooManyRequests:
|
||||
cooldown := ParseRetryAfter(resp)
|
||||
if cooldown <= 0 {
|
||||
cooldown = defaultCooldown
|
||||
}
|
||||
if key.MarkTemporary(cooldown) {
|
||||
logger.Info(ctx, "key marked temporary",
|
||||
slog.F("provider", providerName),
|
||||
slog.F("api_key_hint", utils.MaskSecret(key.Value())),
|
||||
slog.F("status", statusCode),
|
||||
slog.F("cooldown", cooldown))
|
||||
}
|
||||
return true
|
||||
case http.StatusUnauthorized, http.StatusForbidden:
|
||||
if key.MarkPermanent() {
|
||||
logger.Warn(ctx, "key marked permanent",
|
||||
slog.F("provider", providerName),
|
||||
slog.F("api_key_hint", utils.MaskSecret(key.Value())),
|
||||
slog.F("status", statusCode))
|
||||
}
|
||||
return true
|
||||
default:
|
||||
logger.Debug(ctx, "status is not a key failover trigger",
|
||||
slog.F("provider", providerName),
|
||||
slog.F("status", statusCode))
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package keypool_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/aibridge/keypool"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func TestMarkKeyOnStatus(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
headers map[string]string
|
||||
expectedReturn bool
|
||||
expectedState keypool.KeyState
|
||||
expectedCooldown time.Duration
|
||||
}{
|
||||
{
|
||||
// 429 with standard Retry-After header (seconds).
|
||||
name: "429_with_retry_after_seconds",
|
||||
statusCode: http.StatusTooManyRequests,
|
||||
headers: map[string]string{"Retry-After": "5"},
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStateTemporary,
|
||||
expectedCooldown: 5 * time.Second,
|
||||
},
|
||||
{
|
||||
// 429 with retry-after-ms header (milliseconds).
|
||||
name: "429_with_retry_after_ms",
|
||||
statusCode: http.StatusTooManyRequests,
|
||||
headers: map[string]string{"retry-after-ms": "1500"},
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStateTemporary,
|
||||
expectedCooldown: 1500 * time.Millisecond,
|
||||
},
|
||||
{
|
||||
// 429 without headers falls back to default cooldown.
|
||||
name: "429_no_headers_uses_default",
|
||||
statusCode: http.StatusTooManyRequests,
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStateTemporary,
|
||||
expectedCooldown: 60 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "401_marks_permanent",
|
||||
statusCode: http.StatusUnauthorized,
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStatePermanent,
|
||||
},
|
||||
{
|
||||
name: "403_marks_permanent",
|
||||
statusCode: http.StatusForbidden,
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStatePermanent,
|
||||
},
|
||||
{
|
||||
name: "200_does_not_mark",
|
||||
statusCode: http.StatusOK,
|
||||
expectedReturn: false,
|
||||
expectedState: keypool.KeyStateValid,
|
||||
},
|
||||
{
|
||||
name: "500_does_not_mark",
|
||||
statusCode: http.StatusInternalServerError,
|
||||
expectedReturn: false,
|
||||
expectedState: keypool.KeyStateValid,
|
||||
},
|
||||
{
|
||||
// 529 is the Anthropic overloaded status, handled by
|
||||
// the circuit breaker, not key failover.
|
||||
name: "529_does_not_mark",
|
||||
statusCode: 529,
|
||||
expectedReturn: false,
|
||||
expectedState: keypool.KeyStateValid,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
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)
|
||||
|
||||
resp := &http.Response{
|
||||
StatusCode: tc.statusCode,
|
||||
Header: make(http.Header),
|
||||
}
|
||||
for k, v := range tc.headers {
|
||||
resp.Header.Set(k, v)
|
||||
}
|
||||
|
||||
got := keypool.MarkKeyOnStatus(
|
||||
context.Background(),
|
||||
key,
|
||||
resp,
|
||||
// 401 and 403 cases legitimately log at error
|
||||
// level when marking a key permanent.
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
"test",
|
||||
)
|
||||
|
||||
assert.Equal(t, tc.expectedReturn, got)
|
||||
assert.Equal(t, tc.expectedState, key.State())
|
||||
|
||||
// Verify cooldown was set to the expected duration:
|
||||
// advancing by exactly that amount returns the key
|
||||
// to valid.
|
||||
if tc.expectedCooldown > 0 {
|
||||
clk.Advance(tc.expectedCooldown)
|
||||
assert.Equal(t, keypool.KeyStateValid, key.State())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package keypool
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -15,11 +16,24 @@ var (
|
||||
// ErrDuplicateKey is returned when the input contains
|
||||
// duplicate key values.
|
||||
ErrDuplicateKey = xerrors.New("duplicate key")
|
||||
// ErrAllKeysExhausted is returned when the walker has visited
|
||||
// every key in the pool and none are available.
|
||||
ErrAllKeysExhausted = xerrors.New("all keys exhausted")
|
||||
)
|
||||
|
||||
// ErrPermanentKeyPool is returned when every key in the
|
||||
// pool has been permanently marked unavailable.
|
||||
var ErrPermanentKeyPool = xerrors.New("all keys permanently unavailable")
|
||||
|
||||
// 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 {
|
||||
RetryAfter time.Duration
|
||||
}
|
||||
|
||||
func (e *TransientKeyPoolError) Error() string {
|
||||
return fmt.Sprintf("all keys exhausted (retry after %s)", e.RetryAfter)
|
||||
}
|
||||
|
||||
// KeyState represents the current state of a key in the pool.
|
||||
type KeyState int
|
||||
|
||||
@@ -101,6 +115,22 @@ func (k *Key) State() KeyState {
|
||||
return KeyStateValid
|
||||
}
|
||||
|
||||
// stateAndCooldown returns the key's state and remaining
|
||||
// cooldown as a single atomic snapshot.
|
||||
func (k *Key) stateAndCooldown() (KeyState, time.Duration) {
|
||||
k.mu.RLock()
|
||||
defer k.mu.RUnlock()
|
||||
|
||||
if k.permanent {
|
||||
return KeyStatePermanent, 0
|
||||
}
|
||||
now := k.clock.Now()
|
||||
if now.Before(k.cooldownUntil) {
|
||||
return KeyStateTemporary, k.cooldownUntil.Sub(now)
|
||||
}
|
||||
return KeyStateValid, 0
|
||||
}
|
||||
|
||||
// MarkTemporary marks the key as temporarily unavailable with
|
||||
// the specified cooldown duration. Returns true if this call
|
||||
// transitions the key to temporary.
|
||||
@@ -146,6 +176,47 @@ 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 {
|
||||
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.
|
||||
case KeyStateValid:
|
||||
return &TransientKeyPoolError{}
|
||||
// Recoverable later: track soonest remaining cooldown.
|
||||
case KeyStateTemporary:
|
||||
if !hasCooldown || cooldown < retryAfter {
|
||||
retryAfter = cooldown
|
||||
hasCooldown = true
|
||||
}
|
||||
// Permanent: keep walking to confirm error type.
|
||||
default:
|
||||
}
|
||||
}
|
||||
if hasCooldown {
|
||||
return &TransientKeyPoolError{RetryAfter: retryAfter}
|
||||
}
|
||||
return ErrPermanentKeyPool
|
||||
}
|
||||
|
||||
// PoolState returns a snapshot of each key's state in the pool's
|
||||
// original order, used by tests and other diagnostic callers. Use
|
||||
// Walker for the failover iteration path.
|
||||
func (p *Pool) PoolState() []KeyState {
|
||||
states := make([]KeyState, len(p.keys))
|
||||
for i := range p.keys {
|
||||
states[i] = p.keys[i].State()
|
||||
}
|
||||
return states
|
||||
}
|
||||
|
||||
// Walker traverses a Pool for a single request. Each request
|
||||
// creates its own walker so that it can independently iterate
|
||||
// through keys without interfering with other requests.
|
||||
@@ -162,14 +233,15 @@ func (p *Pool) Walker() *Walker {
|
||||
return &Walker{pool: p, pos: 0}
|
||||
}
|
||||
|
||||
// Next returns a Key handle for the next available key. This is
|
||||
// a read-only operation; it does not modify the pool state.
|
||||
// Next returns a Key handle for the next available key without
|
||||
// modifying the pool state.
|
||||
//
|
||||
// Returns ErrAllKeysExhausted when no more keys are available.
|
||||
// Returns *TransientKeyPoolError or ErrPermanentKeyPool
|
||||
// when no more keys are available.
|
||||
func (w *Walker) Next() (*Key, error) {
|
||||
pool := w.pool
|
||||
if pool == nil {
|
||||
return nil, ErrAllKeysExhausted
|
||||
return nil, ErrPermanentKeyPool
|
||||
}
|
||||
|
||||
for i := w.pos; i < len(pool.keys); i++ {
|
||||
@@ -183,5 +255,5 @@ func (w *Walker) Next() (*Key, error) {
|
||||
}
|
||||
|
||||
// No keys available.
|
||||
return nil, ErrAllKeysExhausted
|
||||
return nil, pool.keyPoolError()
|
||||
}
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
package keypool_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -49,7 +51,8 @@ func TestNewKeyPool(t *testing.T) {
|
||||
|
||||
// No more keys available.
|
||||
_, err = walker.Next()
|
||||
require.ErrorIs(t, err, keypool.ErrAllKeysExhausted)
|
||||
var transient *keypool.TransientKeyPoolError
|
||||
require.ErrorAs(t, err, &transient, "expected transient exhaustion: walker returned all valid keys, none marked permanent")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -282,19 +285,21 @@ func TestWalkerNext(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
keys []string
|
||||
setup func(t *testing.T, pool *keypool.Pool)
|
||||
advance time.Duration
|
||||
expectValid []string
|
||||
name string
|
||||
keys []string
|
||||
setup func(t *testing.T, pool *keypool.Pool)
|
||||
advance time.Duration
|
||||
expectedValid []string
|
||||
expectedErr error
|
||||
}{
|
||||
{
|
||||
// Given: key-0: valid, key-1: valid, key-2: valid.
|
||||
// Then: key-0: valid, key-1: valid, key-2: valid.
|
||||
name: "all_keys_valid",
|
||||
keys: []string{"key-0", "key-1", "key-2"},
|
||||
setup: func(_ *testing.T, _ *keypool.Pool) {},
|
||||
expectValid: []string{"key-0", "key-1", "key-2"},
|
||||
name: "all_keys_valid",
|
||||
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{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary, key-1: valid, key-2: valid.
|
||||
@@ -306,7 +311,8 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key.MarkTemporary(60 * time.Second)
|
||||
},
|
||||
expectValid: []string{"key-1", "key-2"},
|
||||
expectedValid: []string{"key-1", "key-2"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: permanent, key-1: permanent, key-2: valid.
|
||||
@@ -322,7 +328,8 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key1.MarkPermanent()
|
||||
},
|
||||
expectValid: []string{"key-2"},
|
||||
expectedValid: []string{"key-2"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (30s), key-1: valid.
|
||||
@@ -335,8 +342,9 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key.MarkTemporary(30 * time.Second)
|
||||
},
|
||||
advance: 35 * time.Second,
|
||||
expectValid: []string{"key-0", "key-1"},
|
||||
advance: 35 * time.Second,
|
||||
expectedValid: []string{"key-0", "key-1"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (zero, default 60s), key-1: valid.
|
||||
@@ -349,8 +357,9 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key.MarkTemporary(0)
|
||||
},
|
||||
advance: 50 * time.Second,
|
||||
expectValid: []string{"key-1"},
|
||||
advance: 50 * time.Second,
|
||||
expectedValid: []string{"key-1"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (zero, default 60s), key-1: valid.
|
||||
@@ -363,8 +372,9 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key.MarkTemporary(0)
|
||||
},
|
||||
advance: 65 * time.Second,
|
||||
expectValid: []string{"key-0", "key-1"},
|
||||
advance: 65 * time.Second,
|
||||
expectedValid: []string{"key-0", "key-1"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (negative, default 60s), key-1: valid.
|
||||
@@ -377,13 +387,14 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key.MarkTemporary(-10 * time.Second)
|
||||
},
|
||||
advance: 65 * time.Second,
|
||||
expectValid: []string{"key-0", "key-1"},
|
||||
advance: 65 * time.Second,
|
||||
expectedValid: []string{"key-0", "key-1"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (60s), then marked again with shorter cooldown (10s).
|
||||
// When: 15s pass (past 10s, but not 60s).
|
||||
// Then: key-0: temporary.
|
||||
// Then: key-0: temporary, 45s remaining.
|
||||
name: "shorter_cooldown_preserves_longer_not_expired",
|
||||
keys: []string{"key-0"},
|
||||
setup: func(t *testing.T, pool *keypool.Pool) {
|
||||
@@ -392,8 +403,9 @@ func TestWalkerNext(t *testing.T) {
|
||||
key.MarkTemporary(60 * time.Second)
|
||||
key.MarkTemporary(10 * time.Second)
|
||||
},
|
||||
advance: 15 * time.Second,
|
||||
expectValid: []string{},
|
||||
advance: 15 * time.Second,
|
||||
expectedValid: []string{},
|
||||
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 45 * time.Second},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (60s), then marked again with shorter cooldown (10s).
|
||||
@@ -407,8 +419,30 @@ func TestWalkerNext(t *testing.T) {
|
||||
key.MarkTemporary(60 * time.Second)
|
||||
key.MarkTemporary(10 * time.Second)
|
||||
},
|
||||
advance: 65 * time.Second,
|
||||
expectValid: []string{"key-0"},
|
||||
advance: 65 * time.Second,
|
||||
expectedValid: []string{"key-0"},
|
||||
expectedErr: &keypool.TransientKeyPoolError{},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary (60s), key-1: temporary (10s), key-2: temporary (30s).
|
||||
// Then: key-0: temporary, key-1: temporary, key-2: temporary.
|
||||
// Smallest remaining cooldown is reported on exhaustion.
|
||||
name: "smallest_cooldown_across_temporary_keys",
|
||||
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.MarkTemporary(60 * time.Second)
|
||||
key1, err := walker.Next()
|
||||
require.NoError(t, err)
|
||||
key1.MarkTemporary(10 * time.Second)
|
||||
key2, err := walker.Next()
|
||||
require.NoError(t, err)
|
||||
key2.MarkTemporary(30 * time.Second)
|
||||
},
|
||||
expectedValid: []string{},
|
||||
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 10 * time.Second},
|
||||
},
|
||||
{
|
||||
// Given: key-0: temporary, key-1: temporary.
|
||||
@@ -424,7 +458,8 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key1.MarkTemporary(60 * time.Second)
|
||||
},
|
||||
expectValid: []string{},
|
||||
expectedValid: []string{},
|
||||
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 60 * time.Second},
|
||||
},
|
||||
{
|
||||
// Given: key-0: permanent, key-1: permanent.
|
||||
@@ -440,7 +475,8 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key1.MarkPermanent()
|
||||
},
|
||||
expectValid: []string{},
|
||||
expectedValid: []string{},
|
||||
expectedErr: keypool.ErrPermanentKeyPool,
|
||||
},
|
||||
{
|
||||
// Given: key-0: permanent, key-1: temporary, key-2: permanent.
|
||||
@@ -459,7 +495,8 @@ func TestWalkerNext(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
key2.MarkPermanent()
|
||||
},
|
||||
expectValid: []string{},
|
||||
expectedValid: []string{},
|
||||
expectedErr: &keypool.TransientKeyPoolError{RetryAfter: 60 * time.Second},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -478,7 +515,7 @@ func TestWalkerNext(t *testing.T) {
|
||||
}
|
||||
|
||||
walker := pool.Walker()
|
||||
for _, expectedKey := range tc.expectValid {
|
||||
for _, expectedKey := range tc.expectedValid {
|
||||
key, err := walker.Next()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, expectedKey, key.Value())
|
||||
@@ -486,7 +523,93 @@ func TestWalkerNext(t *testing.T) {
|
||||
|
||||
// After all expected keys, the walker should be exhausted.
|
||||
_, err = walker.Next()
|
||||
require.ErrorIs(t, err, keypool.ErrAllKeysExhausted)
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestKeyConcurrent exercises the documented concurrent-safety
|
||||
// contract by hammering a single key with concurrent Mark calls
|
||||
// and asserting the resulting state honors the pool's invariants.
|
||||
func TestKeyConcurrent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
// run is called concurrently from numGoroutines, each
|
||||
// with its own index.
|
||||
run func(idx int, key *keypool.Key)
|
||||
// verify asserts the final state. May advance the clock.
|
||||
verify func(t *testing.T, key *keypool.Key, clk *quartz.Mock)
|
||||
}{
|
||||
{
|
||||
// Half of the goroutines mark the key as temporary
|
||||
// with 60s, the other half with 10s. The longer
|
||||
// cooldown must win regardless of ordering.
|
||||
name: "longer_cooldown_wins",
|
||||
run: func(idx int, key *keypool.Key) {
|
||||
if idx%2 == 0 {
|
||||
key.MarkTemporary(60 * time.Second)
|
||||
} else {
|
||||
key.MarkTemporary(10 * time.Second)
|
||||
}
|
||||
},
|
||||
verify: func(t *testing.T, key *keypool.Key, clk *quartz.Mock) {
|
||||
// At 50s the 60s cooldown is still active.
|
||||
clk.Advance(50 * time.Second)
|
||||
assert.Equal(t, keypool.KeyStateTemporary, key.State())
|
||||
// At 65s the 60s cooldown has expired.
|
||||
clk.Advance(15 * time.Second)
|
||||
assert.Equal(t, keypool.KeyStateValid, key.State())
|
||||
},
|
||||
},
|
||||
{
|
||||
// Half of the goroutines mark the key as permanent,
|
||||
// the other half mark it as temporary. Permanent is
|
||||
// terminal: any permanent call wins.
|
||||
name: "permanent_wins_over_temporary",
|
||||
run: func(idx int, key *keypool.Key) {
|
||||
if idx%2 == 0 {
|
||||
key.MarkPermanent()
|
||||
} else {
|
||||
key.MarkTemporary(60 * time.Second)
|
||||
}
|
||||
},
|
||||
verify: func(t *testing.T, key *keypool.Key, _ *quartz.Mock) {
|
||||
assert.Equal(t, keypool.KeyStatePermanent, key.State())
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
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)
|
||||
|
||||
const numGoroutines = 10
|
||||
var wg sync.WaitGroup
|
||||
for r := range numGoroutines {
|
||||
wg.Add(1)
|
||||
go func(r int) {
|
||||
defer wg.Done()
|
||||
tc.run(r, key)
|
||||
}(r)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
tc.verify(t, key, clk)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user