summaryrefslogtreecommitdiff
path: root/internal/providerhealth
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 15:58:57 +1200
committerChia <Chia@93.nz>2026-08-06 15:58:57 +1200
commit3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch)
tree517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/providerhealth
parent41e322c53d7b4b796eb377d0df9c29ecd10ba431 (diff)
feat: complete commercial developer workflowspublish-commercial-control-plane
Add tenant-safe usage observability, prepaid billing controls, API key lifecycle management, Embeddings metering, configurable billing alerts, and resilient provider health propagation. Harden Stripe failure handling, migrations, readiness, and the authenticated control-plane UI with end-to-end verification evidence.
Diffstat (limited to 'internal/providerhealth')
-rw-r--r--internal/providerhealth/prober.go171
-rw-r--r--internal/providerhealth/prober_test.go91
-rw-r--r--internal/providerhealth/redis_history.go347
-rw-r--r--internal/providerhealth/tracker.go162
-rw-r--r--internal/providerhealth/tracker_test.go152
5 files changed, 917 insertions, 6 deletions
diff --git a/internal/providerhealth/prober.go b/internal/providerhealth/prober.go
new file mode 100644
index 0000000..820fd04
--- /dev/null
+++ b/internal/providerhealth/prober.go
@@ -0,0 +1,171 @@
+package providerhealth
+
+import (
+ "context"
+ "encoding/json"
+ "io"
+ "log/slog"
+ "net/http"
+ "strings"
+ "time"
+
+ "aigw/internal/catalog"
+ "aigw/internal/domain"
+)
+
+type ProbeMetrics interface {
+ ProviderProbe(success bool)
+}
+
+type ProbeOptions struct {
+ Enabled bool
+ Interval time.Duration
+ Timeout time.Duration
+ Catalog *catalog.Catalog
+ Tracker *Tracker
+ Metrics ProbeMetrics
+ Logger *slog.Logger
+ Client *http.Client
+}
+
+type Prober struct {
+ enabled bool
+ interval time.Duration
+ timeout time.Duration
+ catalog *catalog.Catalog
+ tracker *Tracker
+ metrics ProbeMetrics
+ logger *slog.Logger
+ client *http.Client
+}
+
+type probeTarget struct {
+ provider domain.Provider
+ keys []RouteKey
+}
+
+func NewProber(options ProbeOptions) *Prober {
+ if options.Interval <= 0 {
+ options.Interval = 30 * time.Second
+ }
+ if options.Timeout <= 0 {
+ options.Timeout = 5 * time.Second
+ }
+ if options.Logger == nil {
+ options.Logger = slog.Default()
+ }
+ if options.Client == nil {
+ transport := http.DefaultTransport.(*http.Transport).Clone()
+ transport.Proxy = http.ProxyFromEnvironment
+ options.Client = &http.Client{Transport: transport, Timeout: options.Timeout}
+ }
+ return &Prober{enabled: options.Enabled, interval: options.Interval, timeout: options.Timeout,
+ catalog: options.Catalog, tracker: options.Tracker, metrics: options.Metrics, logger: options.Logger, client: options.Client}
+}
+
+func (p *Prober) Run(ctx context.Context) {
+ if p == nil || !p.enabled || p.catalog == nil || p.tracker == nil {
+ return
+ }
+ p.ProbeOnce(ctx)
+ ticker := time.NewTicker(p.interval)
+ defer ticker.Stop()
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ p.ProbeOnce(ctx)
+ }
+ }
+}
+
+func (p *Prober) ProbeOnce(ctx context.Context) {
+ for _, target := range p.targets() {
+ if ctx.Err() != nil {
+ return
+ }
+ p.probe(ctx, target)
+ }
+}
+
+func (p *Prober) targets() []probeTarget {
+ providerIndexes := make(map[string]int)
+ seen := make(map[RouteKey]struct{})
+ result := make([]probeTarget, 0)
+ for _, model := range p.catalog.AllModels() {
+ for _, route := range model.Routes {
+ key := RouteKey{ModelID: model.ID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}
+ if key.ProviderID == "" {
+ continue
+ }
+ if _, exists := seen[key]; exists {
+ continue
+ }
+ seen[key] = struct{}{}
+ index, exists := providerIndexes[route.Provider.ID]
+ if !exists {
+ index = len(result)
+ providerIndexes[route.Provider.ID] = index
+ result = append(result, probeTarget{provider: route.Provider})
+ }
+ result[index].keys = append(result[index].keys, key)
+ }
+ }
+ return result
+}
+
+func (p *Prober) probe(parent context.Context, target probeTarget) {
+ ctx, cancel := context.WithTimeout(parent, p.timeout)
+ defer cancel()
+ request, err := http.NewRequestWithContext(ctx, http.MethodGet, providerModelsURL(target.provider), nil)
+ if err != nil {
+ p.observe(target, 0, 0, false)
+ return
+ }
+ request.Header.Set("Accept", "application/json")
+ request.Header.Set("User-Agent", "aigw-health/0.1")
+ if target.provider.Protocol == domain.ProtocolAnthropic {
+ request.Header.Set("x-api-key", target.provider.APIKey)
+ request.Header.Set("anthropic-version", "2023-06-01")
+ } else {
+ request.Header.Set("Authorization", "Bearer "+target.provider.APIKey)
+ }
+ started := time.Now()
+ response, err := p.client.Do(request)
+ latency := time.Since(started)
+ if err != nil {
+ p.observe(target, 0, latency, false)
+ p.logger.Warn("provider_probe_failed", "provider_id", target.provider.ID, "error", err)
+ return
+ }
+ body, readErr := io.ReadAll(io.LimitReader(response.Body, 1<<20))
+ _ = response.Body.Close()
+ success := response.StatusCode >= 200 && response.StatusCode < 300 && readErr == nil &&
+ strings.HasPrefix(strings.ToLower(response.Header.Get("Content-Type")), "application/json") && validModelsEnvelope(body)
+ p.observe(target, response.StatusCode, latency, success)
+ if !success {
+ p.logger.Warn("provider_probe_failed", "provider_id", target.provider.ID, "status_code", response.StatusCode)
+ }
+}
+
+func (p *Prober) observe(target probeTarget, statusCode int, latency time.Duration, success bool) {
+ for _, key := range target.keys {
+ p.tracker.Observe(key, Observation{StatusCode: statusCode, Latency: latency, Failed: !success, Active: true})
+ }
+ if p.metrics != nil {
+ p.metrics.ProviderProbe(success)
+ }
+}
+
+func validModelsEnvelope(body []byte) bool {
+ var envelope struct {
+ Data []json.RawMessage `json:"data"`
+ }
+ return json.Unmarshal(body, &envelope) == nil && envelope.Data != nil
+}
+
+func providerModelsURL(provider domain.Provider) string {
+ baseURL := strings.TrimRight(provider.BaseURL, "/")
+ return baseURL + "/models"
+}
diff --git a/internal/providerhealth/prober_test.go b/internal/providerhealth/prober_test.go
new file mode 100644
index 0000000..69b6e80
--- /dev/null
+++ b/internal/providerhealth/prober_test.go
@@ -0,0 +1,91 @@
+package providerhealth
+
+import (
+ "context"
+ "io"
+ "log/slog"
+ "net/http"
+ "net/http/httptest"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ "aigw/internal/catalog"
+ "aigw/internal/domain"
+)
+
+type probeMetricCounter struct {
+ total atomic.Int64
+ failed atomic.Int64
+}
+
+func (m *probeMetricCounter) ProviderProbe(success bool) {
+ m.total.Add(1)
+ if !success {
+ m.failed.Add(1)
+ }
+}
+
+func TestProberAuthenticatesDeduplicatesAndOpensRouteCircuits(t *testing.T) {
+ var calls atomic.Int64
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ calls.Add(1)
+ if r.URL.Path != "/v1/models" || r.Header.Get("Authorization") != "Bearer secret" {
+ t.Errorf("unexpected probe path=%s auth=%q", r.URL.Path, r.Header.Get("Authorization"))
+ }
+ http.Error(w, "unavailable", http.StatusServiceUnavailable)
+ }))
+ defer server.Close()
+
+ provider := domain.Provider{ID: "provider", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: server.URL + "/v1", APIKey: "secret"}
+ model := domain.Model{ID: "public/model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "one"}, {Provider: provider, UpstreamModel: "two"}}}
+ tracker := New(Options{FailureThreshold: 1, OpenDuration: time.Minute})
+ metrics := &probeMetricCounter{}
+ prober := NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker,
+ Metrics: metrics, Timeout: time.Second, Logger: slog.New(slog.NewTextHandler(io.Discard, nil))})
+ prober.ProbeOnce(context.Background())
+
+ key := RouteKey{ModelID: model.ID, ProviderID: provider.ID, WireAPI: provider.WireAPI}
+ if calls.Load() != 1 || metrics.total.Load() != 1 || metrics.failed.Load() != 1 || !tracker.CircuitOpen(key) {
+ t.Fatalf("calls=%d total=%d failed=%d open=%v", calls.Load(), metrics.total.Load(), metrics.failed.Load(), tracker.CircuitOpen(key))
+ }
+ status := tracker.Snapshot()[0]
+ if status.ActiveProbes != 1 || status.LastProbeAt == nil || status.LastStatusCode != http.StatusServiceUnavailable {
+ t.Fatalf("unexpected probe status: %+v", status)
+ }
+}
+
+func TestProberRejectsHTMLSuccessResponse(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "text/html")
+ _, _ = io.WriteString(w, "<html>provider console</html>")
+ }))
+ defer server.Close()
+ provider := domain.Provider{ID: "provider", Protocol: domain.ProtocolOpenAI, BaseURL: server.URL, APIKey: "secret"}
+ model := domain.Model{ID: "model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "model"}}}
+ tracker := New(Options{FailureThreshold: 1})
+ NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, Timeout: time.Second,
+ Logger: slog.New(slog.NewTextHandler(io.Discard, nil))}).ProbeOnce(context.Background())
+ if !tracker.CircuitOpen(RouteKey{ModelID: model.ID, ProviderID: provider.ID, WireAPI: "chat_completions"}) {
+ t.Fatal("HTML success response must not be treated as a healthy API probe")
+ }
+}
+
+func TestProberUsesAnthropicAuthentication(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Header.Get("x-api-key") != "anthropic-secret" || r.Header.Get("anthropic-version") == "" {
+ t.Errorf("unexpected Anthropic headers")
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = io.WriteString(w, `{"data":[]}`)
+ }))
+ defer server.Close()
+ provider := domain.Provider{ID: "anthropic", Protocol: domain.ProtocolAnthropic, WireAPI: "messages", BaseURL: server.URL + "/v1", APIKey: "anthropic-secret"}
+ model := domain.Model{ID: "anthropic/model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "model"}}}
+ tracker := New(Options{})
+ NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, Timeout: time.Second}).ProbeOnce(context.Background())
+ status := tracker.Snapshot()[0]
+ if status.State != "healthy" || status.ActiveProbes != 1 {
+ t.Fatalf("unexpected status: %+v", status)
+ }
+}
diff --git a/internal/providerhealth/redis_history.go b/internal/providerhealth/redis_history.go
new file mode 100644
index 0000000..322acce
--- /dev/null
+++ b/internal/providerhealth/redis_history.go
@@ -0,0 +1,347 @@
+package providerhealth
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "log/slog"
+ "strings"
+ "sync"
+ "sync/atomic"
+ "time"
+
+ "github.com/redis/go-redis/v9"
+)
+
+const (
+ defaultHistoryStream = "aigw:provider-health:events"
+ defaultHistoryQueueSize = 4096
+ defaultHistoryMaxEvents = 20_000
+ defaultHistoryTTL = 15 * time.Minute
+ historyPublishBatchSize = 64
+ historyCommandTimeout = 500 * time.Millisecond
+ historyRetryCooldown = 5 * time.Second
+)
+
+type HistoryOptions struct {
+ Enabled bool
+ RedisURL string
+ Stream string
+ Instance string
+ QueueSize int
+ MaxEvents int64
+ TTL time.Duration
+ Logger *slog.Logger
+ Client *redis.Client
+ Metrics SharedHistoryMetrics
+}
+
+type SharedHistoryMetrics interface {
+ ProviderHealthSharedPublished()
+ ProviderHealthSharedImported()
+ ProviderHealthSharedDropped()
+ ProviderHealthSharedRedisFailure()
+ ProviderHealthSharedConnected(bool)
+}
+
+type RedisHistory struct {
+ redis *redis.Client
+ ownedClient bool
+ stream string
+ instance string
+ queue chan Event
+ maxEvents int64
+ ttl time.Duration
+ logger *slog.Logger
+ metrics SharedHistoryMetrics
+ retryAt atomic.Int64
+ dropped atomic.Uint64
+ closed atomic.Bool
+ failureSeen atomic.Bool
+}
+
+func NewRedisHistory(options HistoryOptions) *RedisHistory {
+ if !options.Enabled || (strings.TrimSpace(options.RedisURL) == "" && options.Client == nil) {
+ return nil
+ }
+ if options.Logger == nil {
+ options.Logger = slog.Default()
+ }
+ if options.Stream == "" {
+ options.Stream = defaultHistoryStream
+ }
+ if options.QueueSize <= 0 {
+ options.QueueSize = defaultHistoryQueueSize
+ }
+ if options.MaxEvents <= 0 {
+ options.MaxEvents = defaultHistoryMaxEvents
+ }
+ if options.TTL <= 0 {
+ options.TTL = defaultHistoryTTL
+ }
+ client := options.Client
+ ownedClient := false
+ if client == nil {
+ redisOptions, err := redis.ParseURL(options.RedisURL)
+ if err != nil {
+ options.Logger.Warn("provider_health_redis_config_invalid", "error", err, "fallback", "local")
+ return nil
+ }
+ redisOptions.MaxRetries = -1
+ redisOptions.DialerRetries = 1
+ redisOptions.DialTimeout = historyCommandTimeout
+ redisOptions.ReadTimeout = historyCommandTimeout
+ redisOptions.WriteTimeout = historyCommandTimeout
+ redisOptions.PoolTimeout = historyCommandTimeout
+ client = redis.NewClient(redisOptions)
+ ownedClient = true
+ }
+ return &RedisHistory{redis: client, ownedClient: ownedClient, stream: options.Stream, instance: options.Instance,
+ queue: make(chan Event, options.QueueSize), maxEvents: options.MaxEvents, ttl: options.TTL,
+ logger: options.Logger, metrics: options.Metrics}
+}
+
+func (h *RedisHistory) Enqueue(event Event) {
+ if h == nil || h.closed.Load() {
+ return
+ }
+ select {
+ case h.queue <- event:
+ default:
+ if h.recordDropped(1) == 1 {
+ h.logger.Warn("provider_health_share_queue_full", "fallback", "local")
+ }
+ }
+}
+
+func (h *RedisHistory) Run(ctx context.Context, tracker *Tracker) {
+ if h == nil || tracker == nil {
+ return
+ }
+ var workers sync.WaitGroup
+ workers.Add(2)
+ go func() {
+ defer workers.Done()
+ h.runPublisher(ctx)
+ }()
+ go func() {
+ defer workers.Done()
+ h.runReader(ctx, tracker)
+ }()
+ workers.Wait()
+}
+
+func (h *RedisHistory) runPublisher(ctx context.Context) {
+ batch := make([]Event, 0, historyPublishBatchSize)
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case event := <-h.queue:
+ batch = append(batch[:0], event)
+ }
+ for len(batch) < historyPublishBatchSize {
+ select {
+ case event := <-h.queue:
+ batch = append(batch, event)
+ default:
+ h.publishBatch(ctx, batch)
+ batch = batch[:0]
+ goto nextBatch
+ }
+ }
+ h.publishBatch(ctx, batch)
+ batch = batch[:0]
+ nextBatch:
+ }
+}
+
+func (h *RedisHistory) runReader(ctx context.Context, tracker *Tracker) {
+ lastID := h.importRecent(ctx, tracker)
+ if lastID == "" {
+ lastID = "0-0"
+ }
+ for ctx.Err() == nil {
+ if time.Now().UnixNano() < h.retryAt.Load() {
+ if !waitHistoryRetry(ctx, 100*time.Millisecond) {
+ return
+ }
+ continue
+ }
+ lastID = h.read(ctx, tracker, lastID)
+ }
+}
+
+func (h *RedisHistory) Close() error {
+ if h == nil || !h.closed.CompareAndSwap(false, true) || !h.ownedClient {
+ return nil
+ }
+ return h.redis.Close()
+}
+
+func (h *RedisHistory) Dropped() uint64 {
+ if h == nil {
+ return 0
+ }
+ return h.dropped.Load()
+}
+
+func (h *RedisHistory) publishBatch(parent context.Context, events []Event) {
+ if len(events) == 0 {
+ return
+ }
+ if time.Now().UnixNano() < h.retryAt.Load() {
+ h.recordDropped(uint64(len(events)))
+ return
+ }
+ ctx, cancel := context.WithTimeout(parent, historyCommandTimeout)
+ defer cancel()
+ pipe := h.redis.Pipeline()
+ published := 0
+ for _, event := range events {
+ payload, err := json.Marshal(event)
+ if err != nil {
+ h.recordDropped(1)
+ continue
+ }
+ pipe.XAdd(ctx, &redis.XAddArgs{Stream: h.stream, MaxLen: h.maxEvents, Approx: true,
+ Values: map[string]any{"instance": h.instance, "event": string(payload)}})
+ published++
+ }
+ if published == 0 {
+ pipe.Discard()
+ return
+ }
+ pipe.Expire(ctx, h.stream, h.ttl)
+ if _, err := pipe.Exec(ctx); err != nil {
+ h.recordDropped(uint64(published))
+ h.fail(err)
+ return
+ }
+ h.recovered()
+ for range published {
+ if h.metrics != nil {
+ h.metrics.ProviderHealthSharedPublished()
+ }
+ }
+}
+
+func (h *RedisHistory) importRecent(parent context.Context, tracker *Tracker) string {
+ ctx, cancel := context.WithTimeout(parent, historyCommandTimeout)
+ defer cancel()
+ messages, err := h.redis.XRevRangeN(ctx, h.stream, "+", "-", h.maxEvents).Result()
+ if err != nil {
+ if !errors.Is(err, redis.Nil) {
+ h.fail(err)
+ }
+ return ""
+ }
+ if len(messages) == 0 {
+ h.recovered()
+ return ""
+ }
+ lastID := messages[0].ID
+ for index := len(messages) - 1; index >= 0; index-- {
+ h.applyMessage(tracker, messages[index])
+ }
+ h.recovered()
+ return lastID
+}
+
+func (h *RedisHistory) read(parent context.Context, tracker *Tracker, lastID string) string {
+ if time.Now().UnixNano() < h.retryAt.Load() {
+ return lastID
+ }
+ ctx, cancel := context.WithTimeout(parent, historyCommandTimeout)
+ defer cancel()
+ streams, err := h.redis.XRead(ctx, &redis.XReadArgs{Streams: []string{h.stream, lastID}, Count: 512, Block: 200 * time.Millisecond}).Result()
+ if errors.Is(err, redis.Nil) {
+ h.recovered()
+ return lastID
+ }
+ if err != nil {
+ h.fail(err)
+ return lastID
+ }
+ h.recovered()
+ for _, stream := range streams {
+ for _, message := range stream.Messages {
+ h.applyMessage(tracker, message)
+ lastID = message.ID
+ }
+ }
+ return lastID
+}
+
+func (h *RedisHistory) applyMessage(tracker *Tracker, message redis.XMessage) {
+ if asString(message.Values["instance"]) == h.instance && h.instance != "" {
+ return
+ }
+ payload := asString(message.Values["event"])
+ var event Event
+ if payload == "" || json.Unmarshal([]byte(payload), &event) != nil {
+ return
+ }
+ if event.ObservedAt.Before(time.Now().Add(-h.ttl)) {
+ return
+ }
+ tracker.ApplyShared(event)
+ if h.metrics != nil {
+ h.metrics.ProviderHealthSharedImported()
+ }
+}
+
+func (h *RedisHistory) fail(err error) {
+ h.retryAt.Store(time.Now().Add(historyRetryCooldown).UnixNano())
+ if h.metrics != nil {
+ h.metrics.ProviderHealthSharedConnected(false)
+ }
+ if h.failureSeen.CompareAndSwap(false, true) {
+ if h.metrics != nil {
+ h.metrics.ProviderHealthSharedRedisFailure()
+ }
+ h.logger.Warn("provider_health_redis_unavailable", "error", err, "fallback", "local", "retry_after", historyRetryCooldown)
+ }
+}
+
+func (h *RedisHistory) recovered() {
+ h.retryAt.Store(0)
+ if h.metrics != nil {
+ h.metrics.ProviderHealthSharedConnected(true)
+ }
+ if h.failureSeen.Swap(false) {
+ h.logger.Info("provider_health_redis_recovered")
+ }
+}
+
+func (h *RedisHistory) recordDropped(count uint64) uint64 {
+ total := h.dropped.Add(count)
+ if h.metrics != nil {
+ for range count {
+ h.metrics.ProviderHealthSharedDropped()
+ }
+ }
+ return total
+}
+
+func waitHistoryRetry(ctx context.Context, delay time.Duration) bool {
+ timer := time.NewTimer(delay)
+ defer timer.Stop()
+ select {
+ case <-ctx.Done():
+ return false
+ case <-timer.C:
+ return true
+ }
+}
+
+func asString(value any) string {
+ switch item := value.(type) {
+ case string:
+ return item
+ case []byte:
+ return string(item)
+ default:
+ return ""
+ }
+}
diff --git a/internal/providerhealth/tracker.go b/internal/providerhealth/tracker.go
index a5af2b2..6394def 100644
--- a/internal/providerhealth/tracker.go
+++ b/internal/providerhealth/tracker.go
@@ -8,31 +8,59 @@ import (
const recentWindow = 100
+type EventKind string
+
+const (
+ EventOutcome EventKind = "outcome"
+ EventTTFT EventKind = "ttft"
+)
+
type RouteKey struct {
- ModelID string
- ProviderID string
- WireAPI string
+ ModelID string `json:"model_id"`
+ ProviderID string `json:"provider_id"`
+ WireAPI string `json:"wire_api"`
}
type Observation struct {
StatusCode int
Latency time.Duration
Failed bool
+ Active bool
ObservedAt time.Time
}
+type Event struct {
+ Kind EventKind `json:"kind"`
+ Key RouteKey `json:"key"`
+ StatusCode int `json:"status_code,omitempty"`
+ LatencyMillis int64 `json:"latency_ms,omitempty"`
+ Failed bool `json:"failed,omitempty"`
+ Active bool `json:"active,omitempty"`
+ ObservedAt time.Time `json:"observed_at"`
+}
+
+type EventSink interface {
+ Enqueue(Event)
+}
+
type Status struct {
ModelID string `json:"model_id"`
ProviderID string `json:"provider_id"`
WireAPI string `json:"wire_api"`
State string `json:"state"`
Attempts uint64 `json:"attempts"`
+ ActiveProbes uint64 `json:"active_probes"`
RecentSamples int `json:"recent_samples"`
AvailabilityPercent float64 `json:"availability_percent"`
HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"`
+ TTFTSamples uint64 `json:"ttft_samples"`
+ TTFTEWMA int64 `json:"ttft_ewma_ms"`
+ SharedAttempts uint64 `json:"shared_attempts"`
+ SharedTTFTSamples uint64 `json:"shared_ttft_samples"`
ConsecutiveFailures uint64 `json:"consecutive_failures"`
LastStatusCode int `json:"last_status_code,omitempty"`
LastObservedAt *time.Time `json:"last_observed_at,omitempty"`
+ LastProbeAt *time.Time `json:"last_probe_at,omitempty"`
LastHealthyAt *time.Time `json:"last_healthy_at,omitempty"`
CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"`
}
@@ -41,6 +69,7 @@ type Options struct {
FailureThreshold uint64
OpenDuration time.Duration
Now func() time.Time
+ Sink EventSink
}
type Tracker struct {
@@ -48,23 +77,63 @@ type Tracker struct {
failureThreshold uint64
openDuration time.Duration
now func() time.Time
+ sinkMu sync.RWMutex
+ sink EventSink
}
type routeState struct {
mu sync.RWMutex
attempts uint64
+ activeProbes uint64
consecutiveFailures uint64
lastStatusCode int
lastObservedAt time.Time
+ lastProbeAt time.Time
lastHealthyAt time.Time
openUntil time.Time
headerLatencyEWMA float64
+ ttftSamples uint64
+ ttftEWMA float64
+ sharedAttempts uint64
+ sharedTTFTSamples uint64
recent [recentWindow]bool
recentCount int
recentPosition int
recentHealthy int
}
+// ObserveTTFT records user-visible response latency without counting a second
+// request outcome. Forwarder observations already update availability and the
+// circuit when response headers arrive.
+func (t *Tracker) ObserveTTFT(key RouteKey, latency time.Duration) {
+ if t == nil || key.ProviderID == "" || latency <= 0 {
+ return
+ }
+ observedAt := t.now()
+ t.observeTTFT(key, latency, false)
+ t.publish(Event{Kind: EventTTFT, Key: key, LatencyMillis: durationMillis(latency), ObservedAt: observedAt})
+}
+
+func (t *Tracker) observeTTFT(key RouteKey, latency time.Duration, shared bool) {
+ value, _ := t.states.LoadOrStore(key, &routeState{})
+ state := value.(*routeState)
+ state.mu.Lock()
+ defer state.mu.Unlock()
+ valueMS := float64(durationMillis(latency))
+ if valueMS < 1 {
+ valueMS = 1
+ }
+ state.ttftSamples++
+ if shared {
+ state.sharedTTFTSamples++
+ }
+ if state.ttftEWMA == 0 {
+ state.ttftEWMA = valueMS
+ } else {
+ state.ttftEWMA = state.ttftEWMA*0.8 + valueMS*0.2
+ }
+}
+
func New(options Options) *Tracker {
if options.FailureThreshold == 0 {
options.FailureThreshold = 3
@@ -75,7 +144,7 @@ func New(options Options) *Tracker {
if options.Now == nil {
options.Now = time.Now
}
- return &Tracker{failureThreshold: options.FailureThreshold, openDuration: options.OpenDuration, now: options.Now}
+ return &Tracker{failureThreshold: options.FailureThreshold, openDuration: options.OpenDuration, now: options.Now, sink: options.Sink}
}
func (t *Tracker) Observe(key RouteKey, observation Observation) {
@@ -85,12 +154,26 @@ func (t *Tracker) Observe(key RouteKey, observation Observation) {
if observation.ObservedAt.IsZero() {
observation.ObservedAt = t.now()
}
+ t.observe(key, observation, false)
+ t.publish(Event{Kind: EventOutcome, Key: key, StatusCode: observation.StatusCode,
+ LatencyMillis: durationMillis(observation.Latency), Failed: observation.Failed,
+ Active: observation.Active, ObservedAt: observation.ObservedAt})
+}
+
+func (t *Tracker) observe(key RouteKey, observation Observation, shared bool) {
value, _ := t.states.LoadOrStore(key, &routeState{})
state := value.(*routeState)
state.mu.Lock()
defer state.mu.Unlock()
state.attempts++
+ if shared {
+ state.sharedAttempts++
+ }
+ if observation.Active {
+ state.activeProbes++
+ state.lastProbeAt = observation.ObservedAt
+ }
state.lastStatusCode = observation.StatusCode
state.lastObservedAt = observation.ObservedAt
if observation.Latency > 0 {
@@ -117,6 +200,53 @@ func (t *Tracker) Observe(key RouteKey, observation Observation) {
state.lastHealthyAt = observation.ObservedAt
}
+func (t *Tracker) SetSink(sink EventSink) {
+ if t == nil {
+ return
+ }
+ t.sinkMu.Lock()
+ t.sink = sink
+ t.sinkMu.Unlock()
+}
+
+func (t *Tracker) ApplyShared(event Event) {
+ if t == nil || event.Key.ProviderID == "" {
+ return
+ }
+ if event.ObservedAt.IsZero() {
+ event.ObservedAt = t.now()
+ }
+ switch event.Kind {
+ case EventOutcome:
+ t.observe(event.Key, Observation{StatusCode: event.StatusCode, Latency: time.Duration(event.LatencyMillis) * time.Millisecond,
+ Failed: event.Failed, Active: event.Active, ObservedAt: event.ObservedAt}, true)
+ case EventTTFT:
+ if event.LatencyMillis > 0 {
+ t.observeTTFT(event.Key, time.Duration(event.LatencyMillis)*time.Millisecond, true)
+ }
+ }
+}
+
+func (t *Tracker) publish(event Event) {
+ t.sinkMu.RLock()
+ sink := t.sink
+ t.sinkMu.RUnlock()
+ if sink != nil {
+ sink.Enqueue(event)
+ }
+}
+
+func durationMillis(value time.Duration) int64 {
+ if value <= 0 {
+ return 0
+ }
+ milliseconds := value.Milliseconds()
+ if milliseconds < 1 {
+ return 1
+ }
+ return milliseconds
+}
+
func (s *routeState) addRecent(healthy bool) {
if s.recentCount == recentWindow {
if s.recent[s.recentPosition] {
@@ -146,6 +276,21 @@ func (t *Tracker) CircuitOpen(key RouteKey) bool {
return state.openUntil.After(t.now())
}
+// StatusFor returns one immutable route-health snapshot for routing decisions.
+func (t *Tracker) StatusFor(key RouteKey) (Status, bool) {
+ if t == nil {
+ return Status{}, false
+ }
+ value, ok := t.states.Load(key)
+ if !ok {
+ return Status{}, false
+ }
+ state := value.(*routeState)
+ state.mu.RLock()
+ defer state.mu.RUnlock()
+ return statusFromState(key, state, t.now()), true
+}
+
func (t *Tracker) Snapshot() []Status {
if t == nil {
return []Status{}
@@ -172,8 +317,9 @@ func (t *Tracker) Snapshot() []Status {
func statusFromState(key RouteKey, state *routeState, now time.Time) Status {
item := Status{ModelID: key.ModelID, ProviderID: key.ProviderID, WireAPI: key.WireAPI,
- Attempts: state.attempts, RecentSamples: state.recentCount, HeaderLatencyEWMA: int64(state.headerLatencyEWMA + 0.5),
- ConsecutiveFailures: state.consecutiveFailures, LastStatusCode: state.lastStatusCode}
+ Attempts: state.attempts, ActiveProbes: state.activeProbes, RecentSamples: state.recentCount, HeaderLatencyEWMA: int64(state.headerLatencyEWMA + 0.5),
+ TTFTSamples: state.ttftSamples, TTFTEWMA: int64(state.ttftEWMA + 0.5), SharedAttempts: state.sharedAttempts,
+ SharedTTFTSamples: state.sharedTTFTSamples, ConsecutiveFailures: state.consecutiveFailures, LastStatusCode: state.lastStatusCode}
if state.recentCount > 0 {
item.AvailabilityPercent = float64(state.recentHealthy) / float64(state.recentCount) * 100
}
@@ -181,6 +327,10 @@ func statusFromState(key RouteKey, state *routeState, now time.Time) Status {
value := state.lastObservedAt
item.LastObservedAt = &value
}
+ if !state.lastProbeAt.IsZero() {
+ value := state.lastProbeAt
+ item.LastProbeAt = &value
+ }
if !state.lastHealthyAt.IsZero() {
value := state.lastHealthyAt
item.LastHealthyAt = &value
diff --git a/internal/providerhealth/tracker_test.go b/internal/providerhealth/tracker_test.go
index 73bd6d5..19752df 100644
--- a/internal/providerhealth/tracker_test.go
+++ b/internal/providerhealth/tracker_test.go
@@ -1,11 +1,23 @@
package providerhealth
import (
+ "context"
+ "fmt"
+ "log/slog"
+ "os"
"sync"
"testing"
"time"
+
+ "github.com/redis/go-redis/v9"
)
+type captureSink struct {
+ events []Event
+}
+
+func (s *captureSink) Enqueue(event Event) { s.events = append(s.events, event) }
+
func TestTrackerOpensAndRecoversCircuit(t *testing.T) {
now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC)
tracker := New(Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }})
@@ -31,6 +43,146 @@ func TestTrackerOpensAndRecoversCircuit(t *testing.T) {
}
}
+func TestTrackerRecordsActiveProbeMetadata(t *testing.T) {
+ now := time.Date(2026, time.August, 6, 1, 0, 0, 0, time.UTC)
+ tracker := New(Options{Now: func() time.Time { return now }})
+ key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "embeddings"}
+ tracker.Observe(key, Observation{StatusCode: 200, Latency: 20 * time.Millisecond, Active: true})
+ status := tracker.Snapshot()[0]
+ if status.ActiveProbes != 1 || status.LastProbeAt == nil || !status.LastProbeAt.Equal(now) || status.State != "healthy" {
+ t.Fatalf("unexpected active probe status: %+v", status)
+ }
+}
+
+func TestTrackerRecordsTTFTWithoutDoubleCountingAvailability(t *testing.T) {
+ tracker := New(Options{})
+ key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"}
+ tracker.Observe(key, Observation{StatusCode: 200, Latency: 20 * time.Millisecond})
+ tracker.ObserveTTFT(key, 100*time.Millisecond)
+ tracker.ObserveTTFT(key, 200*time.Millisecond)
+ status, found := tracker.StatusFor(key)
+ if !found || status.Attempts != 1 || status.RecentSamples != 1 || status.TTFTSamples != 2 || status.TTFTEWMA != 120 {
+ t.Fatalf("unexpected TTFT status: %+v", status)
+ }
+}
+
+func TestTrackerSharesLocalEventsWithoutRepublishingImportedEvents(t *testing.T) {
+ sink := &captureSink{}
+ tracker := New(Options{Sink: sink})
+ key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"}
+ tracker.Observe(key, Observation{StatusCode: 200, Latency: 20 * time.Millisecond})
+ tracker.ObserveTTFT(key, 100*time.Millisecond)
+ if len(sink.events) != 2 || sink.events[0].Kind != EventOutcome || sink.events[1].Kind != EventTTFT {
+ t.Fatalf("unexpected published events %+v", sink.events)
+ }
+ tracker.ApplyShared(Event{Kind: EventOutcome, Key: key, StatusCode: 503, Failed: true, ObservedAt: time.Now()})
+ tracker.ApplyShared(Event{Kind: EventTTFT, Key: key, LatencyMillis: 250, ObservedAt: time.Now()})
+ if len(sink.events) != 2 {
+ t.Fatalf("imported observations were republished: %d events", len(sink.events))
+ }
+ status, found := tracker.StatusFor(key)
+ if !found || status.Attempts != 2 || status.SharedAttempts != 1 || status.TTFTSamples != 2 || status.SharedTTFTSamples != 1 {
+ t.Fatalf("unexpected shared status %+v", status)
+ }
+}
+
+func TestSharedFailuresOpenLocalCircuit(t *testing.T) {
+ now := time.Date(2026, time.August, 6, 2, 0, 0, 0, time.UTC)
+ tracker := New(Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }})
+ key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"}
+ for range 3 {
+ tracker.ApplyShared(Event{Kind: EventOutcome, Key: key, StatusCode: 503, Failed: true, ObservedAt: now})
+ }
+ status, found := tracker.StatusFor(key)
+ if !found || !tracker.CircuitOpen(key) || status.State != "open" || status.Attempts != 3 || status.SharedAttempts != 3 {
+ t.Fatalf("shared failures did not open the circuit: %+v", status)
+ }
+}
+
+func TestRedisHistoryDropsWithoutBlockingWhenQueueIsFull(t *testing.T) {
+ client := redis.NewClient(&redis.Options{Addr: "127.0.0.1:1"})
+ defer client.Close()
+ history := NewRedisHistory(HistoryOptions{Enabled: true, Client: client, QueueSize: 1,
+ Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))})
+ event := Event{Kind: EventOutcome, Key: RouteKey{ModelID: "model", ProviderID: "provider"}, ObservedAt: time.Now()}
+ history.Enqueue(event)
+ history.Enqueue(event)
+ if history.Dropped() != 1 {
+ t.Fatalf("dropped = %d, want 1", history.Dropped())
+ }
+}
+
+func TestRedisHistorySharesAndReplaysObservations(t *testing.T) {
+ redisURL := os.Getenv("AIGW_TEST_REDIS_URL")
+ if redisURL == "" {
+ t.Skip("AIGW_TEST_REDIS_URL is not set")
+ }
+ stream := fmt.Sprintf("aigw:test:provider-health:%d", time.Now().UnixNano())
+ options, err := redis.ParseURL(redisURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ cleanupClient := redis.NewClient(options)
+ t.Cleanup(func() {
+ _, _ = cleanupClient.Del(context.Background(), stream).Result()
+ _ = cleanupClient.Close()
+ })
+
+ ctx, cancel := context.WithCancel(context.Background())
+ defer cancel()
+ key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"}
+ first := New(Options{})
+ second := New(Options{})
+ firstHistory := NewRedisHistory(HistoryOptions{Enabled: true, RedisURL: redisURL, Stream: stream, Instance: "first",
+ TTL: time.Minute, MaxEvents: 100, Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))})
+ secondHistory := NewRedisHistory(HistoryOptions{Enabled: true, RedisURL: redisURL, Stream: stream, Instance: "second",
+ TTL: time.Minute, MaxEvents: 100, Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))})
+ if firstHistory == nil || secondHistory == nil {
+ t.Fatal("Redis histories were not configured")
+ }
+ t.Cleanup(func() { _ = firstHistory.Close() })
+ t.Cleanup(func() { _ = secondHistory.Close() })
+ first.SetSink(firstHistory)
+ go firstHistory.Run(ctx, first)
+ go secondHistory.Run(ctx, second)
+
+ first.Observe(key, Observation{StatusCode: 200, Latency: 10 * time.Millisecond})
+ first.ObserveTTFT(key, 40*time.Millisecond)
+ waitForSharedStatus(t, second, key, 1, 1)
+ status, _ := first.StatusFor(key)
+ if status.SharedAttempts != 0 || status.SharedTTFTSamples != 0 {
+ t.Fatalf("publisher imported its own events: %+v", status)
+ }
+
+ restarted := New(Options{})
+ restartedHistory := NewRedisHistory(HistoryOptions{Enabled: true, RedisURL: redisURL, Stream: stream, Instance: "restarted",
+ TTL: time.Minute, MaxEvents: 100, Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))})
+ if restartedHistory == nil {
+ t.Fatal("restart history was not configured")
+ }
+ t.Cleanup(func() { _ = restartedHistory.Close() })
+ go restartedHistory.Run(ctx, restarted)
+ waitForSharedStatus(t, restarted, key, 1, 1)
+}
+
+func waitForSharedStatus(t *testing.T, tracker *Tracker, key RouteKey, attempts, ttft uint64) {
+ t.Helper()
+ deadline := time.Now().Add(5 * time.Second)
+ for time.Now().Before(deadline) {
+ status, found := tracker.StatusFor(key)
+ if found && status.SharedAttempts >= attempts && status.SharedTTFTSamples >= ttft {
+ return
+ }
+ time.Sleep(25 * time.Millisecond)
+ }
+ status, _ := tracker.StatusFor(key)
+ t.Fatalf("shared status did not converge: %+v", status)
+}
+
+type ioDiscard struct{}
+
+func (ioDiscard) Write(data []byte) (int, error) { return len(data), nil }
+
func TestTrackerConcurrentObservations(t *testing.T) {
tracker := New(Options{FailureThreshold: 1000})
key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "chat_completions"}