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:
@@ -15,8 +15,7 @@ type KeyFailoverConfig struct {
|
||||
// Pool is the key pool to walk. Nil disables key failover.
|
||||
Pool *Pool
|
||||
|
||||
ProviderName string
|
||||
Logger slog.Logger
|
||||
Logger slog.Logger
|
||||
|
||||
// IsBYOK returns true when the request already carries
|
||||
// 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.
|
||||
walker := t.config.Pool.Walker()
|
||||
defer func() { t.config.Pool.RecordAttempts(walker.Attempts()) }()
|
||||
for {
|
||||
key, keyPoolErr := walker.Next()
|
||||
if keyPoolErr != nil {
|
||||
@@ -95,7 +95,7 @@ func (t *keyFailoverTransport) RoundTrip(req *http.Request) (*http.Response, err
|
||||
return resp, rtErr
|
||||
}
|
||||
// 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.
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
|
||||
@@ -28,7 +28,7 @@ func (*fakeRoundTripper) RoundTrip(*http.Request) (*http.Response, error) {
|
||||
func TestNewKeyFailoverTransport(t *testing.T) {
|
||||
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)
|
||||
|
||||
tests := []struct {
|
||||
|
||||
@@ -11,12 +11,11 @@ import (
|
||||
// status code from resp (429 for temporary, 401 or 403 for
|
||||
// permanent). Returns true if the status was a key-specific
|
||||
// failover trigger so callers can retry with the next key.
|
||||
func MarkKeyOnStatus(
|
||||
func (p *Pool) MarkKeyOnStatus(
|
||||
ctx context.Context,
|
||||
key *Key,
|
||||
resp *http.Response,
|
||||
logger slog.Logger,
|
||||
providerName string,
|
||||
) bool {
|
||||
if resp == nil {
|
||||
return false
|
||||
@@ -29,8 +28,11 @@ func MarkKeyOnStatus(
|
||||
cooldown = defaultCooldown
|
||||
}
|
||||
if key.MarkTemporary(cooldown) {
|
||||
if p.metrics != nil {
|
||||
p.metrics.KeyPoolStateTransitions.WithLabelValues(p.providerName, reasonRateLimited).Inc()
|
||||
}
|
||||
logger.Info(ctx, "key marked temporary",
|
||||
slog.F("provider", providerName),
|
||||
slog.F("provider", p.providerName),
|
||||
slog.F("api_key_hint", key.Hint()),
|
||||
slog.F("status", statusCode),
|
||||
slog.F("cooldown", cooldown))
|
||||
@@ -38,15 +40,22 @@ func MarkKeyOnStatus(
|
||||
return true
|
||||
case http.StatusUnauthorized, http.StatusForbidden:
|
||||
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",
|
||||
slog.F("provider", providerName),
|
||||
slog.F("provider", p.providerName),
|
||||
slog.F("api_key_hint", key.Hint()),
|
||||
slog.F("status", statusCode))
|
||||
}
|
||||
return true
|
||||
default:
|
||||
logger.Debug(ctx, "status is not a key failover trigger",
|
||||
slog.F("provider", providerName),
|
||||
slog.F("provider", p.providerName),
|
||||
slog.F("status", statusCode))
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -6,11 +6,14 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/aibridge/keypool"
|
||||
"github.com/coder/coder/v2/aibridge/metrics"
|
||||
codertestutil "github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
@@ -24,6 +27,9 @@ func TestMarkKeyOnStatus(t *testing.T) {
|
||||
expectedReturn bool
|
||||
expectedState keypool.KeyState
|
||||
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).
|
||||
@@ -33,6 +39,7 @@ func TestMarkKeyOnStatus(t *testing.T) {
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStateTemporary,
|
||||
expectedCooldown: 5 * time.Second,
|
||||
expectedReason: "rate_limited",
|
||||
},
|
||||
{
|
||||
// 429 with retry-after-ms header (milliseconds).
|
||||
@@ -42,6 +49,7 @@ func TestMarkKeyOnStatus(t *testing.T) {
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStateTemporary,
|
||||
expectedCooldown: 1500 * time.Millisecond,
|
||||
expectedReason: "rate_limited",
|
||||
},
|
||||
{
|
||||
// 429 without headers falls back to default cooldown.
|
||||
@@ -50,18 +58,21 @@ func TestMarkKeyOnStatus(t *testing.T) {
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStateTemporary,
|
||||
expectedCooldown: 60 * time.Second,
|
||||
expectedReason: "rate_limited",
|
||||
},
|
||||
{
|
||||
name: "401_marks_permanent",
|
||||
statusCode: http.StatusUnauthorized,
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStatePermanent,
|
||||
expectedReason: "unauthorized",
|
||||
},
|
||||
{
|
||||
name: "403_marks_permanent",
|
||||
statusCode: http.StatusForbidden,
|
||||
expectedReturn: true,
|
||||
expectedState: keypool.KeyStatePermanent,
|
||||
expectedReason: "forbidden",
|
||||
},
|
||||
{
|
||||
name: "200_does_not_mark",
|
||||
@@ -85,11 +96,15 @@ func TestMarkKeyOnStatus(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
const providerName = "test-provider"
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
clk := quartz.NewMock(t)
|
||||
pool, err := keypool.New([]string{"key-0"}, clk)
|
||||
reg := prometheus.NewRegistry()
|
||||
m := metrics.NewMetrics(reg)
|
||||
pool, err := keypool.New(providerName, []string{"key-0"}, clk, m)
|
||||
require.NoError(t, err)
|
||||
key, keyPoolErr := pool.Walker().Next()
|
||||
require.Nil(t, keyPoolErr)
|
||||
@@ -102,19 +117,30 @@ func TestMarkKeyOnStatus(t *testing.T) {
|
||||
resp.Header.Set(k, v)
|
||||
}
|
||||
|
||||
got := keypool.MarkKeyOnStatus(
|
||||
got := pool.MarkKeyOnStatus(
|
||||
context.Background(),
|
||||
key,
|
||||
resp,
|
||||
// 401 and 403 cases legitimately log at error
|
||||
// level when marking a key permanent.
|
||||
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
|
||||
"test",
|
||||
)
|
||||
|
||||
assert.Equal(t, tc.expectedReturn, got)
|
||||
assert.Equal(t, tc.expectedState, key.State())
|
||||
|
||||
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:
|
||||
// advancing by exactly that amount returns the key
|
||||
// to valid.
|
||||
|
||||
+67
-13
@@ -7,6 +7,7 @@ import (
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/aibridge/metrics"
|
||||
"github.com/coder/coder/v2/aibridge/utils"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
@@ -53,23 +54,35 @@ func (e *Error) Error() string {
|
||||
}
|
||||
|
||||
// KeyState represents the current state of a key in the pool.
|
||||
type KeyState int
|
||||
type KeyState string
|
||||
|
||||
const (
|
||||
// KeyStateValid means the key is available for use.
|
||||
KeyStateValid KeyState = iota
|
||||
KeyStateValid KeyState = "valid"
|
||||
// KeyStateTemporary means the key is temporarily unavailable
|
||||
// (e.g. rate-limited) and will recover after a cooldown.
|
||||
KeyStateTemporary
|
||||
KeyStateTemporary KeyState = "temporary"
|
||||
// KeyStatePermanent means the key is permanently unavailable
|
||||
// (e.g. revoked or unauthorized) until process restart.
|
||||
KeyStatePermanent
|
||||
KeyStatePermanent KeyState = "permanent"
|
||||
)
|
||||
|
||||
// defaultCooldown is applied when a key is marked temporary
|
||||
// with a zero or negative cooldown duration.
|
||||
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.
|
||||
type Key struct {
|
||||
value string
|
||||
@@ -83,18 +96,33 @@ type Key struct {
|
||||
// Pool manages a set of keys with state tracking and
|
||||
// cooldown expiry. It is safe for concurrent use.
|
||||
type Pool struct {
|
||||
keys []Key
|
||||
keys []Key
|
||||
metrics *metrics.Metrics
|
||||
providerName string
|
||||
}
|
||||
|
||||
// New creates a pool from the given keys. All keys start in
|
||||
// the valid state. Returns ErrNoKeys if keys is empty and
|
||||
// ErrDuplicateKey if any key appears more than once.
|
||||
func New(keys []string, clk quartz.Clock) (*Pool, error) {
|
||||
// RecordAttempts records the total number of keys tried across an
|
||||
// interception. Each upstream request uses its own walker, so the
|
||||
// total sums the attempts across those per-request walkers. Call it
|
||||
// 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 {
|
||||
return nil, ErrNoKeys
|
||||
}
|
||||
pool := &Pool{
|
||||
keys: make([]Key, len(keys)),
|
||||
keys: make([]Key, len(keys)),
|
||||
metrics: m,
|
||||
providerName: providerName,
|
||||
}
|
||||
|
||||
seen := make(map[string]struct{}, len(keys))
|
||||
@@ -231,6 +259,20 @@ func (p *Pool) keyPoolError() *Error {
|
||||
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
|
||||
// original order, used by tests and other diagnostic callers. Use
|
||||
// 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
|
||||
// through keys without interfering with other requests.
|
||||
type Walker struct {
|
||||
pool *Pool
|
||||
pos int // Next index to consider.
|
||||
pool *Pool
|
||||
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
|
||||
@@ -270,9 +313,20 @@ func (w *Walker) Next() (*Key, *Error) {
|
||||
}
|
||||
// Key is available.
|
||||
w.pos = i + 1
|
||||
w.attempts++
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// 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"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"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"
|
||||
)
|
||||
|
||||
@@ -31,7 +34,7 @@ func TestNewKeyPool(t *testing.T) {
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
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 {
|
||||
require.ErrorIs(t, err, tc.expectedErr)
|
||||
return
|
||||
@@ -125,7 +128,7 @@ func TestState(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
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)
|
||||
|
||||
key := tc.setup(t, pool, clk)
|
||||
@@ -204,7 +207,7 @@ func TestMarkTemporary(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
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)
|
||||
|
||||
key := tc.setup(t, pool, clk)
|
||||
@@ -267,7 +270,7 @@ func TestMarkPermanent(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
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)
|
||||
|
||||
key := tc.setup(t, pool)
|
||||
@@ -498,11 +501,15 @@ func TestWalkerNext(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
const providerName = "test-provider"
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
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)
|
||||
|
||||
tc.setup(t, pool)
|
||||
@@ -522,6 +529,26 @@ func TestWalkerNext(t *testing.T) {
|
||||
// After all expected keys, the walker should be exhausted.
|
||||
_, keyPoolErr := walker.Next()
|
||||
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()
|
||||
|
||||
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)
|
||||
key, keyPoolErr := pool.Walker().Next()
|
||||
require.Nil(t, keyPoolErr)
|
||||
@@ -613,7 +640,7 @@ func TestWalkerIndependence(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
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)
|
||||
|
||||
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))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user