summaryrefslogtreecommitdiff
path: root/internal/providerhealth/redis_history.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/providerhealth/redis_history.go')
-rw-r--r--internal/providerhealth/redis_history.go347
1 files changed, 347 insertions, 0 deletions
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 ""
+ }
+}