mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add key pool failover metrics to aibridge (#25901)
## Description This PR adds Prometheus metrics for aibridge's API-key failover, giving visibility into key pool health and failover behavior per provider. The following metrics are introduced: - **`key_pool_state`** (gauge): number of keys currently in each state (`valid`, `temporary`, `permanent`) per provider, sampled at scrape time. - **`key_pool_state_transitions_total`** (counter): key state transitions during failover, labeled by `reason` (`rate_limited`, `unauthorized`, `forbidden`). - **`key_pool_exhaustions_total`** (counter): times a pool ran out of usable keys, labeled by `outcome` (`rate_limited`, `auth_failed`). - **`key_pool_failover_attempts`** (histogram): keys attempted before success or exhaustion (per interception for bridged requests, per request for passthrough). ## Changes - Moves `MarkKeyOnStatus` and key-pool error handling onto `*keypool.Pool`. - Attaches metrics to each provider's key pool at install time, on construction and on provider reload. - Adds a scrape-time state collector and a `KeyPools()` accessor on the bridge pool to feed it. - Tracks per-request key attempts in the bridged and passthrough failover paths. - Adds test coverage for the new metrics across the keypool unit tests, the bridged intercept failover tests, and the passthrough failover test. Closes https://github.com/coder/internal/issues/1447 Closes https://linear.app/codercom/issue/AIGOV-198/aibridge-key-failover-observability > [!NOTE] > Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira
This commit is contained in:
@@ -226,9 +226,8 @@ func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key,
|
||||
if !errors.As(err, &apiErr) {
|
||||
return false
|
||||
}
|
||||
return keypool.MarkKeyOnStatus(
|
||||
ctx, key, apiErr.Response,
|
||||
i.logger, i.providerName,
|
||||
return i.cfg.KeyPool.MarkKeyOnStatus(
|
||||
ctx, key, apiErr.Response, i.logger,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -136,7 +136,7 @@ func TestMarkKeyOnError(t *testing.T) {
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
pool, err := keypool.New([]string{"key-0"}, quartz.NewMock(t))
|
||||
pool, err := keypool.New(config.ProviderOpenAI, []string{"key-0"}, quartz.NewMock(t), nil)
|
||||
require.NoError(t, err)
|
||||
key, keyPoolErr := pool.Walker().Next()
|
||||
require.Nil(t, keyPoolErr)
|
||||
|
||||
@@ -88,6 +88,11 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
|
||||
logger.Warn(ctx, "failed to retrieve last user prompt", slog.Error(err))
|
||||
}
|
||||
|
||||
// Sum the key attempts across all iterations and record once when the
|
||||
// interception completes.
|
||||
var totalKeyAttempts int
|
||||
defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }()
|
||||
|
||||
for {
|
||||
// TODO add outer loop span (https://github.com/coder/aibridge/issues/67)
|
||||
|
||||
@@ -100,7 +105,9 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
|
||||
opts = append(opts, intercept.ActorHeadersAsOpenAIOpts(actor)...)
|
||||
}
|
||||
|
||||
completion, err = i.newChatCompletion(ctx, svc, opts)
|
||||
var keyAttempts int
|
||||
completion, keyAttempts, err = i.newChatCompletion(ctx, svc, opts)
|
||||
totalKeyAttempts += keyAttempts
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
@@ -267,12 +274,14 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
|
||||
return nil
|
||||
}
|
||||
|
||||
// newChatCompletion routes between BYOK (single attempt) and
|
||||
// centralized failover.
|
||||
func (i *BlockingInterception) newChatCompletion(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, error) {
|
||||
// newChatCompletion routes between BYOK (single attempt) and centralized
|
||||
// failover, returning the upstream completion, the number of key attempts
|
||||
// made for this call, and any error.
|
||||
func (i *BlockingInterception) newChatCompletion(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, int, error) {
|
||||
// BYOK: single attempt, no failover.
|
||||
if i.cfg.KeyPool == nil {
|
||||
return i.newChatCompletionWithKey(ctx, svc, opts)
|
||||
completion, err := i.newChatCompletionWithKey(ctx, svc, opts)
|
||||
return completion, 0, err
|
||||
}
|
||||
return i.newChatCompletionWithKeyFailover(ctx, svc, opts)
|
||||
}
|
||||
@@ -285,17 +294,17 @@ func (i *BlockingInterception) newChatCompletionWithKey(ctx context.Context, svc
|
||||
return svc.New(ctx, i.req.ChatCompletionNewParams, opts...)
|
||||
}
|
||||
|
||||
// newChatCompletionWithKeyFailover walks the centralized key
|
||||
// pool, trying each key until one succeeds or the pool is
|
||||
// exhausted. Keys are marked temporary on 429 and permanent on
|
||||
// 401/403. Errors that aren't key-specific don't trigger
|
||||
// failover and are returned to the caller.
|
||||
func (i *BlockingInterception) newChatCompletionWithKeyFailover(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, error) {
|
||||
// newChatCompletionWithKeyFailover walks the centralized key pool, trying each
|
||||
// key until one succeeds or the pool is exhausted. Keys are marked temporary
|
||||
// on 429 and permanent on 401/403. Errors that aren't key-specific don't
|
||||
// trigger failover and are returned to the caller. It returns the upstream
|
||||
// completion, the number of key attempts made for this call, and any error.
|
||||
func (i *BlockingInterception) newChatCompletionWithKeyFailover(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, int, error) {
|
||||
walker := i.cfg.KeyPool.Walker()
|
||||
for {
|
||||
key, keyPoolErr := walker.Next()
|
||||
if keyPoolErr != nil {
|
||||
return nil, keyPoolErr
|
||||
return nil, walker.Attempts(), keyPoolErr
|
||||
}
|
||||
// Record the key in use so the hint reflects the last attempted key.
|
||||
i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value())
|
||||
@@ -316,6 +325,6 @@ func (i *BlockingInterception) newChatCompletionWithKeyFailover(ctx context.Cont
|
||||
}
|
||||
// Either success (completion, nil) or a non-key error
|
||||
// (nil, err): nothing to retry, return as-is.
|
||||
return completion, err
|
||||
return completion, walker.Attempts(), err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -128,6 +128,11 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re
|
||||
interceptionErr error
|
||||
)
|
||||
|
||||
// Sum the key attempts across all iterations and record once when the
|
||||
// interception completes.
|
||||
var totalKeyAttempts int
|
||||
defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }()
|
||||
|
||||
for {
|
||||
// TODO add outer loop span (https://github.com/coder/aibridge/issues/67)
|
||||
|
||||
@@ -177,6 +182,8 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re
|
||||
)
|
||||
}
|
||||
|
||||
totalKeyAttempts += walker.Attempts()
|
||||
|
||||
// TODO(ssncferreira): inject actor headers directly in the client-header
|
||||
// middleware instead of using SDK options.
|
||||
if actor := aibcontext.ActorFromContext(r.Context()); actor != nil && i.cfg.SendActorHeaders {
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/sjson"
|
||||
@@ -21,7 +22,10 @@ import (
|
||||
"github.com/coder/coder/v2/aibridge/intercept/responses"
|
||||
"github.com/coder/coder/v2/aibridge/internal/testutil"
|
||||
"github.com/coder/coder/v2/aibridge/keypool"
|
||||
"github.com/coder/coder/v2/aibridge/metrics"
|
||||
"github.com/coder/coder/v2/aibridge/utils"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest/promhelp"
|
||||
codertestutil "github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
@@ -31,6 +35,9 @@ import (
|
||||
type interceptorCase struct {
|
||||
// name labels the subtest.
|
||||
name string
|
||||
// provider is the provider name used to build the key pool and to label its
|
||||
// failover metrics.
|
||||
provider string
|
||||
// path is the route the interceptor handles.
|
||||
path string
|
||||
// authHeader is the header the upstream key is carried in. It is also used
|
||||
@@ -69,6 +76,7 @@ func keyFromHeader(name string, h http.Header) string {
|
||||
var interceptorCases = []interceptorCase{
|
||||
{
|
||||
name: "messages",
|
||||
provider: config.ProviderAnthropic,
|
||||
path: "/v1/messages",
|
||||
authHeader: "X-Api-Key",
|
||||
fixture: func(_, agentic bool) []byte {
|
||||
@@ -101,6 +109,7 @@ var interceptorCases = []interceptorCase{
|
||||
},
|
||||
{
|
||||
name: "chatcompletions",
|
||||
provider: config.ProviderOpenAI,
|
||||
path: "/v1/chat/completions",
|
||||
authHeader: "Authorization",
|
||||
fixture: func(_, agentic bool) []byte {
|
||||
@@ -133,6 +142,7 @@ var interceptorCases = []interceptorCase{
|
||||
},
|
||||
{
|
||||
name: "responses",
|
||||
provider: config.ProviderOpenAI,
|
||||
path: "/v1/responses",
|
||||
authHeader: "Authorization",
|
||||
fixture: func(streaming, agentic bool) []byte {
|
||||
@@ -196,6 +206,10 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
expectedKeyStates []keypool.KeyState
|
||||
expectedSeenKeys []string
|
||||
expectedBodyContains string
|
||||
// Expected key_pool_state_transitions_total counts by reason.
|
||||
expectedTransitions map[string]int
|
||||
// Expected key_pool_exhaustions_total counts by outcome.
|
||||
expectedExhaustions map[string]int
|
||||
}{
|
||||
{
|
||||
// One valid key succeeds on the first attempt.
|
||||
@@ -213,9 +227,10 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
responses: func(s testutil.UpstreamResponse) []testutil.UpstreamResponse {
|
||||
return []testutil.UpstreamResponse{errResp(http.StatusTooManyRequests, "5"), s}
|
||||
},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedTransitions: map[string]int{"rate_limited": 1},
|
||||
},
|
||||
{
|
||||
// A 401 marks the key permanent and fails over to the next one.
|
||||
@@ -224,9 +239,10 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
responses: func(s testutil.UpstreamResponse) []testutil.UpstreamResponse {
|
||||
return []testutil.UpstreamResponse{errResp(http.StatusUnauthorized, ""), s}
|
||||
},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedTransitions: map[string]int{"unauthorized": 1},
|
||||
},
|
||||
{
|
||||
// A 403 marks the key permanent and fails over to the next one.
|
||||
@@ -235,9 +251,10 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
responses: func(s testutil.UpstreamResponse) []testutil.UpstreamResponse {
|
||||
return []testutil.UpstreamResponse{errResp(http.StatusForbidden, ""), s}
|
||||
},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedTransitions: map[string]int{"forbidden": 1},
|
||||
},
|
||||
{
|
||||
// Every key is rate-limited, so the pool is exhausted and the
|
||||
@@ -259,7 +276,9 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
keypool.KeyStateTemporary,
|
||||
keypool.KeyStateTemporary,
|
||||
},
|
||||
expectedSeenKeys: []string{k0, k1, k2},
|
||||
expectedSeenKeys: []string{k0, k1, k2},
|
||||
expectedTransitions: map[string]int{"rate_limited": 3},
|
||||
expectedExhaustions: map[string]int{"rate_limited": 1},
|
||||
},
|
||||
{
|
||||
// Every key is unauthorized, so the pool is permanently exhausted.
|
||||
@@ -271,9 +290,11 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
errResp(http.StatusUnauthorized, ""),
|
||||
}
|
||||
},
|
||||
expectedStatus: http.StatusBadGateway,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStatePermanent},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedStatus: http.StatusBadGateway,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStatePermanent},
|
||||
expectedSeenKeys: []string{k0, k1},
|
||||
expectedTransitions: map[string]int{"unauthorized": 2},
|
||||
expectedExhaustions: map[string]int{"auth_failed": 1},
|
||||
},
|
||||
{
|
||||
// A 500 is not a key-specific failure, so it does not fail over.
|
||||
@@ -306,10 +327,12 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
t.Run(ic.name+"/"+mode+"/"+tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
reg := prometheus.NewRegistry()
|
||||
m := metrics.NewMetrics(reg)
|
||||
var pool *keypool.Pool
|
||||
if len(tc.keys) > 0 {
|
||||
var err error
|
||||
pool, err = keypool.New(tc.keys, quartz.NewMock(t))
|
||||
pool, err = keypool.New(ic.provider, tc.keys, quartz.NewMock(t), m)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
@@ -353,6 +376,38 @@ func TestInterception_KeyFailover(t *testing.T) {
|
||||
if tc.expectedBodyContains != "" {
|
||||
assert.Contains(t, w.Body.String(), tc.expectedBodyContains, "response body")
|
||||
}
|
||||
|
||||
// A centralized interception records one failover-attempts
|
||||
// observation, labeled with the provider, summing the keys
|
||||
// tried (one per upstream attempt). BYOK has no pool, so none.
|
||||
if pool != nil {
|
||||
hist := promhelp.HistogramValue(t, reg, "key_pool_failover_attempts",
|
||||
prometheus.Labels{"provider": ic.provider})
|
||||
assert.Equal(t, uint64(1), hist.GetSampleCount())
|
||||
assert.Equal(t, float64(len(tc.expectedSeenKeys)), hist.GetSampleSum())
|
||||
} else {
|
||||
assert.Nil(t, promhelp.MetricValue(t, reg, "key_pool_failover_attempts",
|
||||
prometheus.Labels{"provider": ic.provider}))
|
||||
}
|
||||
|
||||
gathered, err := reg.Gather()
|
||||
require.NoError(t, err)
|
||||
// One transition per marked key, by reason.
|
||||
for _, reason := range []string{"rate_limited", "unauthorized", "forbidden"} {
|
||||
if want := tc.expectedTransitions[reason]; want > 0 {
|
||||
assert.True(t, codertestutil.PromCounterHasValue(t, gathered, float64(want), "key_pool_state_transitions_total", ic.provider, reason))
|
||||
} else {
|
||||
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_state_transitions_total", ic.provider, reason))
|
||||
}
|
||||
}
|
||||
// Exhaustion outcome when no usable key remains.
|
||||
for _, outcome := range []string{"rate_limited", "auth_failed"} {
|
||||
if want := tc.expectedExhaustions[outcome]; want > 0 {
|
||||
assert.True(t, codertestutil.PromCounterHasValue(t, gathered, float64(want), "key_pool_exhaustions_total", outcome, ic.provider))
|
||||
} else {
|
||||
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_exhaustions_total", outcome, ic.provider))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -380,6 +435,10 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
|
||||
expectedKeyStates []keypool.KeyState
|
||||
expectedSeenKeys []string
|
||||
expectedBodyContains string
|
||||
// Expected key_pool_state_transitions_total counts by reason.
|
||||
expectedTransitions map[string]int
|
||||
// Expected key_pool_exhaustions_total counts by outcome.
|
||||
expectedExhaustions map[string]int
|
||||
// expectErr is true when ProcessRequest returns an error because the
|
||||
// pool is exhausted.
|
||||
expectErr bool
|
||||
@@ -403,9 +462,10 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
|
||||
responses: func(toolCall, final testutil.UpstreamResponse) []testutil.UpstreamResponse {
|
||||
return []testutil.UpstreamResponse{toolCall, errResp(http.StatusTooManyRequests, "5"), final}
|
||||
},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k0, k1},
|
||||
expectedStatus: http.StatusOK,
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid},
|
||||
expectedSeenKeys: []string{k0, k0, k1},
|
||||
expectedTransitions: map[string]int{"rate_limited": 1},
|
||||
},
|
||||
{
|
||||
// The continuation is rate-limited on every key, exhausting the pool.
|
||||
@@ -423,6 +483,8 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
|
||||
expectedBodyContains: "all configured keys are rate-limited",
|
||||
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateTemporary},
|
||||
expectedSeenKeys: []string{k0, k0, k1},
|
||||
expectedTransitions: map[string]int{"rate_limited": 2},
|
||||
expectedExhaustions: map[string]int{"rate_limited": 1},
|
||||
expectErr: true,
|
||||
},
|
||||
}
|
||||
@@ -434,7 +496,9 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
|
||||
t.Run(ic.name+"/"+mode+"/"+tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
pool, err := keypool.New(tc.keys, quartz.NewMock(t))
|
||||
reg := prometheus.NewRegistry()
|
||||
m := metrics.NewMetrics(reg)
|
||||
pool, err := keypool.New(ic.provider, tc.keys, quartz.NewMock(t), m)
|
||||
require.NoError(t, err)
|
||||
|
||||
fixture := fixtures.Parse(t, ic.fixture(streaming, true))
|
||||
@@ -487,6 +551,32 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
|
||||
if tc.expectedBodyContains != "" {
|
||||
assert.Contains(t, w.Body.String(), tc.expectedBodyContains, "response body")
|
||||
}
|
||||
|
||||
// One observation per interception, summing keys tried across
|
||||
// all agentic-loop iterations (one per upstream attempt).
|
||||
hist := promhelp.HistogramValue(t, reg, "key_pool_failover_attempts",
|
||||
prometheus.Labels{"provider": ic.provider})
|
||||
assert.Equal(t, uint64(1), hist.GetSampleCount())
|
||||
assert.Equal(t, float64(len(tc.expectedSeenKeys)), hist.GetSampleSum())
|
||||
|
||||
gathered, err := reg.Gather()
|
||||
require.NoError(t, err)
|
||||
// One transition per marked key, by reason.
|
||||
for _, reason := range []string{"rate_limited", "unauthorized", "forbidden"} {
|
||||
if want := tc.expectedTransitions[reason]; want > 0 {
|
||||
assert.True(t, codertestutil.PromCounterHasValue(t, gathered, float64(want), "key_pool_state_transitions_total", ic.provider, reason))
|
||||
} else {
|
||||
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_state_transitions_total", ic.provider, reason))
|
||||
}
|
||||
}
|
||||
// Exhaustion outcome when no usable key remains.
|
||||
for _, outcome := range []string{"rate_limited", "auth_failed"} {
|
||||
if want := tc.expectedExhaustions[outcome]; want > 0 {
|
||||
assert.True(t, codertestutil.PromCounterHasValue(t, gathered, float64(want), "key_pool_exhaustions_total", outcome, ic.provider))
|
||||
} else {
|
||||
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_exhaustions_total", outcome, ic.provider))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -583,9 +583,8 @@ func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key,
|
||||
if !errors.As(err, &apiErr) {
|
||||
return false
|
||||
}
|
||||
return keypool.MarkKeyOnStatus(
|
||||
ctx, key, apiErr.Response,
|
||||
i.logger, i.providerName,
|
||||
return i.cfg.KeyPool.MarkKeyOnStatus(
|
||||
ctx, key, apiErr.Response, i.logger,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -1122,7 +1122,7 @@ func TestMarkKeyOnError(t *testing.T) {
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
pool, err := keypool.New([]string{"key-0"}, quartz.NewMock(t))
|
||||
pool, err := keypool.New(config.ProviderAnthropic, []string{"key-0"}, quartz.NewMock(t), nil)
|
||||
require.NoError(t, err)
|
||||
key, keyPoolErr := pool.Walker().Next()
|
||||
require.Nil(t, keyPoolErr)
|
||||
|
||||
@@ -105,9 +105,16 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
|
||||
// Accumulate usage across the entire streaming interaction (including tool reinvocations).
|
||||
var cumulativeUsage anthropic.Usage
|
||||
|
||||
// Sum the key attempts across all iterations and record once when the
|
||||
// interception completes.
|
||||
var totalKeyAttempts int
|
||||
defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }()
|
||||
|
||||
for {
|
||||
// TODO add outer loop span (https://github.com/coder/aibridge/issues/67)
|
||||
resp, err = i.newMessage(ctx, svc)
|
||||
var keyAttempts int
|
||||
resp, keyAttempts, err = i.newMessage(ctx, svc)
|
||||
totalKeyAttempts += keyAttempts
|
||||
if err != nil {
|
||||
if eventstream.IsConnError(err) {
|
||||
// Can't write a response, just error out.
|
||||
@@ -343,11 +350,13 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
|
||||
}
|
||||
|
||||
// newMessage routes between BYOK (single attempt) and centralized
|
||||
// failover.
|
||||
func (i *BlockingInterception) newMessage(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, error) {
|
||||
// failover, returning the upstream message, the number of key attempts
|
||||
// made for this call, and any error.
|
||||
func (i *BlockingInterception) newMessage(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, int, error) {
|
||||
// BYOK: single attempt, no failover.
|
||||
if i.cfg.KeyPool == nil {
|
||||
return i.newMessageWithKey(ctx, svc)
|
||||
msg, err := i.newMessageWithKey(ctx, svc)
|
||||
return msg, 0, err
|
||||
}
|
||||
return i.newMessageWithKeyFailover(ctx, svc)
|
||||
}
|
||||
@@ -361,17 +370,17 @@ func (i *BlockingInterception) newMessageWithKey(ctx context.Context, svc anthro
|
||||
return svc.New(ctx, anthropic.MessageNewParams{}, opts...)
|
||||
}
|
||||
|
||||
// newMessageWithKeyFailover walks the centralized key pool,
|
||||
// trying each key until one succeeds or the pool is exhausted.
|
||||
// Keys are marked temporary on 429 and permanent on 401/403.
|
||||
// Errors that aren't key-specific don't trigger failover and
|
||||
// are returned to the caller.
|
||||
func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, error) {
|
||||
// newMessageWithKeyFailover walks the centralized key pool, trying each key
|
||||
// until one succeeds or the pool is exhausted. Keys are marked temporary on
|
||||
// 429 and permanent on 401/403. Errors that aren't key-specific don't trigger
|
||||
// failover and are returned to the caller. It returns the upstream message,
|
||||
// the number of key attempts made for this call, and any error.
|
||||
func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, int, error) {
|
||||
walker := i.cfg.KeyPool.Walker()
|
||||
for {
|
||||
key, keyPoolErr := walker.Next()
|
||||
if keyPoolErr != nil {
|
||||
return nil, keyPoolErr
|
||||
return nil, walker.Attempts(), keyPoolErr
|
||||
}
|
||||
// Record the key in use so the hint reflects the last attempted key.
|
||||
i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value())
|
||||
@@ -390,6 +399,6 @@ func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, sv
|
||||
}
|
||||
// Either success (msg, nil) or a non-key error (nil, err):
|
||||
// nothing to retry, return as-is.
|
||||
return msg, err
|
||||
return msg, walker.Attempts(), err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -153,6 +153,11 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re
|
||||
var lastErr error
|
||||
var interceptionErr error
|
||||
|
||||
// Sum the key attempts across all iterations and record once when the
|
||||
// interception completes.
|
||||
var totalKeyAttempts int
|
||||
defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }()
|
||||
|
||||
isFirst := true
|
||||
newStream:
|
||||
for {
|
||||
@@ -208,6 +213,8 @@ newStream:
|
||||
)
|
||||
}
|
||||
|
||||
totalKeyAttempts += walker.Attempts()
|
||||
|
||||
stream := i.newStream(streamCtx, svc, streamOpts...)
|
||||
|
||||
var message anthropic.Message
|
||||
|
||||
@@ -181,9 +181,8 @@ func (i *responsesInterceptionBase) markKeyOnError(ctx context.Context, key *key
|
||||
if !errors.As(err, &apiErr) {
|
||||
return false
|
||||
}
|
||||
return keypool.MarkKeyOnStatus(
|
||||
ctx, key, apiErr.Response,
|
||||
i.logger, i.providerName,
|
||||
return i.cfg.KeyPool.MarkKeyOnStatus(
|
||||
ctx, key, apiErr.Response, i.logger,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -440,7 +440,7 @@ func TestMarkKeyOnError(t *testing.T) {
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
pool, err := keypool.New([]string{"key-0"}, quartz.NewMock(t))
|
||||
pool, err := keypool.New(config.ProviderOpenAI, []string{"key-0"}, quartz.NewMock(t), nil)
|
||||
require.NoError(t, err)
|
||||
key, keyPoolErr := pool.Walker().Next()
|
||||
require.Nil(t, keyPoolErr)
|
||||
|
||||
@@ -86,6 +86,11 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r *
|
||||
}
|
||||
shouldLoop := true
|
||||
|
||||
// Sum the key attempts across all iterations and record once when the
|
||||
// interception completes.
|
||||
var totalKeyAttempts int
|
||||
defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }()
|
||||
|
||||
for shouldLoop {
|
||||
srv := i.newResponsesService()
|
||||
respCopy = responseCopier{}
|
||||
@@ -99,7 +104,9 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r *
|
||||
opts = append(opts, intercept.ActorHeadersAsOpenAIOpts(actor)...)
|
||||
}
|
||||
|
||||
response, upstreamErr = i.newResponse(ctx, srv, opts)
|
||||
var keyAttempts int
|
||||
response, keyAttempts, upstreamErr = i.newResponse(ctx, srv, opts)
|
||||
totalKeyAttempts += keyAttempts
|
||||
|
||||
// The failover loop may return a keypool exhaustion
|
||||
// error. Render it here.
|
||||
@@ -146,12 +153,14 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r *
|
||||
return errors.Join(upstreamErr, err)
|
||||
}
|
||||
|
||||
// newResponse routes between BYOK (single attempt) and
|
||||
// centralized failover.
|
||||
func (i *BlockingResponsesInterceptor) newResponse(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, error) {
|
||||
// newResponse routes between BYOK (single attempt) and centralized failover,
|
||||
// returning the upstream response, the number of key attempts made for this
|
||||
// call, and any error.
|
||||
func (i *BlockingResponsesInterceptor) newResponse(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, int, error) {
|
||||
// BYOK: single attempt, no failover.
|
||||
if i.cfg.KeyPool == nil {
|
||||
return i.newResponseWithKey(ctx, srv, opts)
|
||||
response, err := i.newResponseWithKey(ctx, srv, opts)
|
||||
return response, 0, err
|
||||
}
|
||||
return i.newResponseWithKeyFailover(ctx, srv, opts)
|
||||
}
|
||||
@@ -165,17 +174,17 @@ func (i *BlockingResponsesInterceptor) newResponseWithKey(ctx context.Context, s
|
||||
return srv.New(ctx, responses.ResponseNewParams{}, opts...)
|
||||
}
|
||||
|
||||
// newResponseWithKeyFailover walks the centralized key pool,
|
||||
// trying each key until one succeeds or the pool is exhausted.
|
||||
// Keys are marked temporary on 429 and permanent on 401/403.
|
||||
// Errors that aren't key-specific don't trigger failover and
|
||||
// are returned to the caller.
|
||||
func (i *BlockingResponsesInterceptor) newResponseWithKeyFailover(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, error) {
|
||||
// newResponseWithKeyFailover walks the centralized key pool, trying each key
|
||||
// until one succeeds or the pool is exhausted. Keys are marked temporary on
|
||||
// 429 and permanent on 401/403. Errors that aren't key-specific don't trigger
|
||||
// failover and are returned to the caller. It returns the upstream response,
|
||||
// the number of key attempts made for this call, and any error.
|
||||
func (i *BlockingResponsesInterceptor) newResponseWithKeyFailover(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, int, error) {
|
||||
walker := i.cfg.KeyPool.Walker()
|
||||
for {
|
||||
key, keyPoolErr := walker.Next()
|
||||
if keyPoolErr != nil {
|
||||
return nil, keyPoolErr
|
||||
return nil, walker.Attempts(), keyPoolErr
|
||||
}
|
||||
// Record the key in use so the hint reflects the last attempted key.
|
||||
i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value())
|
||||
@@ -196,6 +205,6 @@ func (i *BlockingResponsesInterceptor) newResponseWithKeyFailover(ctx context.Co
|
||||
}
|
||||
// Either success (response, nil) or a non-key error
|
||||
// (nil, err): nothing to retry, return as-is.
|
||||
return response, err
|
||||
return response, walker.Attempts(), err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -106,6 +106,11 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r
|
||||
shouldLoop := true
|
||||
srv := i.newResponsesService()
|
||||
|
||||
// Sum the key attempts across all iterations and record once when the
|
||||
// interception completes.
|
||||
var totalKeyAttempts int
|
||||
defer func() { i.cfg.KeyPool.RecordAttempts(totalKeyAttempts) }()
|
||||
|
||||
for shouldLoop {
|
||||
shouldLoop = false
|
||||
|
||||
@@ -140,6 +145,7 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r
|
||||
// agentic mode the inner loop buffers events
|
||||
// instead of streaming them downstream, so the
|
||||
// SSE connection has not been opened yet.
|
||||
totalKeyAttempts += walker.Attempts()
|
||||
i.writeUpstreamError(w, intercept.ResponseErrorFromKeyPool(keyPoolErr))
|
||||
return xerrors.Errorf("key pool exhausted: %w", keyPoolErr)
|
||||
}
|
||||
@@ -175,6 +181,8 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r
|
||||
break
|
||||
}
|
||||
|
||||
totalKeyAttempts += walker.Attempts()
|
||||
|
||||
// func scope to defer steam.Close()
|
||||
err := func() error {
|
||||
defer stream.Close()
|
||||
|
||||
Reference in New Issue
Block a user