mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add key pool failover metrics to aibridge (#25901)
## Description This PR adds Prometheus metrics for aibridge's API-key failover, giving visibility into key pool health and failover behavior per provider. The following metrics are introduced: - **`key_pool_state`** (gauge): number of keys currently in each state (`valid`, `temporary`, `permanent`) per provider, sampled at scrape time. - **`key_pool_state_transitions_total`** (counter): key state transitions during failover, labeled by `reason` (`rate_limited`, `unauthorized`, `forbidden`). - **`key_pool_exhaustions_total`** (counter): times a pool ran out of usable keys, labeled by `outcome` (`rate_limited`, `auth_failed`). - **`key_pool_failover_attempts`** (histogram): keys attempted before success or exhaustion (per interception for bridged requests, per request for passthrough). ## Changes - Moves `MarkKeyOnStatus` and key-pool error handling onto `*keypool.Pool`. - Attaches metrics to each provider's key pool at install time, on construction and on provider reload. - Adds a scrape-time state collector and a `KeyPools()` accessor on the bridge pool to feed it. - Tracks per-request key attempts in the bridged and passthrough failover paths. - Adds test coverage for the new metrics across the keypool unit tests, the bridged intercept failover tests, and the passthrough failover test. Closes https://github.com/coder/internal/issues/1447 Closes https://linear.app/codercom/issue/AIGOV-198/aibridge-key-failover-observability > [!NOTE] > Initially generated by Claude Opus 4.7, modified and reviewed by @ssncferreira
This commit is contained in:
@@ -226,9 +226,8 @@ func (i *interceptionBase) markKeyOnError(ctx context.Context, key *keypool.Key,
|
|||||||
if !errors.As(err, &apiErr) {
|
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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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))
|
||||||
|
}
|
||||||
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,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"}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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") != ""
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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") != ""
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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"))
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user