diff options
| author | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
| commit | 3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch) | |
| tree | 517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/providerhealth | |
| parent | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (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.go | 171 | ||||
| -rw-r--r-- | internal/providerhealth/prober_test.go | 91 | ||||
| -rw-r--r-- | internal/providerhealth/redis_history.go | 347 | ||||
| -rw-r--r-- | internal/providerhealth/tracker.go | 162 | ||||
| -rw-r--r-- | internal/providerhealth/tracker_test.go | 152 |
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"} |
