summaryrefslogtreecommitdiff
path: root/internal/controlplane/manager.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/controlplane/manager.go')
-rw-r--r--internal/controlplane/manager.go213
1 files changed, 213 insertions, 0 deletions
diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go
new file mode 100644
index 0000000..42f5efe
--- /dev/null
+++ b/internal/controlplane/manager.go
@@ -0,0 +1,213 @@
+package controlplane
+
+import (
+ "context"
+ "encoding/json"
+ "log/slog"
+ "sync"
+ "sync/atomic"
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/catalog"
+)
+
+const broadcastQueueSize = 128
+
+type managerStore interface {
+ LoadSnapshot(context.Context) (Snapshot, error)
+ DatabaseGeneration(context.Context) (int64, error)
+ PublishChange(context.Context, ChangeEvent) error
+ Subscribe(context.Context) (<-chan ChangeMessage, func() error, error)
+ RedisEnabled() bool
+}
+
+type Manager struct {
+ store managerStore
+ catalog *catalog.Catalog
+ authenticator *auth.StaticAuthenticator
+ logger *slog.Logger
+ pollInterval time.Duration
+ generation atomic.Int64
+ redisConnected atomic.Bool
+ reloadMu sync.Mutex
+ broadcasts chan ChangeEvent
+}
+
+func NewManager(store managerStore, modelCatalog *catalog.Catalog, authenticator *auth.StaticAuthenticator, logger *slog.Logger, pollInterval time.Duration) *Manager {
+ if pollInterval <= 0 {
+ pollInterval = 30 * time.Second
+ }
+ return &Manager{
+ store: store, catalog: modelCatalog, authenticator: authenticator,
+ logger: logger, pollInterval: pollInterval, broadcasts: make(chan ChangeEvent, broadcastQueueSize),
+ }
+}
+
+func (m *Manager) Reload(ctx context.Context) (int64, error) {
+ m.reloadMu.Lock()
+ defer m.reloadMu.Unlock()
+ snapshot, err := m.store.LoadSnapshot(ctx)
+ if err != nil {
+ return m.generation.Load(), err
+ }
+ m.catalog.Replace(snapshot.Models)
+ m.authenticator.ReplaceHashed(snapshot.APIKeys)
+ m.generation.Store(snapshot.Generation)
+ m.logger.Info("control_plane_reloaded", "generation", snapshot.Generation, "models", len(snapshot.Models), "api_keys", len(snapshot.APIKeys))
+ return snapshot.Generation, nil
+}
+
+func (m *Manager) AfterMutation(ctx context.Context, generation int64, resource, id string) error {
+ loadedGeneration, err := m.Reload(ctx)
+ if err != nil {
+ return err
+ }
+ if !m.store.RedisEnabled() {
+ return nil
+ }
+ if loadedGeneration > generation {
+ generation = loadedGeneration
+ }
+ event := newChange(generation, resource, id)
+ select {
+ case m.broadcasts <- event:
+ default:
+ m.logger.Warn("control_plane_publish_dropped", "generation", generation, "resource", resource, "id", id, "error", "broadcast queue is full")
+ }
+ return nil
+}
+
+func (m *Manager) Generation() int64 {
+ return m.generation.Load()
+}
+
+func (m *Manager) RedisConfigured() bool {
+ return m.store.RedisEnabled()
+}
+
+func (m *Manager) RedisConnected() bool {
+ return m.redisConnected.Load()
+}
+
+func (m *Manager) Run(ctx context.Context) {
+ var workers sync.WaitGroup
+ if m.store.RedisEnabled() {
+ workers.Add(2)
+ go func() {
+ defer workers.Done()
+ m.runSubscriptions(ctx)
+ }()
+ go func() {
+ defer workers.Done()
+ m.runBroadcasts(ctx)
+ }()
+ } else {
+ m.logger.Info("control_plane_redis_disabled", "fallback", "postgres_polling")
+ }
+ m.runPolling(ctx)
+ workers.Wait()
+}
+
+func (m *Manager) runPolling(ctx context.Context) {
+ ticker := time.NewTicker(m.pollInterval)
+ defer ticker.Stop()
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ generation, err := m.store.DatabaseGeneration(ctx)
+ if err != nil {
+ m.logger.Warn("control_plane_generation_check_failed", "error", err)
+ continue
+ }
+ if generation > m.generation.Load() {
+ if _, err := m.Reload(ctx); err != nil {
+ m.logger.Error("control_plane_reload_failed", "source", "postgres", "error", err)
+ }
+ }
+ }
+ }
+}
+
+func (m *Manager) runSubscriptions(ctx context.Context) {
+ retryDelay := m.pollInterval
+ if retryDelay > time.Second {
+ retryDelay = time.Second
+ }
+ if retryDelay <= 0 {
+ retryDelay = time.Second
+ }
+
+ for ctx.Err() == nil {
+ messages, closeSubscription, err := m.store.Subscribe(ctx)
+ if err != nil {
+ m.redisConnected.Store(false)
+ m.logger.Warn("control_plane_subscription_failed", "error", err)
+ if !waitForRetry(ctx, retryDelay) {
+ return
+ }
+ continue
+ }
+ m.redisConnected.Store(true)
+ m.logger.Info("control_plane_subscription_connected")
+
+ closed := false
+ for !closed {
+ select {
+ case <-ctx.Done():
+ closed = true
+ case message, ok := <-messages:
+ if !ok {
+ closed = true
+ continue
+ }
+ var event ChangeEvent
+ if json.Unmarshal([]byte(message.Payload), &event) != nil || event.Generation <= m.generation.Load() {
+ continue
+ }
+ if _, err := m.Reload(ctx); err != nil {
+ m.logger.Error("control_plane_reload_failed", "source", "redis", "error", err)
+ }
+ }
+ }
+ m.redisConnected.Store(false)
+ if closeSubscription != nil {
+ _ = closeSubscription()
+ }
+ if ctx.Err() == nil {
+ m.logger.Warn("control_plane_subscription_disconnected")
+ if !waitForRetry(ctx, retryDelay) {
+ return
+ }
+ }
+ }
+}
+
+func (m *Manager) runBroadcasts(ctx context.Context) {
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case event := <-m.broadcasts:
+ publishContext, cancel := context.WithTimeout(ctx, 2*time.Second)
+ err := m.store.PublishChange(publishContext, event)
+ cancel()
+ if err != nil {
+ m.logger.Warn("control_plane_publish_failed", "generation", event.Generation, "resource", event.Resource, "id", event.ID, "error", err)
+ }
+ }
+ }
+}
+
+func waitForRetry(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
+ }
+}