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:
Susana Ferreira
2026-06-09 10:49:47 +01:00
committed by GitHub
parent f8c736f859
commit 01ec5e4577
39 changed files with 774 additions and 149 deletions
+2 -3
View File
@@ -226,9 +226,8 @@ func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key,
if !errors.As(err, &apiErr) { if !errors.As(err, &apiErr) {
return false return false
} }
return keypool.MarkKeyOnStatus( return i.cfg.KeyPool.MarkKeyOnStatus(
ctx, key, apiErr.Response, ctx, key, apiErr.Response, i.logger,
i.logger, i.providerName,
) )
} }
@@ -136,7 +136,7 @@ func TestMarkKeyOnError(t *testing.T) {
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() 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) require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next() key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr) require.Nil(t, keyPoolErr)
+22 -13
View File
@@ -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)) 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 { for {
// TODO add outer loop span (https://github.com/coder/aibridge/issues/67) // 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)...) 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 { if err != nil {
break break
} }
@@ -267,12 +274,14 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
return nil return nil
} }
// newChatCompletion routes between BYOK (single attempt) and // newChatCompletion routes between BYOK (single attempt) and centralized
// centralized failover. // failover, returning the upstream completion, the number of key attempts
func (i *BlockingInterception) newChatCompletion(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, error) { // 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. // BYOK: single attempt, no failover.
if i.cfg.KeyPool == nil { 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) 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...) return svc.New(ctx, i.req.ChatCompletionNewParams, opts...)
} }
// newChatCompletionWithKeyFailover walks the centralized key // newChatCompletionWithKeyFailover walks the centralized key pool, trying each
// pool, trying each key until one succeeds or the pool is // key until one succeeds or the pool is exhausted. Keys are marked temporary
// exhausted. Keys are marked temporary on 429 and permanent on // on 429 and permanent on 401/403. Errors that aren't key-specific don't
// 401/403. Errors that aren't key-specific don't trigger // trigger failover and are returned to the caller. It returns the upstream
// failover and are returned to the caller. // 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, error) { func (i *BlockingInterception) newChatCompletionWithKeyFailover(ctx context.Context, svc openai.ChatCompletionService, opts []option.RequestOption) (*openai.ChatCompletion, int, error) {
walker := i.cfg.KeyPool.Walker() walker := i.cfg.KeyPool.Walker()
for { for {
key, keyPoolErr := walker.Next() key, keyPoolErr := walker.Next()
if keyPoolErr != nil { if keyPoolErr != nil {
return nil, keyPoolErr return nil, walker.Attempts(), keyPoolErr
} }
// Record the key in use so the hint reflects the last attempted key. // Record the key in use so the hint reflects the last attempted key.
i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) 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 // Either success (completion, nil) or a non-key error
// (nil, err): nothing to retry, return as-is. // (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 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 { for {
// TODO add outer loop span (https://github.com/coder/aibridge/issues/67) // 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 // TODO(ssncferreira): inject actor headers directly in the client-header
// middleware instead of using SDK options. // middleware instead of using SDK options.
if actor := aibcontext.ActorFromContext(r.Context()); actor != nil && i.cfg.SendActorHeaders { if actor := aibcontext.ActorFromContext(r.Context()); actor != nil && i.cfg.SendActorHeaders {
+108 -18
View File
@@ -7,6 +7,7 @@ import (
"testing" "testing"
"github.com/google/uuid" "github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/tidwall/sjson" "github.com/tidwall/sjson"
@@ -21,7 +22,10 @@ import (
"github.com/coder/coder/v2/aibridge/intercept/responses" "github.com/coder/coder/v2/aibridge/intercept/responses"
"github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/internal/testutil"
"github.com/coder/coder/v2/aibridge/keypool" "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/aibridge/utils"
"github.com/coder/coder/v2/coderd/coderdtest/promhelp"
codertestutil "github.com/coder/coder/v2/testutil"
"github.com/coder/quartz" "github.com/coder/quartz"
) )
@@ -31,6 +35,9 @@ import (
type interceptorCase struct { type interceptorCase struct {
// name labels the subtest. // name labels the subtest.
name string 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 is the route the interceptor handles.
path string path string
// authHeader is the header the upstream key is carried in. It is also used // 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{ var interceptorCases = []interceptorCase{
{ {
name: "messages", name: "messages",
provider: config.ProviderAnthropic,
path: "/v1/messages", path: "/v1/messages",
authHeader: "X-Api-Key", authHeader: "X-Api-Key",
fixture: func(_, agentic bool) []byte { fixture: func(_, agentic bool) []byte {
@@ -101,6 +109,7 @@ var interceptorCases = []interceptorCase{
}, },
{ {
name: "chatcompletions", name: "chatcompletions",
provider: config.ProviderOpenAI,
path: "/v1/chat/completions", path: "/v1/chat/completions",
authHeader: "Authorization", authHeader: "Authorization",
fixture: func(_, agentic bool) []byte { fixture: func(_, agentic bool) []byte {
@@ -133,6 +142,7 @@ var interceptorCases = []interceptorCase{
}, },
{ {
name: "responses", name: "responses",
provider: config.ProviderOpenAI,
path: "/v1/responses", path: "/v1/responses",
authHeader: "Authorization", authHeader: "Authorization",
fixture: func(streaming, agentic bool) []byte { fixture: func(streaming, agentic bool) []byte {
@@ -196,6 +206,10 @@ func TestInterception_KeyFailover(t *testing.T) {
expectedKeyStates []keypool.KeyState expectedKeyStates []keypool.KeyState
expectedSeenKeys []string expectedSeenKeys []string
expectedBodyContains 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. // 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 { responses: func(s testutil.UpstreamResponse) []testutil.UpstreamResponse {
return []testutil.UpstreamResponse{errResp(http.StatusTooManyRequests, "5"), s} return []testutil.UpstreamResponse{errResp(http.StatusTooManyRequests, "5"), s}
}, },
expectedStatus: http.StatusOK, expectedStatus: http.StatusOK,
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid}, expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid},
expectedSeenKeys: []string{k0, k1}, expectedSeenKeys: []string{k0, k1},
expectedTransitions: map[string]int{"rate_limited": 1},
}, },
{ {
// A 401 marks the key permanent and fails over to the next one. // 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 { responses: func(s testutil.UpstreamResponse) []testutil.UpstreamResponse {
return []testutil.UpstreamResponse{errResp(http.StatusUnauthorized, ""), s} return []testutil.UpstreamResponse{errResp(http.StatusUnauthorized, ""), s}
}, },
expectedStatus: http.StatusOK, expectedStatus: http.StatusOK,
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid}, expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid},
expectedSeenKeys: []string{k0, k1}, expectedSeenKeys: []string{k0, k1},
expectedTransitions: map[string]int{"unauthorized": 1},
}, },
{ {
// A 403 marks the key permanent and fails over to the next one. // 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 { responses: func(s testutil.UpstreamResponse) []testutil.UpstreamResponse {
return []testutil.UpstreamResponse{errResp(http.StatusForbidden, ""), s} return []testutil.UpstreamResponse{errResp(http.StatusForbidden, ""), s}
}, },
expectedStatus: http.StatusOK, expectedStatus: http.StatusOK,
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid}, expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStateValid},
expectedSeenKeys: []string{k0, k1}, expectedSeenKeys: []string{k0, k1},
expectedTransitions: map[string]int{"forbidden": 1},
}, },
{ {
// Every key is rate-limited, so the pool is exhausted and the // 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,
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. // Every key is unauthorized, so the pool is permanently exhausted.
@@ -271,9 +290,11 @@ func TestInterception_KeyFailover(t *testing.T) {
errResp(http.StatusUnauthorized, ""), errResp(http.StatusUnauthorized, ""),
} }
}, },
expectedStatus: http.StatusBadGateway, expectedStatus: http.StatusBadGateway,
expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStatePermanent}, expectedKeyStates: []keypool.KeyState{keypool.KeyStatePermanent, keypool.KeyStatePermanent},
expectedSeenKeys: []string{k0, k1}, 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. // 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.Run(ic.name+"/"+mode+"/"+tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
reg := prometheus.NewRegistry()
m := metrics.NewMetrics(reg)
var pool *keypool.Pool var pool *keypool.Pool
if len(tc.keys) > 0 { if len(tc.keys) > 0 {
var err error 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) require.NoError(t, err)
} }
@@ -353,6 +376,38 @@ func TestInterception_KeyFailover(t *testing.T) {
if tc.expectedBodyContains != "" { if tc.expectedBodyContains != "" {
assert.Contains(t, w.Body.String(), tc.expectedBodyContains, "response body") 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 expectedKeyStates []keypool.KeyState
expectedSeenKeys []string expectedSeenKeys []string
expectedBodyContains 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 // expectErr is true when ProcessRequest returns an error because the
// pool is exhausted. // pool is exhausted.
expectErr bool expectErr bool
@@ -403,9 +462,10 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
responses: func(toolCall, final testutil.UpstreamResponse) []testutil.UpstreamResponse { responses: func(toolCall, final testutil.UpstreamResponse) []testutil.UpstreamResponse {
return []testutil.UpstreamResponse{toolCall, errResp(http.StatusTooManyRequests, "5"), final} return []testutil.UpstreamResponse{toolCall, errResp(http.StatusTooManyRequests, "5"), final}
}, },
expectedStatus: http.StatusOK, expectedStatus: http.StatusOK,
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid}, expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateValid},
expectedSeenKeys: []string{k0, k0, k1}, expectedSeenKeys: []string{k0, k0, k1},
expectedTransitions: map[string]int{"rate_limited": 1},
}, },
{ {
// The continuation is rate-limited on every key, exhausting the pool. // 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", expectedBodyContains: "all configured keys are rate-limited",
expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateTemporary}, expectedKeyStates: []keypool.KeyState{keypool.KeyStateTemporary, keypool.KeyStateTemporary},
expectedSeenKeys: []string{k0, k0, k1}, expectedSeenKeys: []string{k0, k0, k1},
expectedTransitions: map[string]int{"rate_limited": 2},
expectedExhaustions: map[string]int{"rate_limited": 1},
expectErr: true, expectErr: true,
}, },
} }
@@ -434,7 +496,9 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
t.Run(ic.name+"/"+mode+"/"+tc.name, func(t *testing.T) { t.Run(ic.name+"/"+mode+"/"+tc.name, func(t *testing.T) {
t.Parallel() 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) require.NoError(t, err)
fixture := fixtures.Parse(t, ic.fixture(streaming, true)) fixture := fixtures.Parse(t, ic.fixture(streaming, true))
@@ -487,6 +551,32 @@ func TestInterception_AgenticLoopFailover(t *testing.T) {
if tc.expectedBodyContains != "" { if tc.expectedBodyContains != "" {
assert.Contains(t, w.Body.String(), tc.expectedBodyContains, "response body") 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))
}
}
}) })
} }
} }
+2 -3
View File
@@ -583,9 +583,8 @@ func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key,
if !errors.As(err, &apiErr) { if !errors.As(err, &apiErr) {
return false return false
} }
return keypool.MarkKeyOnStatus( return i.cfg.KeyPool.MarkKeyOnStatus(
ctx, key, apiErr.Response, ctx, key, apiErr.Response, i.logger,
i.logger, i.providerName,
) )
} }
@@ -1122,7 +1122,7 @@ func TestMarkKeyOnError(t *testing.T) {
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() 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) require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next() key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr) require.Nil(t, keyPoolErr)
+21 -12
View File
@@ -105,9 +105,16 @@ func (i *BlockingInterception) ProcessRequest(w http.ResponseWriter, r *http.Req
// Accumulate usage across the entire streaming interaction (including tool reinvocations). // Accumulate usage across the entire streaming interaction (including tool reinvocations).
var cumulativeUsage anthropic.Usage 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 { for {
// TODO add outer loop span (https://github.com/coder/aibridge/issues/67) // 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 err != nil {
if eventstream.IsConnError(err) { if eventstream.IsConnError(err) {
// Can't write a response, just error out. // 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 // newMessage routes between BYOK (single attempt) and centralized
// failover. // failover, returning the upstream message, the number of key attempts
func (i *BlockingInterception) newMessage(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, error) { // 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. // BYOK: single attempt, no failover.
if i.cfg.KeyPool == nil { 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) 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...) return svc.New(ctx, anthropic.MessageNewParams{}, opts...)
} }
// newMessageWithKeyFailover walks the centralized key pool, // newMessageWithKeyFailover walks the centralized key pool, trying each key
// trying each key until one succeeds or the pool is exhausted. // until one succeeds or the pool is exhausted. Keys are marked temporary on
// Keys are marked temporary on 429 and permanent on 401/403. // 429 and permanent on 401/403. Errors that aren't key-specific don't trigger
// Errors that aren't key-specific don't trigger failover and // failover and are returned to the caller. It returns the upstream message,
// are returned to the caller. // the number of key attempts made for this call, and any error.
func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, error) { func (i *BlockingInterception) newMessageWithKeyFailover(ctx context.Context, svc anthropic.MessageService) (*anthropic.Message, int, error) {
walker := i.cfg.KeyPool.Walker() walker := i.cfg.KeyPool.Walker()
for { for {
key, keyPoolErr := walker.Next() key, keyPoolErr := walker.Next()
if keyPoolErr != nil { if keyPoolErr != nil {
return nil, keyPoolErr return nil, walker.Attempts(), keyPoolErr
} }
// Record the key in use so the hint reflects the last attempted key. // Record the key in use so the hint reflects the last attempted key.
i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) 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): // Either success (msg, nil) or a non-key error (nil, err):
// nothing to retry, return as-is. // nothing to retry, return as-is.
return msg, err return msg, walker.Attempts(), err
} }
} }
+7
View File
@@ -153,6 +153,11 @@ func (i *StreamingInterception) ProcessRequest(w http.ResponseWriter, r *http.Re
var lastErr error var lastErr error
var interceptionErr 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 isFirst := true
newStream: newStream:
for { for {
@@ -208,6 +213,8 @@ newStream:
) )
} }
totalKeyAttempts += walker.Attempts()
stream := i.newStream(streamCtx, svc, streamOpts...) stream := i.newStream(streamCtx, svc, streamOpts...)
var message anthropic.Message var message anthropic.Message
+2 -3
View File
@@ -181,9 +181,8 @@ func (i *responsesInterceptionBase) markKeyOnError(ctx context.Context, key *key
if !errors.As(err, &apiErr) { if !errors.As(err, &apiErr) {
return false return false
} }
return keypool.MarkKeyOnStatus( return i.cfg.KeyPool.MarkKeyOnStatus(
ctx, key, apiErr.Response, ctx, key, apiErr.Response, i.logger,
i.logger, i.providerName,
) )
} }
@@ -440,7 +440,7 @@ func TestMarkKeyOnError(t *testing.T) {
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() 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) require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next() key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr) require.Nil(t, keyPoolErr)
+22 -13
View File
@@ -86,6 +86,11 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r *
} }
shouldLoop := true 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 { for shouldLoop {
srv := i.newResponsesService() srv := i.newResponsesService()
respCopy = responseCopier{} respCopy = responseCopier{}
@@ -99,7 +104,9 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r *
opts = append(opts, intercept.ActorHeadersAsOpenAIOpts(actor)...) 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 // The failover loop may return a keypool exhaustion
// error. Render it here. // error. Render it here.
@@ -146,12 +153,14 @@ func (i *BlockingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r *
return errors.Join(upstreamErr, err) return errors.Join(upstreamErr, err)
} }
// newResponse routes between BYOK (single attempt) and // newResponse routes between BYOK (single attempt) and centralized failover,
// centralized failover. // returning the upstream response, the number of key attempts made for this
func (i *BlockingResponsesInterceptor) newResponse(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, error) { // 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. // BYOK: single attempt, no failover.
if i.cfg.KeyPool == nil { 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) 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...) return srv.New(ctx, responses.ResponseNewParams{}, opts...)
} }
// newResponseWithKeyFailover walks the centralized key pool, // newResponseWithKeyFailover walks the centralized key pool, trying each key
// trying each key until one succeeds or the pool is exhausted. // until one succeeds or the pool is exhausted. Keys are marked temporary on
// Keys are marked temporary on 429 and permanent on 401/403. // 429 and permanent on 401/403. Errors that aren't key-specific don't trigger
// Errors that aren't key-specific don't trigger failover and // failover and are returned to the caller. It returns the upstream response,
// are returned to the caller. // 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, error) { func (i *BlockingResponsesInterceptor) newResponseWithKeyFailover(ctx context.Context, srv responses.ResponseService, opts []option.RequestOption) (*responses.Response, int, error) {
walker := i.cfg.KeyPool.Walker() walker := i.cfg.KeyPool.Walker()
for { for {
key, keyPoolErr := walker.Next() key, keyPoolErr := walker.Next()
if keyPoolErr != nil { if keyPoolErr != nil {
return nil, keyPoolErr return nil, walker.Attempts(), keyPoolErr
} }
// Record the key in use so the hint reflects the last attempted key. // Record the key in use so the hint reflects the last attempted key.
i.credential = intercept.NewCredentialInfo(intercept.CredentialKindCentralized, key.Value()) 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 // Either success (response, nil) or a non-key error
// (nil, err): nothing to retry, return as-is. // (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 shouldLoop := true
srv := i.newResponsesService() 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 { for shouldLoop {
shouldLoop = false shouldLoop = false
@@ -140,6 +145,7 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r
// agentic mode the inner loop buffers events // agentic mode the inner loop buffers events
// instead of streaming them downstream, so the // instead of streaming them downstream, so the
// SSE connection has not been opened yet. // SSE connection has not been opened yet.
totalKeyAttempts += walker.Attempts()
i.writeUpstreamError(w, intercept.ResponseErrorFromKeyPool(keyPoolErr)) i.writeUpstreamError(w, intercept.ResponseErrorFromKeyPool(keyPoolErr))
return xerrors.Errorf("key pool exhausted: %w", keyPoolErr) return xerrors.Errorf("key pool exhausted: %w", keyPoolErr)
} }
@@ -175,6 +181,8 @@ func (i *StreamingResponsesInterceptor) ProcessRequest(w http.ResponseWriter, r
break break
} }
totalKeyAttempts += walker.Attempts()
// func scope to defer steam.Close() // func scope to defer steam.Close()
err := func() error { err := func() error {
defer stream.Close() defer stream.Close()
@@ -78,7 +78,7 @@ func TestOpenAI_KeyFailover(t *testing.T) {
successBody = fix.NonStreaming() successBody = fix.NonStreaming()
} }
pool, err := keypool.New([]string{"k0", "k1"}, quartz.NewMock(t)) pool, err := keypool.New(config.ProviderOpenAI, []string{"k0", "k1"}, quartz.NewMock(t), nil)
require.NoError(t, err) require.NoError(t, err)
var requestCount atomic.Int32 var requestCount atomic.Int32
@@ -185,7 +185,7 @@ func TestAnthropic_KeyFailover(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
pool, err := keypool.New([]string{"k0", "k1"}, quartz.NewMock(t)) pool, err := keypool.New(config.ProviderAnthropic, []string{"k0", "k1"}, quartz.NewMock(t), nil)
require.NoError(t, err) require.NoError(t, err)
var requestCount atomic.Int32 var requestCount atomic.Int32
@@ -30,6 +30,7 @@ func (m *MockProvider) BridgedRoutes() []string { return m.Bridged }
func (m *MockProvider) PassthroughRoutes() []string { return m.Passthrough } func (m *MockProvider) PassthroughRoutes() []string { return m.Passthrough }
func (*MockProvider) AuthHeader() string { return "Authorization" } func (*MockProvider) AuthHeader() string { return "Authorization" }
func (*MockProvider) KeyPool() *keypool.Pool { return nil }
func (*MockProvider) KeyFailoverConfig(_ slog.Logger) keypool.KeyFailoverConfig { func (*MockProvider) KeyFailoverConfig(_ slog.Logger) keypool.KeyFailoverConfig {
return keypool.KeyFailoverConfig{} return keypool.KeyFailoverConfig{}
} }
+3 -3
View File
@@ -15,8 +15,7 @@ type KeyFailoverConfig struct {
// Pool is the key pool to walk. Nil disables key failover. // Pool is the key pool to walk. Nil disables key failover.
Pool *Pool Pool *Pool
ProviderName string Logger slog.Logger
Logger slog.Logger
// IsBYOK returns true when the request already carries // IsBYOK returns true when the request already carries
// user-supplied auth. BYOK requests skip key failover. // user-supplied auth. BYOK requests skip key failover.
@@ -70,6 +69,7 @@ func (t *keyFailoverTransport) RoundTrip(req *http.Request) (*http.Response, err
// Fresh walker per request, independent of other inflight requests. // Fresh walker per request, independent of other inflight requests.
walker := t.config.Pool.Walker() walker := t.config.Pool.Walker()
defer func() { t.config.Pool.RecordAttempts(walker.Attempts()) }()
for { for {
key, keyPoolErr := walker.Next() key, keyPoolErr := walker.Next()
if keyPoolErr != nil { if keyPoolErr != nil {
@@ -95,7 +95,7 @@ func (t *keyFailoverTransport) RoundTrip(req *http.Request) (*http.Response, err
return resp, rtErr return resp, rtErr
} }
// MarkKeyOnStatus returns true on key-specific failures (e.g. 401/403/429). // MarkKeyOnStatus returns true on key-specific failures (e.g. 401/403/429).
if MarkKeyOnStatus(req.Context(), key, resp, t.config.Logger, t.config.ProviderName) { if t.config.Pool.MarkKeyOnStatus(req.Context(), key, resp, t.config.Logger) {
// Drain and retry with the next key. // Drain and retry with the next key.
_, _ = io.Copy(io.Discard, resp.Body) _, _ = io.Copy(io.Discard, resp.Body)
_ = resp.Body.Close() _ = resp.Body.Close()
+1 -1
View File
@@ -28,7 +28,7 @@ func (*fakeRoundTripper) RoundTrip(*http.Request) (*http.Response, error) {
func TestNewKeyFailoverTransport(t *testing.T) { func TestNewKeyFailoverTransport(t *testing.T) {
t.Parallel() t.Parallel()
pool, err := keypool.New([]string{"k0"}, quartz.NewMock(t)) pool, err := keypool.New("test-provider", []string{"k0"}, quartz.NewMock(t), nil)
require.NoError(t, err) require.NoError(t, err)
tests := []struct { tests := []struct {
+14 -5
View File
@@ -11,12 +11,11 @@ import (
// status code from resp (429 for temporary, 401 or 403 for // status code from resp (429 for temporary, 401 or 403 for
// permanent). Returns true if the status was a key-specific // permanent). Returns true if the status was a key-specific
// failover trigger so callers can retry with the next key. // failover trigger so callers can retry with the next key.
func MarkKeyOnStatus( func (p *Pool) MarkKeyOnStatus(
ctx context.Context, ctx context.Context,
key *Key, key *Key,
resp *http.Response, resp *http.Response,
logger slog.Logger, logger slog.Logger,
providerName string,
) bool { ) bool {
if resp == nil { if resp == nil {
return false return false
@@ -29,8 +28,11 @@ func MarkKeyOnStatus(
cooldown = defaultCooldown cooldown = defaultCooldown
} }
if key.MarkTemporary(cooldown) { if key.MarkTemporary(cooldown) {
if p.metrics != nil {
p.metrics.KeyPoolStateTransitions.WithLabelValues(p.providerName, reasonRateLimited).Inc()
}
logger.Info(ctx, "key marked temporary", logger.Info(ctx, "key marked temporary",
slog.F("provider", providerName), slog.F("provider", p.providerName),
slog.F("api_key_hint", key.Hint()), slog.F("api_key_hint", key.Hint()),
slog.F("status", statusCode), slog.F("status", statusCode),
slog.F("cooldown", cooldown)) slog.F("cooldown", cooldown))
@@ -38,15 +40,22 @@ func MarkKeyOnStatus(
return true return true
case http.StatusUnauthorized, http.StatusForbidden: case http.StatusUnauthorized, http.StatusForbidden:
if key.MarkPermanent() { if key.MarkPermanent() {
if p.metrics != nil {
reason := reasonUnauthorized
if statusCode == http.StatusForbidden {
reason = reasonForbidden
}
p.metrics.KeyPoolStateTransitions.WithLabelValues(p.providerName, reason).Inc()
}
logger.Warn(ctx, "key marked permanent", logger.Warn(ctx, "key marked permanent",
slog.F("provider", providerName), slog.F("provider", p.providerName),
slog.F("api_key_hint", key.Hint()), slog.F("api_key_hint", key.Hint()),
slog.F("status", statusCode)) slog.F("status", statusCode))
} }
return true return true
default: default:
logger.Debug(ctx, "status is not a key failover trigger", logger.Debug(ctx, "status is not a key failover trigger",
slog.F("provider", providerName), slog.F("provider", p.providerName),
slog.F("status", statusCode)) slog.F("status", statusCode))
return false return false
} }
+29 -3
View File
@@ -6,11 +6,14 @@ import (
"testing" "testing"
"time" "time"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"cdr.dev/slog/v3/sloggers/slogtest" "cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/aibridge/keypool" "github.com/coder/coder/v2/aibridge/keypool"
"github.com/coder/coder/v2/aibridge/metrics"
codertestutil "github.com/coder/coder/v2/testutil"
"github.com/coder/quartz" "github.com/coder/quartz"
) )
@@ -24,6 +27,9 @@ func TestMarkKeyOnStatus(t *testing.T) {
expectedReturn bool expectedReturn bool
expectedState keypool.KeyState expectedState keypool.KeyState
expectedCooldown time.Duration expectedCooldown time.Duration
// expectedReason is the transition metric's reason label, or
// empty when no transition is expected.
expectedReason string
}{ }{
{ {
// 429 with standard Retry-After header (seconds). // 429 with standard Retry-After header (seconds).
@@ -33,6 +39,7 @@ func TestMarkKeyOnStatus(t *testing.T) {
expectedReturn: true, expectedReturn: true,
expectedState: keypool.KeyStateTemporary, expectedState: keypool.KeyStateTemporary,
expectedCooldown: 5 * time.Second, expectedCooldown: 5 * time.Second,
expectedReason: "rate_limited",
}, },
{ {
// 429 with retry-after-ms header (milliseconds). // 429 with retry-after-ms header (milliseconds).
@@ -42,6 +49,7 @@ func TestMarkKeyOnStatus(t *testing.T) {
expectedReturn: true, expectedReturn: true,
expectedState: keypool.KeyStateTemporary, expectedState: keypool.KeyStateTemporary,
expectedCooldown: 1500 * time.Millisecond, expectedCooldown: 1500 * time.Millisecond,
expectedReason: "rate_limited",
}, },
{ {
// 429 without headers falls back to default cooldown. // 429 without headers falls back to default cooldown.
@@ -50,18 +58,21 @@ func TestMarkKeyOnStatus(t *testing.T) {
expectedReturn: true, expectedReturn: true,
expectedState: keypool.KeyStateTemporary, expectedState: keypool.KeyStateTemporary,
expectedCooldown: 60 * time.Second, expectedCooldown: 60 * time.Second,
expectedReason: "rate_limited",
}, },
{ {
name: "401_marks_permanent", name: "401_marks_permanent",
statusCode: http.StatusUnauthorized, statusCode: http.StatusUnauthorized,
expectedReturn: true, expectedReturn: true,
expectedState: keypool.KeyStatePermanent, expectedState: keypool.KeyStatePermanent,
expectedReason: "unauthorized",
}, },
{ {
name: "403_marks_permanent", name: "403_marks_permanent",
statusCode: http.StatusForbidden, statusCode: http.StatusForbidden,
expectedReturn: true, expectedReturn: true,
expectedState: keypool.KeyStatePermanent, expectedState: keypool.KeyStatePermanent,
expectedReason: "forbidden",
}, },
{ {
name: "200_does_not_mark", name: "200_does_not_mark",
@@ -85,11 +96,15 @@ func TestMarkKeyOnStatus(t *testing.T) {
}, },
} }
const providerName = "test-provider"
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0"}, clk) reg := prometheus.NewRegistry()
m := metrics.NewMetrics(reg)
pool, err := keypool.New(providerName, []string{"key-0"}, clk, m)
require.NoError(t, err) require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next() key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr) require.Nil(t, keyPoolErr)
@@ -102,19 +117,30 @@ func TestMarkKeyOnStatus(t *testing.T) {
resp.Header.Set(k, v) resp.Header.Set(k, v)
} }
got := keypool.MarkKeyOnStatus( got := pool.MarkKeyOnStatus(
context.Background(), context.Background(),
key, key,
resp, resp,
// 401 and 403 cases legitimately log at error // 401 and 403 cases legitimately log at error
// level when marking a key permanent. // level when marking a key permanent.
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
"test",
) )
assert.Equal(t, tc.expectedReturn, got) assert.Equal(t, tc.expectedReturn, got)
assert.Equal(t, tc.expectedState, key.State()) assert.Equal(t, tc.expectedState, key.State())
gathered, err := reg.Gather()
require.NoError(t, err)
// A state transition records one event under its reason,
// and other reasons record none.
for _, reason := range []string{"rate_limited", "unauthorized", "forbidden"} {
if reason == tc.expectedReason {
assert.True(t, codertestutil.PromCounterHasValue(t, gathered, 1, "key_pool_state_transitions_total", providerName, reason))
} else {
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_state_transitions_total", providerName, reason))
}
}
// Verify cooldown was set to the expected duration: // Verify cooldown was set to the expected duration:
// advancing by exactly that amount returns the key // advancing by exactly that amount returns the key
// to valid. // to valid.
+67 -13
View File
@@ -7,6 +7,7 @@ import (
"golang.org/x/xerrors" "golang.org/x/xerrors"
"github.com/coder/coder/v2/aibridge/metrics"
"github.com/coder/coder/v2/aibridge/utils" "github.com/coder/coder/v2/aibridge/utils"
"github.com/coder/quartz" "github.com/coder/quartz"
) )
@@ -53,23 +54,35 @@ func (e *Error) Error() string {
} }
// KeyState represents the current state of a key in the pool. // KeyState represents the current state of a key in the pool.
type KeyState int type KeyState string
const ( const (
// KeyStateValid means the key is available for use. // KeyStateValid means the key is available for use.
KeyStateValid KeyState = iota KeyStateValid KeyState = "valid"
// KeyStateTemporary means the key is temporarily unavailable // KeyStateTemporary means the key is temporarily unavailable
// (e.g. rate-limited) and will recover after a cooldown. // (e.g. rate-limited) and will recover after a cooldown.
KeyStateTemporary KeyStateTemporary KeyState = "temporary"
// KeyStatePermanent means the key is permanently unavailable // KeyStatePermanent means the key is permanently unavailable
// (e.g. revoked or unauthorized) until process restart. // (e.g. revoked or unauthorized) until process restart.
KeyStatePermanent KeyStatePermanent KeyState = "permanent"
) )
// defaultCooldown is applied when a key is marked temporary // defaultCooldown is applied when a key is marked temporary
// with a zero or negative cooldown duration. // with a zero or negative cooldown duration.
const defaultCooldown = 60 * time.Second const defaultCooldown = 60 * time.Second
// Metric label values for the key pool failover metrics.
const (
// Reasons for a key_pool_state_transitions_total event.
reasonRateLimited = "rate_limited"
reasonUnauthorized = "unauthorized"
reasonForbidden = "forbidden"
// Outcomes for a key_pool_exhaustions_total event.
outcomeRateLimited = "rate_limited"
outcomeAuthFailed = "auth_failed"
)
// Key holds a key value and its runtime state. // Key holds a key value and its runtime state.
type Key struct { type Key struct {
value string value string
@@ -83,18 +96,33 @@ type Key struct {
// Pool manages a set of keys with state tracking and // Pool manages a set of keys with state tracking and
// cooldown expiry. It is safe for concurrent use. // cooldown expiry. It is safe for concurrent use.
type Pool struct { type Pool struct {
keys []Key keys []Key
metrics *metrics.Metrics
providerName string
} }
// New creates a pool from the given keys. All keys start in // RecordAttempts records the total number of keys tried across an
// the valid state. Returns ErrNoKeys if keys is empty and // interception. Each upstream request uses its own walker, so the
// ErrDuplicateKey if any key appears more than once. // total sums the attempts across those per-request walkers. Call it
func New(keys []string, clk quartz.Clock) (*Pool, error) { // once when the interception finishes.
func (p *Pool) RecordAttempts(attempts int) {
if p == nil || p.metrics == nil || attempts == 0 {
return
}
p.metrics.KeyPoolFailoverAttempts.WithLabelValues(p.providerName).Observe(float64(attempts))
}
// New creates a pool from the given keys, labeled by providerName in its
// metrics and logs. All keys start in the valid state. Returns ErrNoKeys
// if keys is empty and ErrDuplicateKey if any key appears more than once.
func New(providerName string, keys []string, clk quartz.Clock, m *metrics.Metrics) (*Pool, error) {
if len(keys) == 0 { if len(keys) == 0 {
return nil, ErrNoKeys return nil, ErrNoKeys
} }
pool := &Pool{ pool := &Pool{
keys: make([]Key, len(keys)), keys: make([]Key, len(keys)),
metrics: m,
providerName: providerName,
} }
seen := make(map[string]struct{}, len(keys)) seen := make(map[string]struct{}, len(keys))
@@ -231,6 +259,20 @@ func (p *Pool) keyPoolError() *Error {
return &Error{Kind: ErrorKindPermanent} return &Error{Kind: ErrorKindPermanent}
} }
// recordExhaustion increments the exhaustion counter for the outcome
// implied by err.Kind: a rate-limited pool can recover, a permanent
// one cannot.
func (p *Pool) recordExhaustion(err *Error) {
if p.metrics == nil {
return
}
outcome := outcomeRateLimited
if err.Kind == ErrorKindPermanent {
outcome = outcomeAuthFailed
}
p.metrics.KeyPoolExhaustions.WithLabelValues(p.providerName, outcome).Inc()
}
// PoolState returns a snapshot of each key's state in the pool's // PoolState returns a snapshot of each key's state in the pool's
// original order, used by tests and other diagnostic callers. Use // original order, used by tests and other diagnostic callers. Use
// Walker for the failover iteration path. // Walker for the failover iteration path.
@@ -246,8 +288,9 @@ func (p *Pool) PoolState() []KeyState {
// creates its own walker so that it can independently iterate // creates its own walker so that it can independently iterate
// through keys without interfering with other requests. // through keys without interfering with other requests.
type Walker struct { type Walker struct {
pool *Pool pool *Pool
pos int // Next index to consider. pos int // Next index to consider.
attempts int // Number of attempts, one per upstream HTTP request.
} }
// Walker creates a new Walker that follows a primary-with-fallback // Walker creates a new Walker that follows a primary-with-fallback
@@ -270,9 +313,20 @@ func (w *Walker) Next() (*Key, *Error) {
} }
// Key is available. // Key is available.
w.pos = i + 1 w.pos = i + 1
w.attempts++
return key, nil return key, nil
} }
// No keys available. // No keys available.
return nil, w.pool.keyPoolError() err := w.pool.keyPoolError()
w.pool.recordExhaustion(err)
return nil, err
}
// Attempts returns the number of keys this walker handed out.
func (w *Walker) Attempts() int {
if w == nil {
return 0
}
return w.attempts
} }
+34 -7
View File
@@ -5,10 +5,13 @@ import (
"testing" "testing"
"time" "time"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/coder/coder/v2/aibridge/keypool" "github.com/coder/coder/v2/aibridge/keypool"
"github.com/coder/coder/v2/aibridge/metrics"
codertestutil "github.com/coder/coder/v2/testutil"
"github.com/coder/quartz" "github.com/coder/quartz"
) )
@@ -31,7 +34,7 @@ func TestNewKeyPool(t *testing.T) {
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
pool, err := keypool.New(tc.keys, quartz.NewMock(t)) pool, err := keypool.New("test-provider", tc.keys, quartz.NewMock(t), nil)
if tc.expectedErr != nil { if tc.expectedErr != nil {
require.ErrorIs(t, err, tc.expectedErr) require.ErrorIs(t, err, tc.expectedErr)
return return
@@ -125,7 +128,7 @@ func TestState(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0"}, clk) pool, err := keypool.New("test-provider", []string{"key-0"}, clk, nil)
require.NoError(t, err) require.NoError(t, err)
key := tc.setup(t, pool, clk) key := tc.setup(t, pool, clk)
@@ -204,7 +207,7 @@ func TestMarkTemporary(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0", "key-1"}, clk) pool, err := keypool.New("test-provider", []string{"key-0", "key-1"}, clk, nil)
require.NoError(t, err) require.NoError(t, err)
key := tc.setup(t, pool, clk) key := tc.setup(t, pool, clk)
@@ -267,7 +270,7 @@ func TestMarkPermanent(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0", "key-1"}, clk) pool, err := keypool.New("test-provider", []string{"key-0", "key-1"}, clk, nil)
require.NoError(t, err) require.NoError(t, err)
key := tc.setup(t, pool) key := tc.setup(t, pool)
@@ -498,11 +501,15 @@ func TestWalkerNext(t *testing.T) {
}, },
} }
const providerName = "test-provider"
for _, tc := range tests { for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New(tc.keys, clk) reg := prometheus.NewRegistry()
m := metrics.NewMetrics(reg)
pool, err := keypool.New(providerName, tc.keys, clk, m)
require.NoError(t, err) require.NoError(t, err)
tc.setup(t, pool) tc.setup(t, pool)
@@ -522,6 +529,26 @@ func TestWalkerNext(t *testing.T) {
// After all expected keys, the walker should be exhausted. // After all expected keys, the walker should be exhausted.
_, keyPoolErr := walker.Next() _, keyPoolErr := walker.Next()
require.Equal(t, tc.expectedErr, keyPoolErr) require.Equal(t, tc.expectedErr, keyPoolErr)
// The walker hands out one attempt per valid key before
// exhaustion.
assert.Equal(t, len(tc.expectedValid), walker.Attempts())
// Exhaustion records one event whose outcome reflects the
// error kind: rate-limited keys can recover, permanent cannot.
wantOutcome := "rate_limited"
if tc.expectedErr.Kind == keypool.ErrorKindPermanent {
wantOutcome = "auth_failed"
}
gathered, err := reg.Gather()
require.NoError(t, err)
for _, outcome := range []string{"rate_limited", "auth_failed"} {
if outcome == wantOutcome {
assert.True(t, codertestutil.PromCounterHasValue(t, gathered, 1, "key_pool_exhaustions_total", outcome, providerName))
} else {
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_exhaustions_total", outcome, providerName))
}
}
}) })
} }
} }
@@ -584,7 +611,7 @@ func TestKeyConcurrent(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0"}, clk) pool, err := keypool.New("test-provider", []string{"key-0"}, clk, nil)
require.NoError(t, err) require.NoError(t, err)
key, keyPoolErr := pool.Walker().Next() key, keyPoolErr := pool.Walker().Next()
require.Nil(t, keyPoolErr) require.Nil(t, keyPoolErr)
@@ -613,7 +640,7 @@ func TestWalkerIndependence(t *testing.T) {
t.Parallel() t.Parallel()
clk := quartz.NewMock(t) clk := quartz.NewMock(t)
pool, err := keypool.New([]string{"key-0", "key-1", "key-2"}, clk) pool, err := keypool.New("test-provider", []string{"key-0", "key-1", "key-2"}, clk, nil)
require.NoError(t, err) require.NoError(t, err)
walker := pool.Walker() walker := pool.Walker()
+55
View File
@@ -0,0 +1,55 @@
package keypool
import (
"github.com/prometheus/client_golang/prometheus"
)
// stateCollector reports the number of keys currently in each state per
// provider. State is read at scrape time rather than tracked via events
// because key recovery (cooldown expiry) happens lazily and is not observable
// as an event.
type stateCollector struct {
// pools returns the pools to report on. It is called on every scrape so
// reloaded pools are reflected.
pools func() []*Pool
desc *prometheus.Desc
}
// NewStateCollector returns a collector reporting the number of keys in
// each state, per provider.
func NewStateCollector(pools func() []*Pool) prometheus.Collector {
return &stateCollector{
pools: pools,
desc: prometheus.NewDesc(
"key_pool_state",
"The number of keys currently in each state (state: valid, temporary, permanent).",
[]string{"provider", "state"},
nil,
),
}
}
func (c *stateCollector) Describe(ch chan<- *prometheus.Desc) {
ch <- c.desc
}
func (c *stateCollector) Collect(ch chan<- prometheus.Metric) {
for _, pool := range c.pools() {
if pool == nil {
continue
}
counts := map[KeyState]int{
KeyStateValid: 0,
KeyStateTemporary: 0,
KeyStatePermanent: 0,
}
for _, state := range pool.PoolState() {
counts[state]++
}
for _, state := range []KeyState{KeyStateValid, KeyStateTemporary, KeyStatePermanent} {
ch <- prometheus.MustNewConstMetric(c.desc, prometheus.GaugeValue, float64(counts[state]), pool.providerName, string(state))
}
}
}
+114
View File
@@ -0,0 +1,114 @@
package keypool_test
import (
"fmt"
"testing"
"time"
"github.com/prometheus/client_golang/prometheus"
promtest "github.com/prometheus/client_golang/prometheus/testutil"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/aibridge/keypool"
codertestutil "github.com/coder/coder/v2/testutil"
"github.com/coder/quartz"
)
// newPool builds a pool named name with the given number of valid, temporary,
// and permanent keys.
func newPool(t *testing.T, clk quartz.Clock, name string, valid, temporary, permanent int) *keypool.Pool {
t.Helper()
keys := make([]string, valid+temporary+permanent)
for i := range keys {
keys[i] = fmt.Sprintf("%s-key-%d", name, i)
}
pool, err := keypool.New(name, keys, clk, nil)
require.NoError(t, err)
walker := pool.Walker()
for range temporary {
key, kpErr := walker.Next()
require.Nil(t, kpErr)
key.MarkTemporary(time.Minute)
}
for range permanent {
key, kpErr := walker.Next()
require.Nil(t, kpErr)
key.MarkPermanent()
}
return pool
}
func TestStateCollector(t *testing.T) {
t.Parallel()
type stateCount struct {
provider string
state string
count int
}
tests := []struct {
name string
pools func(t *testing.T, clk quartz.Clock) []*keypool.Pool
expectedStateCounts []stateCount
}{
{
name: "no_pools",
pools: func(*testing.T, quartz.Clock) []*keypool.Pool { return nil },
expectedStateCounts: nil,
},
{
name: "single_provider_mixed_states",
pools: func(t *testing.T, clk quartz.Clock) []*keypool.Pool {
return []*keypool.Pool{newPool(t, clk, "anthropic", 2, 1, 1)}
},
expectedStateCounts: []stateCount{
{"anthropic", "valid", 2},
{"anthropic", "temporary", 1},
{"anthropic", "permanent", 1},
},
},
{
name: "multiple_providers_nil_skipped",
pools: func(t *testing.T, clk quartz.Clock) []*keypool.Pool {
return []*keypool.Pool{
newPool(t, clk, "anthropic", 2, 1, 0),
nil,
newPool(t, clk, "openai", 1, 0, 1),
}
},
expectedStateCounts: []stateCount{
{"anthropic", "valid", 2},
{"anthropic", "temporary", 1},
{"anthropic", "permanent", 0},
{"openai", "valid", 1},
{"openai", "temporary", 0},
{"openai", "permanent", 1},
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
clk := quartz.NewMock(t)
pools := tc.pools(t, clk)
collector := keypool.NewStateCollector(func() []*keypool.Pool { return pools })
reg := prometheus.NewRegistry()
require.NoError(t, reg.Register(collector))
if len(tc.expectedStateCounts) == 0 {
require.Equal(t, 0, promtest.CollectAndCount(collector), "no key_pool_state series expected for empty pool list")
}
gathered, err := reg.Gather()
require.NoError(t, err)
for _, s := range tc.expectedStateCounts {
assert.True(t, codertestutil.PromGaugeHasValue(t, gathered, float64(s.count),
"key_pool_state", s.provider, s.state))
}
})
}
}
+33
View File
@@ -33,6 +33,13 @@ type Metrics struct {
CircuitBreakerState *prometheus.GaugeVec // Current state (0=closed, 0.5=half-open, 1=open) CircuitBreakerState *prometheus.GaugeVec // Current state (0=closed, 0.5=half-open, 1=open)
CircuitBreakerTrips *prometheus.CounterVec // Total times circuit opened CircuitBreakerTrips *prometheus.CounterVec // Total times circuit opened
CircuitBreakerRejects *prometheus.CounterVec // Requests rejected due to open circuit CircuitBreakerRejects *prometheus.CounterVec // Requests rejected due to open circuit
// Key pool failover metrics.
KeyPoolStateTransitions *prometheus.CounterVec // Key state transitions during failover.
KeyPoolExhaustions *prometheus.CounterVec // Times the pool ran out of usable keys.
// Keys attempted before success or exhaustion, per interception for
// bridged requests and per request for passthrough requests.
KeyPoolFailoverAttempts *prometheus.HistogramVec
} }
// NewMetrics creates AND registers metrics. It will panic if a collector has already been registered. // NewMetrics creates AND registers metrics. It will panic if a collector has already been registered.
@@ -128,5 +135,31 @@ func NewMetrics(reg prometheus.Registerer) *Metrics {
Name: "rejects_total", Name: "rejects_total",
Help: "Total number of requests rejected due to open circuit breaker.", Help: "Total number of requests rejected due to open circuit breaker.",
}, []string{"provider", "endpoint", "model"}), }, []string{"provider", "endpoint", "model"}),
// Key pool failover metrics.
// Pessimistic cardinality: 2 providers, 3 reasons = up to 6.
KeyPoolStateTransitions: promauto.With(reg).NewCounterVec(prometheus.CounterOpts{
Subsystem: "key_pool",
Name: "state_transitions_total",
Help: "The number of API key state transitions during failover " +
"(reason: rate_limited, unauthorized, forbidden).",
}, []string{"provider", "reason"}),
// Pessimistic cardinality: 2 providers, 2 outcomes = up to 4.
KeyPoolExhaustions: promauto.With(reg).NewCounterVec(prometheus.CounterOpts{
Subsystem: "key_pool",
Name: "exhaustions_total",
Help: "The number of times the key pool was exhausted with no usable key " +
"(outcome: rate_limited, auth_failed).",
}, []string{"provider", "outcome"}),
// Pessimistic cardinality: 2 providers, 7 buckets + 3 extra series (count, sum, +Inf) = up to 20.
KeyPoolFailoverAttempts: promauto.With(reg).NewHistogramVec(prometheus.HistogramOpts{
Subsystem: "key_pool",
Name: "failover_attempts",
Help: "The number of keys attempted before success or exhaustion, " +
"per interception for bridged requests and per request for " +
"passthrough requests.",
Buckets: []float64{1, 2, 3, 4, 5, 10, 25},
}, []string{"provider"}),
} }
} }
+42 -1
View File
@@ -13,6 +13,7 @@ import (
"sync/atomic" "sync/atomic"
"testing" "testing"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.opentelemetry.io/otel" "go.opentelemetry.io/otel"
@@ -22,6 +23,8 @@ import (
"github.com/coder/coder/v2/aibridge/internal/testutil" "github.com/coder/coder/v2/aibridge/internal/testutil"
"github.com/coder/coder/v2/aibridge/keypool" "github.com/coder/coder/v2/aibridge/keypool"
"github.com/coder/coder/v2/aibridge/provider" "github.com/coder/coder/v2/aibridge/provider"
"github.com/coder/coder/v2/coderd/coderdtest/promhelp"
codertestutil "github.com/coder/coder/v2/testutil"
"github.com/coder/quartz" "github.com/coder/quartz"
) )
@@ -384,6 +387,10 @@ func TestPassthrough_KeyFailover(t *testing.T) {
expectedRetryAfter string expectedRetryAfter string
// Expected key states after the request, by index in keys. // Expected key states after the request, by index in keys.
expectedKeyStates []keypool.KeyState expectedKeyStates []keypool.KeyState
// 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
}{ }{
{ {
// Given: 1 valid key returning 200. // Given: 1 valid key returning 200.
@@ -416,6 +423,7 @@ func TestPassthrough_KeyFailover(t *testing.T) {
keypool.KeyStateTemporary, keypool.KeyStateTemporary,
keypool.KeyStateValid, keypool.KeyStateValid,
}, },
expectedTransitions: map[string]int{"rate_limited": 1},
}, },
{ {
// Given: 2 keys; key-0 returns 401, key-1 returns 200. // Given: 2 keys; key-0 returns 401, key-1 returns 200.
@@ -432,6 +440,7 @@ func TestPassthrough_KeyFailover(t *testing.T) {
keypool.KeyStatePermanent, keypool.KeyStatePermanent,
keypool.KeyStateValid, keypool.KeyStateValid,
}, },
expectedTransitions: map[string]int{"unauthorized": 1},
}, },
{ {
// Given: 2 keys; key-0 returns 403, key-1 returns 200. // Given: 2 keys; key-0 returns 403, key-1 returns 200.
@@ -448,6 +457,7 @@ func TestPassthrough_KeyFailover(t *testing.T) {
keypool.KeyStatePermanent, keypool.KeyStatePermanent,
keypool.KeyStateValid, keypool.KeyStateValid,
}, },
expectedTransitions: map[string]int{"forbidden": 1},
}, },
{ {
// Given: 3 keys; all return 429 with cooldowns 5s, 3s, 10s. // Given: 3 keys; all return 429 with cooldowns 5s, 3s, 10s.
@@ -480,6 +490,8 @@ func TestPassthrough_KeyFailover(t *testing.T) {
keypool.KeyStateTemporary, keypool.KeyStateTemporary,
keypool.KeyStateTemporary, keypool.KeyStateTemporary,
}, },
expectedTransitions: map[string]int{"rate_limited": 3},
expectedExhaustions: map[string]int{"rate_limited": 1},
}, },
{ {
// Given: 2 keys; both return 401. // Given: 2 keys; both return 401.
@@ -496,6 +508,8 @@ func TestPassthrough_KeyFailover(t *testing.T) {
keypool.KeyStatePermanent, keypool.KeyStatePermanent,
keypool.KeyStatePermanent, keypool.KeyStatePermanent,
}, },
expectedTransitions: map[string]int{"unauthorized": 2},
expectedExhaustions: map[string]int{"auth_failed": 1},
}, },
{ {
// Given: 2 keys; key-0 returns 500. // Given: 2 keys; key-0 returns 500.
@@ -561,10 +575,13 @@ func TestPassthrough_KeyFailover(t *testing.T) {
})) }))
t.Cleanup(upstream.Close) t.Cleanup(upstream.Close)
reg := prometheus.NewRegistry()
m := NewMetrics(reg)
var pool *keypool.Pool var pool *keypool.Pool
if len(tc.keys) > 0 { if len(tc.keys) > 0 {
var err error var err error
pool, err = keypool.New(tc.keys, quartz.NewMock(t)) pool, err = keypool.New("test", tc.keys, quartz.NewMock(t), m)
require.NoError(t, err) require.NoError(t, err)
} }
@@ -587,6 +604,30 @@ func TestPassthrough_KeyFailover(t *testing.T) {
assert.Equal(t, tc.expectedRetryAfter, w.Header().Get("Retry-After"), "Retry-After header") assert.Equal(t, tc.expectedRetryAfter, w.Header().Get("Retry-After"), "Retry-After header")
if pool != nil { if pool != nil {
assert.Equal(t, tc.expectedKeyStates, pool.PoolState(), "key states") assert.Equal(t, tc.expectedKeyStates, pool.PoolState(), "key states")
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", "test", reason))
} else {
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_state_transitions_total", "test", 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, "test"))
} else {
assert.False(t, codertestutil.PromCounterGathered(t, gathered, "key_pool_exhaustions_total", outcome, "test"))
}
}
// One observation per request, summing the keys tried.
hist := promhelp.HistogramValue(t, reg, "key_pool_failover_attempts", prometheus.Labels{"provider": "test"})
require.NotNil(t, hist)
assert.Equal(t, uint64(1), hist.GetSampleCount())
assert.Equal(t, float64(tc.expectedRequestCount), hist.GetSampleSum())
} }
}) })
} }
+7 -4
View File
@@ -69,7 +69,7 @@ func NewAnthropic(cfg config.Anthropic, bedrockCfg *config.AWSBedrock) *Anthropi
if cfg.KeyPool == nil && cfg.Key != "" { if cfg.KeyPool == nil && cfg.Key != "" {
// keypool.New only fails on empty or duplicate keys, // keypool.New only fails on empty or duplicate keys,
// neither possible with a single non-empty key. // neither possible with a single non-empty key.
pool, err := keypool.New([]string{cfg.Key}, quartz.NewReal()) pool, err := keypool.New(cfg.Name, []string{cfg.Key}, quartz.NewReal(), nil)
if err != nil { if err != nil {
panic(fmt.Sprintf("anthropic provider: build single-key pool: %s", err)) panic(fmt.Sprintf("anthropic provider: build single-key pool: %s", err))
} }
@@ -194,11 +194,14 @@ func (*Anthropic) AuthHeader() string {
return "X-Api-Key" return "X-Api-Key"
} }
func (p *Anthropic) KeyPool() *keypool.Pool {
return p.cfg.KeyPool
}
func (p *Anthropic) KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverConfig { func (p *Anthropic) KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverConfig {
return keypool.KeyFailoverConfig{ return keypool.KeyFailoverConfig{
Pool: p.cfg.KeyPool, Pool: p.cfg.KeyPool,
ProviderName: p.Name(), Logger: logger,
Logger: logger,
IsBYOK: func(r *http.Request) bool { IsBYOK: func(r *http.Request) bool {
return r.Header.Get("X-Api-Key") != "" || r.Header.Get("Authorization") != "" return r.Header.Get("X-Api-Key") != "" || r.Header.Get("Authorization") != ""
}, },
+2 -3
View File
@@ -55,7 +55,7 @@ func TestAnthropic_TypeAndName(t *testing.T) {
func TestNewAnthropic_KeyResolution(t *testing.T) { func TestNewAnthropic_KeyResolution(t *testing.T) {
t.Parallel() t.Parallel()
pool, err := keypool.New([]string{"pool-key-0", "pool-key-1"}, quartz.NewMock(t)) pool, err := keypool.New(config.ProviderAnthropic, []string{"pool-key-0", "pool-key-1"}, quartz.NewMock(t), nil)
require.NoError(t, err) require.NoError(t, err)
tests := []struct { tests := []struct {
@@ -323,7 +323,7 @@ func TestAnthropic_CreateInterceptor_BYOK(t *testing.T) {
func TestAnthropic_KeyFailoverConfig(t *testing.T) { func TestAnthropic_KeyFailoverConfig(t *testing.T) {
t.Parallel() t.Parallel()
pool, err := keypool.New([]string{"k0", "k1"}, quartz.NewMock(t)) pool, err := keypool.New(config.ProviderAnthropic, []string{"k0", "k1"}, quartz.NewMock(t), nil)
require.NoError(t, err) require.NoError(t, err)
p := NewAnthropic(config.Anthropic{KeyPool: pool}, nil) p := NewAnthropic(config.Anthropic{KeyPool: pool}, nil)
@@ -331,7 +331,6 @@ func TestAnthropic_KeyFailoverConfig(t *testing.T) {
cfg := p.KeyFailoverConfig(slog.Make()) cfg := p.KeyFailoverConfig(slog.Make())
assert.Same(t, pool, cfg.Pool, "Pool must be wired from the provider config") assert.Same(t, pool, cfg.Pool, "Pool must be wired from the provider config")
assert.Equal(t, config.ProviderAnthropic, cfg.ProviderName, "ProviderName must match the provider name")
require.NotNil(t, cfg.IsBYOK) require.NotNil(t, cfg.IsBYOK)
require.NotNil(t, cfg.InjectAuthKey) require.NotNil(t, cfg.InjectAuthKey)
require.NotNil(t, cfg.BuildKeyPoolResponse) require.NotNil(t, cfg.BuildKeyPoolResponse)
+5
View File
@@ -109,6 +109,11 @@ func (*Copilot) AuthHeader() string {
return "Authorization" return "Authorization"
} }
// KeyPool returns nil. Copilot is always BYOK and has no key pool.
func (*Copilot) KeyPool() *keypool.Pool {
return nil
}
// KeyFailoverConfig returns a config with a nil Pool, which makes // KeyFailoverConfig returns a config with a nil Pool, which makes
// the KeyFailoverTransport short-circuit. Copilot is always BYOK. // the KeyFailoverTransport short-circuit. Copilot is always BYOK.
func (*Copilot) KeyFailoverConfig(_ slog.Logger) keypool.KeyFailoverConfig { func (*Copilot) KeyFailoverConfig(_ slog.Logger) keypool.KeyFailoverConfig {
+1
View File
@@ -36,6 +36,7 @@ func (d *DisabledStub) RoutePrefix() string {
func (*DisabledStub) BridgedRoutes() []string { return nil } func (*DisabledStub) BridgedRoutes() []string { return nil }
func (*DisabledStub) PassthroughRoutes() []string { return nil } func (*DisabledStub) PassthroughRoutes() []string { return nil }
func (*DisabledStub) AuthHeader() string { return "" } func (*DisabledStub) AuthHeader() string { return "" }
func (*DisabledStub) KeyPool() *keypool.Pool { return nil }
func (*DisabledStub) KeyFailoverConfig(_ slog.Logger) keypool.KeyFailoverConfig { func (*DisabledStub) KeyFailoverConfig(_ slog.Logger) keypool.KeyFailoverConfig {
return keypool.KeyFailoverConfig{} return keypool.KeyFailoverConfig{}
} }
+7 -4
View File
@@ -59,7 +59,7 @@ func NewOpenAI(cfg config.OpenAI) *OpenAI {
if cfg.KeyPool == nil && cfg.Key != "" { if cfg.KeyPool == nil && cfg.Key != "" {
// keypool.New only fails on empty or duplicate keys, // keypool.New only fails on empty or duplicate keys,
// neither possible with a single non-empty key. // neither possible with a single non-empty key.
pool, err := keypool.New([]string{cfg.Key}, quartz.NewReal()) pool, err := keypool.New(cfg.Name, []string{cfg.Key}, quartz.NewReal(), nil)
if err != nil { if err != nil {
panic(fmt.Sprintf("openai provider: build single-key pool: %s", err)) panic(fmt.Sprintf("openai provider: build single-key pool: %s", err))
} }
@@ -194,11 +194,14 @@ func (*OpenAI) AuthHeader() string {
return "Authorization" return "Authorization"
} }
func (p *OpenAI) KeyPool() *keypool.Pool {
return p.cfg.KeyPool
}
func (p *OpenAI) KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverConfig { func (p *OpenAI) KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverConfig {
return keypool.KeyFailoverConfig{ return keypool.KeyFailoverConfig{
Pool: p.cfg.KeyPool, Pool: p.cfg.KeyPool,
ProviderName: p.Name(), Logger: logger,
Logger: logger,
IsBYOK: func(r *http.Request) bool { IsBYOK: func(r *http.Request) bool {
return r.Header.Get("Authorization") != "" return r.Header.Get("Authorization") != ""
}, },
+1 -2
View File
@@ -335,7 +335,7 @@ func TestOpenAI_CreateInterceptor(t *testing.T) {
func TestOpenAI_KeyFailoverConfig(t *testing.T) { func TestOpenAI_KeyFailoverConfig(t *testing.T) {
t.Parallel() t.Parallel()
pool, err := keypool.New([]string{"k0", "k1"}, quartz.NewMock(t)) pool, err := keypool.New(config.ProviderOpenAI, []string{"k0", "k1"}, quartz.NewMock(t), nil)
require.NoError(t, err) require.NoError(t, err)
p := NewOpenAI(config.OpenAI{KeyPool: pool}) p := NewOpenAI(config.OpenAI{KeyPool: pool})
@@ -343,7 +343,6 @@ func TestOpenAI_KeyFailoverConfig(t *testing.T) {
cfg := p.KeyFailoverConfig(slog.Make()) cfg := p.KeyFailoverConfig(slog.Make())
assert.Same(t, pool, cfg.Pool, "Pool must be wired from the provider config") assert.Same(t, pool, cfg.Pool, "Pool must be wired from the provider config")
assert.Equal(t, config.ProviderOpenAI, cfg.ProviderName, "ProviderName must match the provider name")
require.NotNil(t, cfg.IsBYOK) require.NotNil(t, cfg.IsBYOK)
require.NotNil(t, cfg.InjectAuthKey) require.NotNil(t, cfg.InjectAuthKey)
require.NotNil(t, cfg.BuildKeyPoolResponse) require.NotNil(t, cfg.BuildKeyPoolResponse)
+4
View File
@@ -83,6 +83,10 @@ type Provider interface {
// automatic key failover on passthrough routes. // automatic key failover on passthrough routes.
KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverConfig KeyFailoverConfig(logger slog.Logger) keypool.KeyFailoverConfig
// KeyPool returns the provider's key pool for centralized keys, or nil
// when the provider is BYOK only.
KeyPool() *keypool.Pool
// CircuitBreakerConfig returns the circuit breaker configuration for the provider. // CircuitBreakerConfig returns the circuit breaker configuration for the provider.
CircuitBreakerConfig() *config.CircuitBreaker CircuitBreakerConfig() *config.CircuitBreaker
+26 -22
View File
@@ -30,14 +30,12 @@ import (
// database on every ai_providers change event. The returned unsubscribe // database on every ai_providers change event. The returned unsubscribe
// function tears down the subscription; callers must invoke it // function tears down the subscription; callers must invoke it
// alongside Server.Close on shutdown. // alongside Server.Close on shutdown.
func newAIBridgeDaemon(coderAPI *coderd.API, providers []aibridge.Provider, cfg codersdk.AIBridgeConfig) (*aibridged.Server, func(), error) { func newAIBridgeDaemon(coderAPI *coderd.API, providers []aibridge.Provider, cfg codersdk.AIBridgeConfig, reg prometheus.Registerer, metrics *aibridge.Metrics) (*aibridged.Server, func(), error) {
ctx := context.Background() ctx := context.Background()
coderAPI.Logger.Debug(ctx, "starting in-memory aibridge daemon") coderAPI.Logger.Debug(ctx, "starting in-memory aibridge daemon")
logger := coderAPI.Logger.Named("aibridged") logger := coderAPI.Logger.Named("aibridged")
reg := prometheus.WrapRegistererWithPrefix("coder_aibridged_", coderAPI.PrometheusRegistry)
metrics := aibridge.NewMetrics(reg)
providerMetrics := aibridged.NewMetrics(reg) providerMetrics := aibridged.NewMetrics(reg)
tracer := coderAPI.TracerProvider.Tracer(tracing.TracerName) tracer := coderAPI.TracerProvider.Tracer(tracing.TracerName)
@@ -47,16 +45,20 @@ func newAIBridgeDaemon(coderAPI *coderd.API, providers []aibridge.Provider, cfg
return nil, nil, xerrors.Errorf("create request pool: %w", err) return nil, nil, xerrors.Errorf("create request pool: %w", err)
} }
// Report current key pool state per provider at scrape time.
reg.MustRegister(keypool.NewStateCollector(pool.KeyPools))
// Subscribe to ai_providers change events so the pool tracks the // Subscribe to ai_providers change events so the pool tracks the
// database without a restart. The boot-time `providers` snapshot // database without a restart. The boot-time `providers` snapshot
// derives from env config and serves as a fallback if the database // derives from env config and serves as a fallback if the database
// load fails inside the reloader. // load fails inside the reloader.
reloader := &poolDBReloader{ reloader := &poolDBReloader{
pool: pool, pool: pool,
db: coderAPI.Database, db: coderAPI.Database,
cfg: cfg, cfg: cfg,
logger: logger.Named("provider-loader"), logger: logger.Named("provider-loader"),
metrics: providerMetrics, aibridgeMetrics: metrics,
providerMetrics: providerMetrics,
} }
unsubscribe, err := aibridged.SubscribeProviderReload(ctx, coderAPI.Pubsub, reloader, logger.Named("provider-reload")) unsubscribe, err := aibridged.SubscribeProviderReload(ctx, coderAPI.Pubsub, reloader, logger.Named("provider-reload"))
if err != nil { if err != nil {
@@ -81,16 +83,17 @@ func newAIBridgeDaemon(coderAPI *coderd.API, providers []aibridge.Provider, cfg
// the live provider set from the database and forwarding it to the // the live provider set from the database and forwarding it to the
// pool. // pool.
type poolDBReloader struct { type poolDBReloader struct {
pool *aibridged.CachedBridgePool pool *aibridged.CachedBridgePool
db database.Store db database.Store
cfg codersdk.AIBridgeConfig cfg codersdk.AIBridgeConfig
logger slog.Logger logger slog.Logger
metrics *aibridged.Metrics aibridgeMetrics *aibridge.Metrics
providerMetrics *aibridged.Metrics
} }
func (r *poolDBReloader) Reload(ctx context.Context) error { func (r *poolDBReloader) Reload(ctx context.Context) error {
r.metrics.RecordReloadAttempt() r.providerMetrics.RecordReloadAttempt()
providers, outcomes, err := BuildProviders(ctx, r.db, r.cfg, r.logger) providers, outcomes, err := BuildProviders(ctx, r.db, r.cfg, r.logger, r.aibridgeMetrics)
if err != nil { if err != nil {
// Keep the previous snapshot in place: dropping all providers // Keep the previous snapshot in place: dropping all providers
// because the DB read failed would compound the visible failure // because the DB read failed would compound the visible failure
@@ -98,7 +101,7 @@ func (r *poolDBReloader) Reload(ctx context.Context) error {
return xerrors.Errorf("load ai providers from database: %w", err) return xerrors.Errorf("load ai providers from database: %w", err)
} }
r.pool.ReplaceProviders(providers) r.pool.ReplaceProviders(providers)
r.metrics.RecordReloadSuccess(outcomes) r.providerMetrics.RecordReloadSuccess(outcomes)
return nil return nil
} }
@@ -114,7 +117,7 @@ func (r *poolDBReloader) Reload(ctx context.Context) error {
// excluded from the returned snapshot; only a failure of the DB query // excluded from the returned snapshot; only a failure of the DB query
// itself is propagated. This keeps a single misconfigured row from // itself is propagated. This keeps a single misconfigured row from
// taking the whole daemon down. // taking the whole daemon down.
func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridgeConfig, logger slog.Logger) ([]aibridge.Provider, []aibridged.ProviderOutcome, error) { func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridgeConfig, logger slog.Logger, metrics *aibridge.Metrics) ([]aibridge.Provider, []aibridged.ProviderOutcome, error) {
//nolint:gocritic // AsAIBridged has a minimal permission set for this purpose. //nolint:gocritic // AsAIBridged has a minimal permission set for this purpose.
authCtx := dbauthz.AsAIBridged(ctx) authCtx := dbauthz.AsAIBridged(ctx)
@@ -172,7 +175,7 @@ func BuildProviders(ctx context.Context, db database.Store, cfg codersdk.AIBridg
if row.Enabled { if row.Enabled {
enabledCount++ enabledCount++
} }
prov, err := buildAIProviderFromRow(row, keysByProvider[row.ID], cfg) prov, err := buildAIProviderFromRow(row, keysByProvider[row.ID], cfg, metrics)
if err != nil { if err != nil {
outcome.Status = aibridged.ProviderStatusError outcome.Status = aibridged.ProviderStatusError
outcome.Err = err outcome.Err = err
@@ -210,6 +213,7 @@ func buildAIProviderFromRow(
row database.AIProvider, row database.AIProvider,
keys []database.AIProviderKey, keys []database.AIProviderKey,
cfg codersdk.AIBridgeConfig, cfg codersdk.AIBridgeConfig,
metrics *aibridge.Metrics,
) (aibridge.Provider, error) { ) (aibridge.Provider, error) {
if !row.Enabled { if !row.Enabled {
return disabledProviderFromRow(row) return disabledProviderFromRow(row)
@@ -243,7 +247,7 @@ func buildAIProviderFromRow(
var pool *keypool.Pool var pool *keypool.Pool
if len(keys) > 0 { if len(keys) > 0 {
var err error var err error
pool, err = buildAIProviderKeyPool(keys) pool, err = buildAIProviderKeyPool(row.Name, keys, metrics)
if err != nil { if err != nil {
return nil, xerrors.Errorf("%s key pool: %w", row.Type, err) return nil, xerrors.Errorf("%s key pool: %w", row.Type, err)
} }
@@ -275,7 +279,7 @@ func buildAIProviderFromRow(
var pool *keypool.Pool var pool *keypool.Pool
if len(keys) > 0 { if len(keys) > 0 {
var err error var err error
pool, err = buildAIProviderKeyPool(keys) pool, err = buildAIProviderKeyPool(row.Name, keys, metrics)
if err != nil { if err != nil {
return nil, xerrors.Errorf("anthropic key pool: %w", err) return nil, xerrors.Errorf("anthropic key pool: %w", err)
} }
@@ -314,12 +318,12 @@ func disabledProviderFromRow(row database.AIProvider) (aibridge.Provider, error)
// buildAIProviderKeyPool builds a [keypool.Pool]. Callers must check // buildAIProviderKeyPool builds a [keypool.Pool]. Callers must check
// len(keys) > 0 first; keypool.New rejects empty input. // len(keys) > 0 first; keypool.New rejects empty input.
func buildAIProviderKeyPool(keys []database.AIProviderKey) (*keypool.Pool, error) { func buildAIProviderKeyPool(providerName string, keys []database.AIProviderKey, metrics *aibridge.Metrics) (*keypool.Pool, error) {
raw := make([]string, 0, len(keys)) raw := make([]string, 0, len(keys))
for _, k := range keys { for _, k := range keys {
raw = append(raw, k.APIKey) raw = append(raw, k.APIKey)
} }
return keypool.New(raw, quartz.NewReal()) return keypool.New(providerName, raw, quartz.NewReal(), metrics)
} }
// bedrockConfigFromRow returns nil when the settings have no Bedrock // bedrockConfigFromRow returns nil when the settings have no Bedrock
+5 -5
View File
@@ -36,7 +36,7 @@ func buildFromEnv(t *testing.T, cfg codersdk.AIBridgeConfig) ([]aibridge.Provide
if err := coderd.SeedAIProvidersFromEnv(ctx, db, cfg, logger); err != nil { if err := coderd.SeedAIProvidersFromEnv(ctx, db, cfg, logger); err != nil {
return nil, err return nil, err
} }
providers, _, err := BuildProviders(ctx, db, cfg, logger) providers, _, err := BuildProviders(ctx, db, cfg, logger, nil)
return providers, err return providers, err
} }
@@ -325,7 +325,7 @@ func TestBuildProvidersSkipsBadRows(t *testing.T) {
Settings: sql.NullString{String: "not-json", Valid: true}, Settings: sql.NullString{String: "not-json", Valid: true},
}) })
providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger) providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger, nil)
require.NoError(t, err) require.NoError(t, err)
assert.Empty(t, providers) assert.Empty(t, providers)
require.Len(t, outcomes, 1) require.Len(t, outcomes, 1)
@@ -349,7 +349,7 @@ func TestBuildProvidersSkipsBadRows(t *testing.T) {
BaseUrl: "https://example.openai.azure.com/", BaseUrl: "https://example.openai.azure.com/",
}) })
providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger) providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger, nil)
require.NoError(t, err) require.NoError(t, err)
assert.Empty(t, providers) assert.Empty(t, providers)
require.Len(t, outcomes, 1) require.Len(t, outcomes, 1)
@@ -378,7 +378,7 @@ func TestBuildProvidersSkipsBadRows(t *testing.T) {
APIKey: "sk-good", APIKey: "sk-good",
}) })
providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger) providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger, nil)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, providers, 1) require.Len(t, providers, 1)
assert.Equal(t, "openai-good", providers[0].Name()) assert.Equal(t, "openai-good", providers[0].Name())
@@ -436,7 +436,7 @@ func TestBuildProvidersSkipsBadRows(t *testing.T) {
p.Enabled = false p.Enabled = false
}) })
providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger) providers, outcomes, err := BuildProviders(ctx, db, codersdk.AIBridgeConfig{}, logger, nil)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, providers, 1, "disabled providers stay in the snapshot so the bridge can serve a 503 sentinel") require.Len(t, providers, 1, "disabled providers stay in the snapshot so the bridge can serve a 503 sentinel")
assert.Equal(t, tc.row.Name, providers[0].Name()) assert.Equal(t, tc.row.Name, providers[0].Name())
+5 -2
View File
@@ -57,6 +57,7 @@ import (
"cdr.dev/slog/v3" "cdr.dev/slog/v3"
"cdr.dev/slog/v3/sloggers/sloghuman" "cdr.dev/slog/v3/sloggers/sloghuman"
"github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/buildinfo" "github.com/coder/coder/v2/buildinfo"
"github.com/coder/coder/v2/cli/clilog" "github.com/coder/coder/v2/cli/clilog"
"github.com/coder/coder/v2/cli/cliui" "github.com/coder/coder/v2/cli/cliui"
@@ -1061,12 +1062,14 @@ func (r *RootCmd) Server(newAPI func(context.Context, *coderd.Options) (*coderd.
// unconditionally when the bridge feature is enabled by config so // unconditionally when the bridge feature is enabled by config so
// chatd can use it regardless of license entitlement. // chatd can use it regardless of license entitlement.
if vals.AI.BridgeConfig.Enabled.Value() { if vals.AI.BridgeConfig.Enabled.Value() {
aibridgeProviders, _, err := BuildProviders(aibridgeInitCtx, options.Database, vals.AI.BridgeConfig, logger.Named("aibridge.providers")) aibridgeReg := prometheus.WrapRegistererWithPrefix("coder_aibridged_", coderAPI.PrometheusRegistry)
aibridgeMetrics := aibridge.NewMetrics(aibridgeReg)
aibridgeProviders, _, err := BuildProviders(aibridgeInitCtx, options.Database, vals.AI.BridgeConfig, logger.Named("aibridge.providers"), aibridgeMetrics)
if err != nil { if err != nil {
return xerrors.Errorf("build AI providers: %w", err) return xerrors.Errorf("build AI providers: %w", err)
} }
var unsubscribeProviderReload func() var unsubscribeProviderReload func()
aibridgeDaemon, unsubscribeProviderReload, err = newAIBridgeDaemon(coderAPI, aibridgeProviders, vals.AI.BridgeConfig) aibridgeDaemon, unsubscribeProviderReload, err = newAIBridgeDaemon(coderAPI, aibridgeProviders, vals.AI.BridgeConfig, aibridgeReg, aibridgeMetrics)
if err != nil { if err != nil {
return xerrors.Errorf("create aibridged: %w", err) return xerrors.Errorf("create aibridged: %w", err)
} }
+2 -2
View File
@@ -744,7 +744,7 @@ func TestBuildAIProviderFromRowSetsAPIDumpDir(t *testing.T) {
provider, err := buildAIProviderFromRow(tt.row, nil, codersdk.AIBridgeConfig{ provider, err := buildAIProviderFromRow(tt.row, nil, codersdk.AIBridgeConfig{
AllowBYOK: serpent.Bool(true), AllowBYOK: serpent.Bool(true),
APIDumpDir: serpent.String(dumpDir), APIDumpDir: serpent.String(dumpDir),
}) }, nil)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, dumpDir, provider.APIDumpDir()) assert.Equal(t, dumpDir, provider.APIDumpDir())
assert.Equal(t, tt.expectedType, provider.Type()) assert.Equal(t, tt.expectedType, provider.Type())
@@ -762,7 +762,7 @@ func TestBuildAIProviderFromRowBedrockWithoutSettings(t *testing.T) {
BaseUrl: "https://bedrock-runtime.us-east-1.amazonaws.com/", BaseUrl: "https://bedrock-runtime.us-east-1.amazonaws.com/",
}, nil, codersdk.AIBridgeConfig{ }, nil, codersdk.AIBridgeConfig{
AllowBYOK: serpent.Bool(true), AllowBYOK: serpent.Bool(true),
}) }, nil)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "bedrock provider has no bedrock credentials configured") assert.Contains(t, err.Error(), "bedrock provider has no bedrock credentials configured")
} }
+13
View File
@@ -17,6 +17,7 @@ import (
"cdr.dev/slog/v3" "cdr.dev/slog/v3"
"github.com/coder/coder/v2/aibridge" "github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/aibridge/keypool"
"github.com/coder/coder/v2/aibridge/mcp" "github.com/coder/coder/v2/aibridge/mcp"
"github.com/coder/coder/v2/aibridge/tracing" "github.com/coder/coder/v2/aibridge/tracing"
) )
@@ -148,6 +149,18 @@ func (p *CachedBridgePool) loadProviders() []aibridge.Provider {
return nil return nil
} }
// KeyPools returns the key pools of the current live providers.
func (p *CachedBridgePool) KeyPools() []*keypool.Pool {
providers := p.loadProviders()
pools := make([]*keypool.Pool, 0, len(providers))
for _, prov := range providers {
if pool := prov.KeyPool(); pool != nil {
pools = append(pools, pool)
}
}
return pools
}
// Acquire retrieves or creates a [*aibridge.RequestBridge] instance per given key. // Acquire retrieves or creates a [*aibridge.RequestBridge] instance per given key.
// //
// Each returned [*aibridge.RequestBridge] is safe for concurrent use. // Each returned [*aibridge.RequestBridge] is safe for concurrent use.
+95
View File
@@ -11,6 +11,8 @@ import (
"time" "time"
"github.com/google/uuid" "github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/trace" "go.opentelemetry.io/otel/trace"
"go.uber.org/mock/gomock" "go.uber.org/mock/gomock"
@@ -18,11 +20,13 @@ import (
"cdr.dev/slog/v3/sloggers/slogtest" "cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/aibridge" "github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/aibridge/config" "github.com/coder/coder/v2/aibridge/config"
"github.com/coder/coder/v2/aibridge/keypool"
"github.com/coder/coder/v2/aibridge/mcp" "github.com/coder/coder/v2/aibridge/mcp"
"github.com/coder/coder/v2/aibridge/mcpmock" "github.com/coder/coder/v2/aibridge/mcpmock"
"github.com/coder/coder/v2/coderd/aibridged" "github.com/coder/coder/v2/coderd/aibridged"
mock "github.com/coder/coder/v2/coderd/aibridged/aibridgedmock" mock "github.com/coder/coder/v2/coderd/aibridged/aibridgedmock"
"github.com/coder/coder/v2/testutil" "github.com/coder/coder/v2/testutil"
"github.com/coder/quartz"
) )
// TestPool validates the published behavior of [aibridged.CachedBridgePool]. // TestPool validates the published behavior of [aibridged.CachedBridgePool].
@@ -394,3 +398,94 @@ func (m *blockingMCPFactory) Build(ctx context.Context, _ aibridged.Request, _ t
} }
return nil, context.Canceled return nil, context.Canceled
} }
// TestPoolKeyPools verifies KeyPools returns the providers' pools, the pool
// wires failover metrics into them, and the state collector reflects live
// pool state, on both the initial set and reload.
func TestPoolKeyPools(t *testing.T) {
t.Parallel()
// Setup.
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
opts := aibridged.PoolOptions{MaxItems: 1, TTL: time.Minute}
clk := quartz.NewMock(t)
reg := prometheus.NewRegistry()
m := aibridge.NewMetrics(reg)
// markRateLimited drives one rate-limit transition on the pool's first
// key, recording a metric only if the pool has metrics attached.
markRateLimited := func(t *testing.T, pool *keypool.Pool) {
key, kpErr := pool.Walker().Next()
require.Nil(t, kpErr)
pool.MarkKeyOnStatus(context.Background(), key,
&http.Response{StatusCode: http.StatusTooManyRequests, Header: make(http.Header)}, logger)
}
// Given: provider "a" (2 keys), a BYOK provider with no key pool, and
// provider "b" (1 key).
poolA, err := keypool.New("a", []string{"a-key-0", "a-key-1"}, clk, m)
require.NoError(t, err)
poolB, err := keypool.New("b", []string{"b-key-0"}, clk, m)
require.NoError(t, err)
// When: the providers are loaded into a new bridge pool.
aibridgePool, err := aibridged.NewCachedBridgePool(opts, []aibridge.Provider{
aibridge.NewOpenAIProvider(config.OpenAI{Name: "a", KeyPool: poolA}),
aibridge.NewOpenAIProvider(config.OpenAI{Name: "byok"}),
aibridge.NewOpenAIProvider(config.OpenAI{Name: "b", KeyPool: poolB}),
}, logger, m, testTracer)
require.NoError(t, err)
t.Cleanup(func() { _ = aibridgePool.Shutdown(context.Background()) })
reg.MustRegister(keypool.NewStateCollector(aibridgePool.KeyPools))
// Then: KeyPools returns the non-BYOK pools, and the collector reports
// every key as valid.
require.Equal(t, []*keypool.Pool{poolA, poolB}, aibridgePool.KeyPools())
gathered, err := reg.Gather()
require.NoError(t, err)
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 2, "key_pool_state", "a", "valid"))
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 1, "key_pool_state", "b", "valid"))
// When: a key in pool "a" is rate-limited.
markRateLimited(t, poolA)
// Then: the transition is recorded (metrics were attached) and the key
// moves to temporary, which the collector reflects.
gathered, err = reg.Gather()
require.NoError(t, err)
assert.True(t, testutil.PromCounterHasValue(t, gathered, 1, "key_pool_state_transitions_total", "a", "rate_limited"))
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 1, "key_pool_state", "a", "valid"))
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 1, "key_pool_state", "a", "temporary"))
// When: the providers reload, dropping a key from "a", adding one to "b",
// and introducing a new provider "c".
poolA, err = keypool.New("a", []string{"a-key-0"}, clk, m)
require.NoError(t, err)
poolB, err = keypool.New("b", []string{"b-key-0", "b-key-1"}, clk, m)
require.NoError(t, err)
poolC, err := keypool.New("c", []string{"c-key-0"}, clk, m)
require.NoError(t, err)
aibridgePool.ReplaceProviders([]aibridge.Provider{
aibridge.NewOpenAIProvider(config.OpenAI{Name: "a", KeyPool: poolA}),
aibridge.NewOpenAIProvider(config.OpenAI{Name: "b", KeyPool: poolB}),
aibridge.NewOpenAIProvider(config.OpenAI{Name: "c", KeyPool: poolC}),
})
// Then: KeyPools, metric wiring, and pool state all follow the new set.
require.Equal(t, []*keypool.Pool{poolA, poolB, poolC}, aibridgePool.KeyPools())
gathered, err = reg.Gather()
require.NoError(t, err)
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 1, "key_pool_state", "a", "valid"))
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 2, "key_pool_state", "b", "valid"))
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 1, "key_pool_state", "c", "valid"))
// When: a key in the new pool "c" is rate-limited.
markRateLimited(t, poolC)
// Then: the transition is recorded and the key moves to temporary.
gathered, err = reg.Gather()
require.NoError(t, err)
assert.True(t, testutil.PromCounterHasValue(t, gathered, 1, "key_pool_state_transitions_total", "c", "rate_limited"))
assert.True(t, testutil.PromGaugeHasValue(t, gathered, 1, "key_pool_state", "c", "temporary"))
}
+2 -2
View File
@@ -64,7 +64,7 @@ func startTestAIBridgeDaemon(t *testing.T, api *coderd.API) *aibridged.Metrics {
cfg := api.DeploymentValues.AI.BridgeConfig cfg := api.DeploymentValues.AI.BridgeConfig
tracer := otel.Tracer("aibridge-reload-test") tracer := otel.Tracer("aibridge-reload-test")
providers, _, err := cli.BuildProviders(ctx, api.Database, cfg, logger) providers, _, err := cli.BuildProviders(ctx, api.Database, cfg, logger, nil)
require.NoError(t, err) require.NoError(t, err)
pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger.Named("pool"), nil, tracer) pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, providers, logger.Named("pool"), nil, tracer)
@@ -97,7 +97,7 @@ type testPoolReloader struct {
func (r *testPoolReloader) Reload(ctx context.Context) error { func (r *testPoolReloader) Reload(ctx context.Context) error {
defer r.metrics.RecordReloadAttempt() defer r.metrics.RecordReloadAttempt()
providers, outcomes, err := cli.BuildProviders(ctx, r.db, r.cfg, r.logger) providers, outcomes, err := cli.BuildProviders(ctx, r.db, r.cfg, r.logger, nil)
if err != nil { if err != nil {
return err return err
} }