summaryrefslogtreecommitdiff
path: root/internal/controlplane
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/controlplane/manager.go213
-rw-r--r--internal/controlplane/manager_test.go212
-rw-r--r--internal/controlplane/mutations.go281
-rw-r--r--internal/controlplane/queries.go146
-rw-r--r--internal/controlplane/schema.sql82
-rw-r--r--internal/controlplane/snapshot.go150
-rw-r--r--internal/controlplane/store.go170
-rw-r--r--internal/controlplane/types.go138
8 files changed, 1392 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
+ }
+}
diff --git a/internal/controlplane/manager_test.go b/internal/controlplane/manager_test.go
new file mode 100644
index 0000000..8dd4012
--- /dev/null
+++ b/internal/controlplane/manager_test.go
@@ -0,0 +1,212 @@
+package controlplane
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "log/slog"
+ "strings"
+ "sync"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/catalog"
+)
+
+type fakeManagerStore struct {
+ redisEnabled bool
+ snapshot atomic.Value
+ databaseGeneration atomic.Int64
+ publishCalls atomic.Int64
+ subscribeCalls atomic.Int64
+ publishErr error
+ published chan ChangeEvent
+ subscribe func(context.Context, int64) (<-chan ChangeMessage, func() error, error)
+}
+
+type safeLogBuffer struct {
+ mu sync.Mutex
+ buf bytes.Buffer
+}
+
+func (b *safeLogBuffer) Write(data []byte) (int, error) {
+ b.mu.Lock()
+ defer b.mu.Unlock()
+ return b.buf.Write(data)
+}
+
+func (b *safeLogBuffer) String() string {
+ b.mu.Lock()
+ defer b.mu.Unlock()
+ return b.buf.String()
+}
+
+func newFakeManagerStore(generation int64) *fakeManagerStore {
+ store := &fakeManagerStore{published: make(chan ChangeEvent, 8)}
+ store.snapshot.Store(Snapshot{Generation: generation})
+ store.databaseGeneration.Store(generation)
+ return store
+}
+
+func (s *fakeManagerStore) LoadSnapshot(context.Context) (Snapshot, error) {
+ return s.snapshot.Load().(Snapshot), nil
+}
+
+func (s *fakeManagerStore) DatabaseGeneration(context.Context) (int64, error) {
+ return s.databaseGeneration.Load(), nil
+}
+
+func (s *fakeManagerStore) PublishChange(_ context.Context, event ChangeEvent) error {
+ s.publishCalls.Add(1)
+ s.published <- event
+ return s.publishErr
+}
+
+func (s *fakeManagerStore) Subscribe(ctx context.Context) (<-chan ChangeMessage, func() error, error) {
+ call := s.subscribeCalls.Add(1)
+ if s.subscribe != nil {
+ return s.subscribe(ctx, call)
+ }
+ channel := make(chan ChangeMessage)
+ return channel, func() error { return nil }, nil
+}
+
+func (s *fakeManagerStore) RedisEnabled() bool {
+ return s.redisEnabled
+}
+
+func newTestManager(store managerStore, logger *slog.Logger, interval time.Duration) *Manager {
+ return NewManager(store, catalog.NewModels(nil), auth.NewDynamic(nil, false), logger, interval)
+}
+
+func TestAfterMutationReloadsLocallyWhenRedisPublishFails(t *testing.T) {
+ store := newFakeManagerStore(7)
+ store.redisEnabled = true
+ store.publishErr = errors.New("redis unavailable")
+ var logs safeLogBuffer
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&logs, nil)), 10*time.Millisecond)
+ ctx, cancel := context.WithCancel(context.Background())
+ done := make(chan struct{})
+ go func() {
+ manager.Run(ctx)
+ close(done)
+ }()
+
+ if err := manager.AfterMutation(context.Background(), 7, "model", "model-1"); err != nil {
+ t.Fatalf("mutation unexpectedly failed: %v", err)
+ }
+ if manager.Generation() != 7 {
+ t.Fatalf("local generation = %d, want 7", manager.Generation())
+ }
+ select {
+ case event := <-store.published:
+ if event.Generation != 7 || event.Resource != "model" {
+ t.Fatalf("unexpected event: %+v", event)
+ }
+ case <-time.After(time.Second):
+ t.Fatal("broadcast was not attempted")
+ }
+ waitUntil(t, time.Second, func() bool { return strings.Contains(logs.String(), "control_plane_publish_failed") })
+
+ cancel()
+ select {
+ case <-done:
+ case <-time.After(time.Second):
+ t.Fatal("manager did not stop")
+ }
+}
+
+func TestPollingContinuesWhileRedisSubscribeIsBlocked(t *testing.T) {
+ store := newFakeManagerStore(1)
+ store.redisEnabled = true
+ store.subscribe = func(ctx context.Context, _ int64) (<-chan ChangeMessage, func() error, error) {
+ <-ctx.Done()
+ return nil, nil, ctx.Err()
+ }
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&safeLogBuffer{}, nil)), 10*time.Millisecond)
+ if _, err := manager.Reload(context.Background()); err != nil {
+ t.Fatal(err)
+ }
+
+ ctx, cancel := context.WithCancel(context.Background())
+ done := make(chan struct{})
+ go func() {
+ manager.Run(ctx)
+ close(done)
+ }()
+ store.snapshot.Store(Snapshot{Generation: 2})
+ store.databaseGeneration.Store(2)
+ waitUntil(t, time.Second, func() bool { return manager.Generation() == 2 })
+
+ cancel()
+ select {
+ case <-done:
+ case <-time.After(time.Second):
+ t.Fatal("manager did not stop")
+ }
+}
+
+func TestSubscriptionReconnectsAfterChannelCloses(t *testing.T) {
+ store := newFakeManagerStore(1)
+ store.redisEnabled = true
+ first := make(chan ChangeMessage)
+ second := make(chan ChangeMessage, 1)
+ store.subscribe = func(_ context.Context, call int64) (<-chan ChangeMessage, func() error, error) {
+ if call == 1 {
+ return first, func() error { return nil }, nil
+ }
+ return second, func() error { return nil }, nil
+ }
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&safeLogBuffer{}, nil)), 10*time.Millisecond)
+ if _, err := manager.Reload(context.Background()); err != nil {
+ t.Fatal(err)
+ }
+
+ ctx, cancel := context.WithCancel(context.Background())
+ done := make(chan struct{})
+ go func() {
+ manager.Run(ctx)
+ close(done)
+ }()
+ waitUntil(t, time.Second, func() bool { return store.subscribeCalls.Load() == 1 })
+ close(first)
+ waitUntil(t, time.Second, func() bool { return store.subscribeCalls.Load() >= 2 && manager.RedisConnected() })
+ store.snapshot.Store(Snapshot{Generation: 2})
+ second <- ChangeMessage{Payload: `{"generation":2,"resource":"model"}`}
+ waitUntil(t, time.Second, func() bool { return manager.Generation() == 2 })
+
+ cancel()
+ select {
+ case <-done:
+ case <-time.After(time.Second):
+ t.Fatal("manager did not stop")
+ }
+}
+
+func TestRedisCanBeDisabled(t *testing.T) {
+ store := newFakeManagerStore(3)
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&bytes.Buffer{}, nil)), 10*time.Millisecond)
+ if err := manager.AfterMutation(context.Background(), 3, "tenant", "tenant-1"); err != nil {
+ t.Fatal(err)
+ }
+ if manager.RedisConfigured() || manager.RedisConnected() {
+ t.Fatal("Redis unexpectedly reported as available")
+ }
+ if store.publishCalls.Load() != 0 || store.subscribeCalls.Load() != 0 {
+ t.Fatal("Redis operations were attempted while disabled")
+ }
+}
+
+func waitUntil(t *testing.T, timeout time.Duration, condition func() bool) {
+ t.Helper()
+ deadline := time.Now().Add(timeout)
+ for time.Now().Before(deadline) {
+ if condition() {
+ return
+ }
+ time.Sleep(5 * time.Millisecond)
+ }
+ t.Fatal("condition was not met before timeout")
+}
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
new file mode 100644
index 0000000..a781994
--- /dev/null
+++ b/internal/controlplane/mutations.go
@@ -0,0 +1,281 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/rand"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "net/url"
+ "regexp"
+ "strings"
+
+ "github.com/jackc/pgx/v5"
+)
+
+var (
+ ErrNotFound = errors.New("control-plane resource not found")
+ slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`)
+)
+
+func (s *Store) CreateTenant(ctx context.Context, input CreateTenantInput) (Tenant, int64, error) {
+ input.Slug = strings.ToLower(strings.TrimSpace(input.Slug))
+ input.Name = strings.TrimSpace(input.Name)
+ if !slugPattern.MatchString(input.Slug) || input.Name == "" {
+ return Tenant{}, 0, errors.New("tenant requires a 3-64 character lowercase slug and a name")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Tenant{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Tenant
+ err = tx.QueryRow(ctx, `
+ INSERT INTO tenants (slug, name) VALUES ($1, $2)
+ RETURNING id::text, slug, name, status, created_at`, input.Slug, input.Name,
+ ).Scan(&result.ID, &result.Slug, &result.Name, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return Tenant{}, 0, fmt.Errorf("create tenant: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Tenant{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Tenant{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) CreateProject(ctx context.Context, input CreateProjectInput) (Project, int64, error) {
+ input.Slug = strings.ToLower(strings.TrimSpace(input.Slug))
+ input.Name = strings.TrimSpace(input.Name)
+ if input.TenantID == "" || !slugPattern.MatchString(input.Slug) || input.Name == "" {
+ return Project{}, 0, errors.New("project requires tenant_id, a 3-64 character lowercase slug, and a name")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Project{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Project
+ err = tx.QueryRow(ctx, `
+ INSERT INTO projects (tenant_id, slug, name) VALUES ($1, $2, $3)
+ RETURNING id::text, tenant_id::text, slug, name, status, created_at`, input.TenantID, input.Slug, input.Name,
+ ).Scan(&result.ID, &result.TenantID, &result.Slug, &result.Name, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return Project{}, 0, fmt.Errorf("create project: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Project{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Project{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (CreatedAPIKey, int64, error) {
+ input.Name = strings.TrimSpace(input.Name)
+ if input.TenantID == "" || input.ProjectID == "" || input.Name == "" {
+ return CreatedAPIKey{}, 0, errors.New("API key requires tenant_id, project_id, and name")
+ }
+ if len(input.Scopes) == 0 {
+ input.Scopes = []string{"inference"}
+ }
+ scopes := uniqueStrings(input.Scopes)
+ scopesJSON, _ := json.Marshal(scopes)
+ random := make([]byte, 32)
+ if _, err := rand.Read(random); err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("generate API key: %w", err)
+ }
+ rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random)
+ hash := sha256.Sum256([]byte(rawKey))
+ prefix := rawKey[:min(18, len(rawKey))] + "..."
+
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result CreatedAPIKey
+ err = tx.QueryRow(ctx, `
+ INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_hash, scopes)
+ VALUES ($1, $2, $3, $4, $5, $6)
+ RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at`,
+ input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON,
+ ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &scopesJSON, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err)
+ }
+ result.Scopes = scopes
+ result.Key = rawKey
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) {
+ return s.toggle(ctx, `UPDATE api_keys SET status = 'revoked', revoked_at = now() WHERE id = $1 AND status <> 'revoked'`, id)
+}
+
+func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) {
+ input.Name = strings.TrimSpace(input.Name)
+ input.BaseURL = strings.TrimRight(strings.TrimSpace(input.BaseURL), "/")
+ if input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") {
+ return Provider{}, 0, errors.New("provider requires name, protocol openai|anthropic, base_url, and api_key")
+ }
+ parsed, err := url.Parse(input.BaseURL)
+ if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") {
+ return Provider{}, 0, errors.New("provider base_url must be an absolute http(s) URL")
+ }
+ ciphertext, err := s.cipher.Encrypt(input.APIKey)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Provider
+ err = tx.QueryRow(ctx, `
+ INSERT INTO providers (name, protocol, base_url, api_key_ciphertext)
+ VALUES ($1, $2, $3, $4)
+ RETURNING id::text, name, protocol, base_url, enabled, created_at`,
+ input.Name, input.Protocol, input.BaseURL, ciphertext,
+ ).Scan(&result.ID, &result.Name, &result.Protocol, &result.BaseURL, &result.Enabled, &result.CreatedAt)
+ if err != nil {
+ return Provider{}, 0, fmt.Errorf("create provider: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Provider{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
+ return s.toggle(ctx, `UPDATE providers SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
+}
+
+func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, int64, error) {
+ input.PublicID = strings.TrimSpace(input.PublicID)
+ input.OwnedBy = strings.TrimSpace(input.OwnedBy)
+ if input.PublicID == "" || len(input.Routes) == 0 {
+ return Model{}, 0, errors.New("model requires public_id and at least one route")
+ }
+ for i := range input.Routes {
+ input.Routes[i].ProviderID = strings.TrimSpace(input.Routes[i].ProviderID)
+ input.Routes[i].UpstreamModel = strings.TrimSpace(input.Routes[i].UpstreamModel)
+ if input.Routes[i].Weight == 0 {
+ input.Routes[i].Weight = 1
+ }
+ if input.Routes[i].ProviderID == "" || input.Routes[i].UpstreamModel == "" || input.Routes[i].Priority < 0 || input.Routes[i].Weight < 1 || input.Routes[i].Weight > 100 {
+ return Model{}, 0, fmt.Errorf("route %d has invalid provider, upstream model, priority, or weight", i+1)
+ }
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Model{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Model
+ err = tx.QueryRow(ctx, `
+ INSERT INTO models (public_id, owned_by) VALUES ($1, $2)
+ RETURNING id::text, public_id, owned_by, enabled, created_at`, input.PublicID, input.OwnedBy,
+ ).Scan(&result.ID, &result.PublicID, &result.OwnedBy, &result.Enabled, &result.CreatedAt)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model: %w", err)
+ }
+ result.Routes = make([]Route, 0, len(input.Routes))
+ for _, route := range input.Routes {
+ var created Route
+ err := tx.QueryRow(ctx, `
+ INSERT INTO model_routes (model_id, provider_id, upstream_model, priority, weight)
+ VALUES ($1, $2, $3, $4, $5)
+ RETURNING id::text, provider_id::text, upstream_model, priority, weight, enabled`,
+ result.ID, route.ProviderID, route.UpstreamModel, route.Priority, route.Weight,
+ ).Scan(&created.ID, &created.ProviderID, &created.UpstreamModel, &created.Priority, &created.Weight, &created.Enabled)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model route: %w", err)
+ }
+ result.Routes = append(result.Routes, created)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Model{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Model{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) SetModelEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
+ return s.toggle(ctx, `UPDATE models SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
+}
+
+func (s *Store) toggle(ctx context.Context, query, id string, args ...any) (int64, error) {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ parameters := append([]any{id}, args...)
+ command, err := tx.Exec(ctx, query, parameters...)
+ if err != nil {
+ return 0, err
+ }
+ if command.RowsAffected() == 0 {
+ return 0, ErrNotFound
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return 0, err
+ }
+ return generation, nil
+}
+
+func bumpGeneration(ctx context.Context, tx pgx.Tx) (int64, error) {
+ var generation int64
+ err := tx.QueryRow(ctx, `
+ UPDATE control_state SET generation = generation + 1, updated_at = now()
+ WHERE singleton = TRUE RETURNING generation`,
+ ).Scan(&generation)
+ if err != nil {
+ return 0, fmt.Errorf("advance control-plane generation: %w", err)
+ }
+ return generation, nil
+}
+
+func uniqueStrings(values []string) []string {
+ seen := make(map[string]struct{}, len(values))
+ result := make([]string, 0, len(values))
+ for _, value := range values {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ continue
+ }
+ if _, exists := seen[value]; exists {
+ continue
+ }
+ seen[value] = struct{}{}
+ result = append(result, value)
+ }
+ return result
+}
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go
new file mode 100644
index 0000000..732cb46
--- /dev/null
+++ b/internal/controlplane/queries.go
@@ -0,0 +1,146 @@
+package controlplane
+
+import (
+ "context"
+ "encoding/json"
+ "fmt"
+)
+
+func (s *Store) Overview(ctx context.Context) (Overview, error) {
+ var result Overview
+ err := s.db.QueryRow(ctx, `
+ SELECT
+ (SELECT generation FROM control_state WHERE singleton = TRUE),
+ (SELECT count(*) FROM tenants WHERE status = 'active'),
+ (SELECT count(*) FROM projects WHERE status = 'active'),
+ (SELECT count(*) FROM api_keys WHERE status = 'active'),
+ (SELECT count(*) FROM providers WHERE enabled = TRUE),
+ (SELECT count(*) FROM models WHERE enabled = TRUE)`,
+ ).Scan(&result.Generation, &result.Tenants, &result.Projects, &result.APIKeys, &result.Providers, &result.Models)
+ if err != nil {
+ return Overview{}, fmt.Errorf("query control-plane overview: %w", err)
+ }
+ return result, nil
+}
+
+func (s *Store) ListTenants(ctx context.Context) ([]Tenant, error) {
+ rows, err := s.db.Query(ctx, `SELECT id::text, slug, name, status, created_at FROM tenants ORDER BY created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query tenants: %w", err)
+ }
+ defer rows.Close()
+ result := make([]Tenant, 0)
+ for rows.Next() {
+ var item Tenant
+ if err := rows.Scan(&item.ID, &item.Slug, &item.Name, &item.Status, &item.CreatedAt); err != nil {
+ return nil, fmt.Errorf("scan tenant: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListProjects(ctx context.Context) ([]Project, error) {
+ rows, err := s.db.Query(ctx, `SELECT id::text, tenant_id::text, slug, name, status, created_at FROM projects ORDER BY created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query projects: %w", err)
+ }
+ defer rows.Close()
+ result := make([]Project, 0)
+ for rows.Next() {
+ var item Project
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.Slug, &item.Name, &item.Status, &item.CreatedAt); err != nil {
+ return nil, fmt.Errorf("scan project: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListAPIKeys(ctx context.Context) ([]APIKey, error) {
+ rows, err := s.db.Query(ctx, `
+ SELECT id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at
+ FROM api_keys ORDER BY created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query API keys: %w", err)
+ }
+ defer rows.Close()
+ result := make([]APIKey, 0)
+ for rows.Next() {
+ var item APIKey
+ var scopesJSON []byte
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, &scopesJSON, &item.Status, &item.CreatedAt); err != nil {
+ return nil, fmt.Errorf("scan API key: %w", err)
+ }
+ if err := json.Unmarshal(scopesJSON, &item.Scopes); err != nil {
+ return nil, fmt.Errorf("decode API key scopes: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) {
+ rows, err := s.db.Query(ctx, `
+ SELECT p.id::text, p.name, p.protocol, p.base_url, p.enabled, count(r.id), p.created_at
+ FROM providers p LEFT JOIN model_routes r ON r.provider_id = p.id
+ GROUP BY p.id ORDER BY p.created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query providers: %w", err)
+ }
+ defer rows.Close()
+ result := make([]Provider, 0)
+ for rows.Next() {
+ var item Provider
+ if err := rows.Scan(&item.ID, &item.Name, &item.Protocol, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil {
+ return nil, fmt.Errorf("scan provider: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
+ rows, err := s.db.Query(ctx, `SELECT id::text, public_id, owned_by, enabled, created_at FROM models ORDER BY public_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query models: %w", err)
+ }
+ models := make([]Model, 0)
+ positions := make(map[string]int)
+ for rows.Next() {
+ var item Model
+ if err := rows.Scan(&item.ID, &item.PublicID, &item.OwnedBy, &item.Enabled, &item.CreatedAt); err != nil {
+ rows.Close()
+ return nil, fmt.Errorf("scan model: %w", err)
+ }
+ item.Routes = []Route{}
+ positions[item.ID] = len(models)
+ models = append(models, item)
+ }
+ if err := rows.Err(); err != nil {
+ rows.Close()
+ return nil, err
+ }
+ rows.Close()
+
+ routeRows, err := s.db.Query(ctx, `
+ SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol,
+ r.upstream_model, r.priority, r.weight, r.enabled
+ FROM model_routes r JOIN providers p ON p.id = r.provider_id
+ ORDER BY r.priority, r.created_at`)
+ if err != nil {
+ return nil, fmt.Errorf("query model routes: %w", err)
+ }
+ defer routeRows.Close()
+ for routeRows.Next() {
+ var route Route
+ var modelID string
+ if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil {
+ return nil, fmt.Errorf("scan model route: %w", err)
+ }
+ if position, ok := positions[modelID]; ok {
+ models[position].Routes = append(models[position].Routes, route)
+ }
+ }
+ return models, routeRows.Err()
+}
diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql
new file mode 100644
index 0000000..25af7c9
--- /dev/null
+++ b/internal/controlplane/schema.sql
@@ -0,0 +1,82 @@
+CREATE EXTENSION IF NOT EXISTS pgcrypto;
+
+CREATE TABLE IF NOT EXISTS control_state (
+ singleton BOOLEAN PRIMARY KEY DEFAULT TRUE CHECK (singleton),
+ generation BIGINT NOT NULL DEFAULT 0,
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+INSERT INTO control_state (singleton) VALUES (TRUE) ON CONFLICT DO NOTHING;
+
+CREATE TABLE IF NOT EXISTS tenants (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ slug TEXT NOT NULL UNIQUE,
+ name TEXT NOT NULL,
+ status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'suspended')),
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS projects (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ slug TEXT NOT NULL,
+ name TEXT NOT NULL,
+ status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'suspended')),
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ UNIQUE (tenant_id, slug),
+ UNIQUE (id, tenant_id)
+);
+
+CREATE TABLE IF NOT EXISTS api_keys (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ project_id UUID NOT NULL,
+ name TEXT NOT NULL,
+ key_prefix TEXT NOT NULL,
+ key_hash BYTEA NOT NULL UNIQUE CHECK (octet_length(key_hash) = 32),
+ scopes JSONB NOT NULL DEFAULT '["inference"]'::jsonb,
+ status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'revoked')),
+ last_used_at TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ revoked_at TIMESTAMPTZ,
+ FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE
+);
+
+CREATE TABLE IF NOT EXISTS providers (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ name TEXT NOT NULL UNIQUE,
+ protocol TEXT NOT NULL CHECK (protocol IN ('openai', 'anthropic')),
+ base_url TEXT NOT NULL,
+ api_key_ciphertext BYTEA NOT NULL,
+ enabled BOOLEAN NOT NULL DEFAULT TRUE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS models (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ public_id TEXT NOT NULL UNIQUE,
+ owned_by TEXT NOT NULL DEFAULT '',
+ enabled BOOLEAN NOT NULL DEFAULT TRUE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS model_routes (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ provider_id UUID NOT NULL REFERENCES providers(id) ON DELETE RESTRICT,
+ upstream_model TEXT NOT NULL,
+ priority INTEGER NOT NULL DEFAULT 0 CHECK (priority >= 0),
+ weight INTEGER NOT NULL DEFAULT 1 CHECK (weight BETWEEN 1 AND 100),
+ enabled BOOLEAN NOT NULL DEFAULT TRUE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ UNIQUE (model_id, provider_id, upstream_model)
+);
+
+CREATE INDEX IF NOT EXISTS api_keys_active_hash_idx ON api_keys (key_hash) WHERE status = 'active';
+CREATE INDEX IF NOT EXISTS projects_tenant_idx ON projects (tenant_id);
+CREATE INDEX IF NOT EXISTS model_routes_model_idx ON model_routes (model_id) WHERE enabled;
+CREATE INDEX IF NOT EXISTS model_routes_provider_idx ON model_routes (provider_id) WHERE enabled;
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go
new file mode 100644
index 0000000..c8931ab
--- /dev/null
+++ b/internal/controlplane/snapshot.go
@@ -0,0 +1,150 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/sha256"
+ "encoding/json"
+ "fmt"
+ "strings"
+
+ "aigw/internal/auth"
+ "aigw/internal/domain"
+
+ "github.com/jackc/pgx/v5"
+)
+
+func (s *Store) LoadSnapshot(ctx context.Context) (Snapshot, error) {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly})
+ if err != nil {
+ return Snapshot{}, fmt.Errorf("begin snapshot transaction: %w", err)
+ }
+ defer tx.Rollback(ctx)
+
+ var result Snapshot
+ if err := tx.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&result.Generation); err != nil {
+ return Snapshot{}, fmt.Errorf("read snapshot generation: %w", err)
+ }
+
+ providers, err := s.loadProviders(ctx, tx)
+ if err != nil {
+ return Snapshot{}, err
+ }
+ result.Models, err = loadModels(ctx, tx, providers)
+ if err != nil {
+ return Snapshot{}, err
+ }
+ result.APIKeys, err = loadAPIKeys(ctx, tx)
+ if err != nil {
+ return Snapshot{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Snapshot{}, fmt.Errorf("commit snapshot transaction: %w", err)
+ }
+ return result, nil
+}
+
+func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain.Provider, error) {
+ rows, err := tx.Query(ctx, `
+ SELECT id::text, name, protocol, base_url, api_key_ciphertext
+ FROM providers
+ WHERE enabled = TRUE
+ ORDER BY name`)
+ if err != nil {
+ return nil, fmt.Errorf("query providers: %w", err)
+ }
+ defer rows.Close()
+ providers := make(map[string]domain.Provider)
+ for rows.Next() {
+ var id, name, protocol, baseURL string
+ var ciphertext []byte
+ if err := rows.Scan(&id, &name, &protocol, &baseURL, &ciphertext); err != nil {
+ return nil, fmt.Errorf("scan provider: %w", err)
+ }
+ apiKey, err := s.cipher.Decrypt(ciphertext)
+ if err != nil {
+ return nil, fmt.Errorf("decrypt provider %q credential: %w", name, err)
+ }
+ providers[id] = domain.Provider{ID: id, Protocol: domain.Protocol(protocol), BaseURL: strings.TrimRight(baseURL, "/"), APIKey: apiKey}
+ }
+ if err := rows.Err(); err != nil {
+ return nil, fmt.Errorf("read providers: %w", err)
+ }
+ return providers, nil
+}
+
+func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) {
+ rows, err := tx.Query(ctx, `
+ SELECT m.public_id, m.owned_by, r.provider_id::text, r.upstream_model, r.priority, r.weight
+ FROM models m
+ JOIN model_routes r ON r.model_id = m.id AND r.enabled = TRUE
+ JOIN providers p ON p.id = r.provider_id AND p.enabled = TRUE
+ WHERE m.enabled = TRUE
+ ORDER BY m.public_id, r.priority, r.created_at`)
+ if err != nil {
+ return nil, fmt.Errorf("query model routes: %w", err)
+ }
+ defer rows.Close()
+ models := make([]domain.Model, 0)
+ index := make(map[string]int)
+ for rows.Next() {
+ var publicID, ownedBy, providerID, upstreamModel string
+ var priority, weight int
+ if err := rows.Scan(&publicID, &ownedBy, &providerID, &upstreamModel, &priority, &weight); err != nil {
+ return nil, fmt.Errorf("scan model route: %w", err)
+ }
+ provider, ok := providers[providerID]
+ if !ok {
+ continue
+ }
+ position, exists := index[publicID]
+ if !exists {
+ position = len(models)
+ index[publicID] = position
+ models = append(models, domain.Model{ID: publicID, OwnedBy: ownedBy})
+ }
+ models[position].Routes = append(models[position].Routes, domain.Route{
+ Provider: provider, UpstreamModel: upstreamModel, Priority: priority, Weight: weight,
+ })
+ }
+ if err := rows.Err(); err != nil {
+ return nil, fmt.Errorf("read model routes: %w", err)
+ }
+ return models, nil
+}
+
+func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) {
+ rows, err := tx.Query(ctx, `
+ SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes
+ FROM api_keys k
+ JOIN tenants t ON t.id = k.tenant_id AND t.status = 'active'
+ JOIN projects p ON p.id = k.project_id AND p.status = 'active'
+ WHERE k.status = 'active'`)
+ if err != nil {
+ return nil, fmt.Errorf("query API keys: %w", err)
+ }
+ defer rows.Close()
+ records := make([]auth.HashedKeyRecord, 0)
+ for rows.Next() {
+ var keyID, tenantID, projectID string
+ var hashBytes, scopesJSON []byte
+ if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON); err != nil {
+ return nil, fmt.Errorf("scan API key: %w", err)
+ }
+ if len(hashBytes) != sha256.Size {
+ return nil, fmt.Errorf("API key %s has invalid hash length", keyID)
+ }
+ var hash [sha256.Size]byte
+ copy(hash[:], hashBytes)
+ var scopes []string
+ if err := json.Unmarshal(scopesJSON, &scopes); err != nil {
+ return nil, fmt.Errorf("decode API key %s scopes: %w", keyID, err)
+ }
+ records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{
+ KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes,
+ }})
+ }
+ if err := rows.Err(); err != nil {
+ return nil, fmt.Errorf("read API keys: %w", err)
+ }
+ return records, nil
+}
diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go
new file mode 100644
index 0000000..833f846
--- /dev/null
+++ b/internal/controlplane/store.go
@@ -0,0 +1,170 @@
+package controlplane
+
+import (
+ "context"
+ _ "embed"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "aigw/internal/security"
+
+ "github.com/jackc/pgx/v5/pgxpool"
+ "github.com/redis/go-redis/v9"
+)
+
+//go:embed schema.sql
+var schemaSQL string
+
+var ErrRedisDisabled = errors.New("Redis propagation is disabled")
+
+type Options struct {
+ DatabaseURL string
+ RedisURL string
+ CredentialKey string
+ RedisChannel string
+ VersionCacheKey string
+}
+
+type Store struct {
+ db *pgxpool.Pool
+ redis *redis.Client
+ cipher *security.CredentialCipher
+ redisChannel string
+ versionCacheKey string
+}
+
+func NewStore(ctx context.Context, options Options) (*Store, error) {
+ cipher, err := security.NewCredentialCipher(options.CredentialKey)
+ if err != nil {
+ return nil, err
+ }
+ db, err := pgxpool.New(ctx, options.DatabaseURL)
+ if err != nil {
+ return nil, fmt.Errorf("configure PostgreSQL: %w", err)
+ }
+ if err := db.Ping(ctx); err != nil {
+ db.Close()
+ return nil, fmt.Errorf("connect PostgreSQL: %w", err)
+ }
+ var redisClient *redis.Client
+ if strings.TrimSpace(options.RedisURL) != "" {
+ redisOptions, err := redis.ParseURL(options.RedisURL)
+ if err != nil {
+ db.Close()
+ return nil, fmt.Errorf("parse Redis URL: %w", err)
+ }
+ redisClient = redis.NewClient(redisOptions)
+ }
+ return &Store{
+ db: db, redis: redisClient, cipher: cipher,
+ redisChannel: options.RedisChannel, versionCacheKey: options.VersionCacheKey,
+ }, nil
+}
+
+func (s *Store) Close() error {
+ s.db.Close()
+ if s.redis == nil {
+ return nil
+ }
+ return s.redis.Close()
+}
+
+func (s *Store) RedisEnabled() bool {
+ return s.redis != nil
+}
+
+func (s *Store) Migrate(ctx context.Context) error {
+ if _, err := s.db.Exec(ctx, schemaSQL); err != nil {
+ return fmt.Errorf("apply control-plane schema: %w", err)
+ }
+ return nil
+}
+
+func MigrateDatabase(ctx context.Context, databaseURL string) error {
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ return fmt.Errorf("configure PostgreSQL: %w", err)
+ }
+ defer db.Close()
+ if _, err := db.Exec(ctx, schemaSQL); err != nil {
+ return fmt.Errorf("apply control-plane schema: %w", err)
+ }
+ return nil
+}
+
+func (s *Store) DatabaseGeneration(ctx context.Context) (int64, error) {
+ var generation int64
+ err := s.db.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&generation)
+ if err != nil {
+ return 0, fmt.Errorf("read control-plane generation: %w", err)
+ }
+ return generation, nil
+}
+
+func (s *Store) RedisGeneration(ctx context.Context) (int64, error) {
+ if s.redis == nil {
+ return 0, ErrRedisDisabled
+ }
+ generation, err := s.redis.Get(ctx, s.versionCacheKey).Int64()
+ if errors.Is(err, redis.Nil) {
+ return 0, nil
+ }
+ return generation, err
+}
+
+func (s *Store) PublishChange(ctx context.Context, event ChangeEvent) error {
+ if s.redis == nil {
+ return ErrRedisDisabled
+ }
+ payload, err := json.Marshal(event)
+ if err != nil {
+ return err
+ }
+ pipeline := s.redis.TxPipeline()
+ pipeline.Set(ctx, s.versionCacheKey, event.Generation, 0)
+ pipeline.Publish(ctx, s.redisChannel, payload)
+ _, err = pipeline.Exec(ctx)
+ if err != nil {
+ return fmt.Errorf("publish control-plane change: %w", err)
+ }
+ return nil
+}
+
+func (s *Store) Subscribe(ctx context.Context) (<-chan ChangeMessage, func() error, error) {
+ if s.redis == nil {
+ return nil, nil, ErrRedisDisabled
+ }
+ pubsub := s.redis.Subscribe(ctx, s.redisChannel)
+ if _, err := pubsub.Receive(ctx); err != nil {
+ _ = pubsub.Close()
+ return nil, nil, fmt.Errorf("subscribe control-plane changes: %w", err)
+ }
+ messages := make(chan ChangeMessage)
+ redisMessages := pubsub.Channel()
+ go func() {
+ defer close(messages)
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case message, ok := <-redisMessages:
+ if !ok {
+ return
+ }
+ select {
+ case messages <- ChangeMessage{Payload: message.Payload}:
+ case <-ctx.Done():
+ return
+ }
+ }
+ }
+ }()
+ return messages, pubsub.Close, nil
+}
+
+func newChange(generation int64, resource, id string) ChangeEvent {
+ return ChangeEvent{Generation: generation, Resource: resource, ID: id, ChangedAt: time.Now().UTC().Format(time.RFC3339Nano)}
+}
diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go
new file mode 100644
index 0000000..24f2843
--- /dev/null
+++ b/internal/controlplane/types.go
@@ -0,0 +1,138 @@
+package controlplane
+
+import (
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/domain"
+)
+
+type Tenant struct {
+ ID string `json:"id"`
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+ Status string `json:"status"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type Project struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+ Status string `json:"status"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type APIKey struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ ProjectID string `json:"project_id"`
+ Name string `json:"name"`
+ KeyPrefix string `json:"key_prefix"`
+ Scopes []string `json:"scopes"`
+ Status string `json:"status"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type CreatedAPIKey struct {
+ APIKey
+ Key string `json:"key"`
+}
+
+type Provider struct {
+ ID string `json:"id"`
+ Name string `json:"name"`
+ Protocol string `json:"protocol"`
+ BaseURL string `json:"base_url"`
+ Enabled bool `json:"enabled"`
+ RouteCount int `json:"route_count"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type Route struct {
+ ID string `json:"id"`
+ ProviderID string `json:"provider_id"`
+ ProviderName string `json:"provider_name"`
+ Protocol string `json:"protocol"`
+ UpstreamModel string `json:"upstream_model"`
+ Priority int `json:"priority"`
+ Weight int `json:"weight"`
+ Enabled bool `json:"enabled"`
+}
+
+type Model struct {
+ ID string `json:"id"`
+ PublicID string `json:"public_id"`
+ OwnedBy string `json:"owned_by"`
+ Enabled bool `json:"enabled"`
+ Routes []Route `json:"routes"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type Overview struct {
+ Generation int64 `json:"generation"`
+ RuntimeGeneration int64 `json:"runtime_generation"`
+ RedisConfigured bool `json:"redis_configured"`
+ RedisConnected bool `json:"redis_connected"`
+ Tenants int64 `json:"tenants"`
+ Projects int64 `json:"projects"`
+ APIKeys int64 `json:"api_keys"`
+ Providers int64 `json:"providers"`
+ Models int64 `json:"models"`
+}
+
+type Snapshot struct {
+ Generation int64
+ Models []domain.Model
+ APIKeys []auth.HashedKeyRecord
+}
+
+type ChangeEvent struct {
+ Generation int64 `json:"generation"`
+ Resource string `json:"resource"`
+ ID string `json:"id,omitempty"`
+ ChangedAt string `json:"changed_at"`
+}
+
+type ChangeMessage struct {
+ Payload string
+}
+
+type CreateTenantInput struct {
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+}
+
+type CreateProjectInput struct {
+ TenantID string `json:"tenant_id"`
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+}
+
+type CreateAPIKeyInput struct {
+ TenantID string `json:"tenant_id"`
+ ProjectID string `json:"project_id"`
+ Name string `json:"name"`
+ Scopes []string `json:"scopes"`
+}
+
+type CreateProviderInput struct {
+ Name string `json:"name"`
+ Protocol string `json:"protocol"`
+ BaseURL string `json:"base_url"`
+ APIKey string `json:"api_key"`
+}
+
+type RouteInput struct {
+ ProviderID string `json:"provider_id"`
+ UpstreamModel string `json:"upstream_model"`
+ Priority int `json:"priority"`
+ Weight int `json:"weight"`
+}
+
+type CreateModelInput struct {
+ PublicID string `json:"public_id"`
+ OwnedBy string `json:"owned_by"`
+ Routes []RouteInput `json:"routes"`
+}