package capacitycontroller import ( "context" "encoding/json" "errors" "log/slog" "net/http" "os" "strings" "sync" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/config" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/jackc/pgx/v5" ) const ( controllerReconcileInterval = 10 * time.Second controllerLeadershipRetry = 5 * time.Second controllerDeletionCost = -1000 ) type capacityStore interface { TryAcquireCapacityControllerLeadership(context.Context) (store.Leadership, bool, error) WorkerQueueRuntime(context.Context) (store.WorkerQueueRuntime, error) ListWorkerInstanceRuntime(context.Context) ([]store.WorkerInstanceRuntime, error) CapacityDatabaseHealth(context.Context) (store.CapacityDatabaseHealth, error) MarkWorkerDraining(context.Context, string) error ReactivateWorkerInstance(context.Context, string) error } type capacityKubernetes interface { SiteState(context.Context, string) (KubernetesSiteState, error) ScaleWorkerDeployment(context.Context, string, int) error SetPodDeletionCost(context.Context, string, int) error } type Status struct { Leader bool `json:"leader"` LastRunAt time.Time `json:"lastRunAt,omitempty"` LastError string `json:"lastError,omitempty"` Queue store.WorkerQueueRuntime `json:"queue"` Plan Plan `json:"plan"` ScaleActions uint64 `json:"scaleActions"` } type Controller struct { cfg config.Config store capacityStore kubernetes capacityKubernetes logger *slog.Logger expectedRevision string now func() time.Time highSince time.Time lowSince time.Time statusMu sync.RWMutex status Status } func New( cfg config.Config, db capacityStore, kubernetes capacityKubernetes, logger *slog.Logger, ) *Controller { return &Controller{ cfg: cfg, store: db, kubernetes: kubernetes, logger: logger, expectedRevision: strings.TrimSpace(os.Getenv("AI_GATEWAY_REVISION")), now: time.Now, } } func (controller *Controller) Run(ctx context.Context) { for ctx.Err() == nil { leadership, acquired, err := controller.store.TryAcquireCapacityControllerLeadership(ctx) if err != nil { controller.setError(err) if !waitContext(ctx, controllerLeadershipRetry) { return } continue } if !acquired { controller.setLeader(false) controller.clearError() if !waitContext(ctx, controllerLeadershipRetry) { return } continue } controller.setLeader(true) controller.clearError() controller.runLeader(ctx, leadership) leadership.Release() controller.setLeader(false) } } func (controller *Controller) runLeader(ctx context.Context, leadership store.Leadership) { ticker := time.NewTicker(controllerReconcileInterval) defer ticker.Stop() if err := controller.reconcile(ctx); err != nil { controller.logError("initial capacity reconciliation failed", err) } for { select { case <-ctx.Done(): return case <-ticker.C: keepAliveCtx, cancel := context.WithTimeout(ctx, 5*time.Second) err := leadership.KeepAlive(keepAliveCtx) cancel() if err != nil { controller.logError("capacity controller leadership lost", err) return } if err := controller.reconcile(ctx); err != nil { controller.logError("capacity reconciliation failed", err) } } } } func (controller *Controller) reconcile(ctx context.Context) error { queue, err := controller.store.WorkerQueueRuntime(ctx) if err != nil { return err } instances, err := controller.store.ListWorkerInstanceRuntime(ctx) if err != nil { return err } database, err := controller.store.CapacityDatabaseHealth(ctx) if err != nil { return err } now := controller.now() sites := make([]SiteResources, 0, 2) kubernetesStates := make(map[string]KubernetesSiteState, 2) for _, site := range []string{"ningbo", "hongkong"} { state, stateErr := controller.kubernetes.SiteState(ctx, site) if stateErr != nil { return stateErr } kubernetesStates[site] = state minReplicas, maxReplicas := controller.siteReplicaBounds(site) sites = append(sites, SiteResources{ Site: site, CurrentReplicas: state.CurrentReplicas, MinReplicas: minReplicas, MaxReplicas: maxReplicas, AllocatableMemoryBytes: state.AllocatableMemoryBytes, UsedMemoryBytes: state.UsedMemoryBytes, WorkerRequestMemoryBytes: state.WorkerRequestMemoryBytes, AllocatableMilliCPU: state.AllocatableMilliCPU, UsedMilliCPU: state.UsedMilliCPU, WorkerRequestMilliCPU: state.WorkerRequestMilliCPU, MemoryPressure: state.MemoryPressure, Nodes: nodeResources(state.Nodes), }) } currentTotal := 0 for _, site := range sites { currentTotal += site.CurrentReplicas } target := controller.cfg.WorkerTargetOutstandingPerReplica if target < 1 { target = 2 * controller.cfg.AsyncWorkerInstanceHardLimit } rawDesired := (max(queue.Queued+queue.Running, 0) + target - 1) / target if rawDesired > currentTotal { if controller.highSince.IsZero() { controller.highSince = now } controller.lowSince = time.Time{} } else if rawDesired < currentTotal && queue.Queued == 0 { if controller.lowSince.IsZero() { controller.lowSince = now } controller.highSince = time.Time{} } else { controller.highSince = time.Time{} controller.lowSince = time.Time{} } revisionHealthy := controller.revisionsMatch(instances, kubernetesStates) scaleUpEligible := !controller.highSince.IsZero() && now.Sub(controller.highSince) >= time.Duration(controller.cfg.WorkerScaleUpWindowSeconds)*time.Second && revisionHealthy scaleDownEligible := !controller.lowSince.IsZero() && now.Sub(controller.lowSince) >= time.Duration(controller.cfg.WorkerScaleDownStabilizationSeconds)*time.Second && revisionHealthy plan := CalculatePlan(PlanInput{ Queued: queue.Queued, Running: queue.Running, InstanceSlots: controller.cfg.AsyncWorkerInstanceHardLimit, TargetOutstandingPerReplica: target, MemoryTargetPercent: controller.cfg.NodeMemoryTargetPercent, MemoryHardPercent: controller.cfg.NodeMemoryHardPercent, CPUTargetPercent: controller.cfg.NodeCPUTargetPercent, DatabaseConnections: database.Connections, DatabaseConnectionBudget: controller.cfg.PostgresConnectionBudget, NonWorkerConnectionBudget: controller.cfg.PostgresNonWorkerConnectionBudget, WorkerDatabasePoolMax: controller.cfg.WorkerDatabaseMaxConns, SynchronousDatabasePeers: database.SynchronousPeers, ScaleUpEligible: scaleUpEligible, ScaleDownEligible: scaleDownEligible, Sites: sites, }) if !revisionHealthy && plan.FrozenReason == "" { plan.FrozenReason = "release_revision_mismatch" } if controller.cfg.WorkerAutoscalingEnabled { for _, sitePlan := range plan.Sites { switch { case sitePlan.DesiredReplicas > sitePlan.CurrentReplicas: if err := controller.kubernetes.ScaleWorkerDeployment(ctx, sitePlan.Site, sitePlan.DesiredReplicas); err != nil { return err } controller.observeScale(sitePlan.Site, sitePlan.CurrentReplicas, sitePlan.DesiredReplicas, "scale_up") case sitePlan.DesiredReplicas < sitePlan.CurrentReplicas: if err := controller.reconcileScaleDown(ctx, now, sitePlan, instances); err != nil { return err } } } } controller.statusMu.Lock() controller.status.LastRunAt = now controller.status.LastError = "" controller.status.Queue = queue controller.status.Plan = plan controller.statusMu.Unlock() return nil } func nodeResources(states []KubernetesNodeState) []NodeResources { nodes := make([]NodeResources, 0, len(states)) for _, state := range states { nodes = append(nodes, NodeResources{ NodeName: state.NodeName, CurrentReplicas: state.CurrentReplicas, AllocatableMemoryBytes: state.AllocatableMemoryBytes, UsedMemoryBytes: state.UsedMemoryBytes, WorkerUsedMemoryBytes: state.WorkerUsedMemoryBytes, AllocatableMilliCPU: state.AllocatableMilliCPU, UsedMilliCPU: state.UsedMilliCPU, WorkerUsedMilliCPU: state.WorkerUsedMilliCPU, MemoryPressure: state.MemoryPressure, }) } return nodes } func (controller *Controller) reconcileScaleDown( ctx context.Context, now time.Time, sitePlan SitePlan, instances []store.WorkerInstanceRuntime, ) error { for _, instance := range instances { if instance.Site != sitePlan.Site || instance.Status != "draining" { continue } if instance.RunningTasks == 0 && instance.ActiveLeases == 0 { if err := controller.kubernetes.SetPodDeletionCost(ctx, instance.PodName, controllerDeletionCost); err != nil { return err } if err := controller.kubernetes.ScaleWorkerDeployment(ctx, sitePlan.Site, sitePlan.CurrentReplicas-1); err != nil { return err } controller.observeScale(sitePlan.Site, sitePlan.CurrentReplicas, sitePlan.CurrentReplicas-1, "drained_scale_down") return nil } if instance.DrainingAt != nil && now.Sub(*instance.DrainingAt) >= time.Duration(controller.cfg.WorkerDrainTimeoutSeconds)*time.Second { if err := controller.store.ReactivateWorkerInstance(ctx, instance.InstanceID); err != nil && !errors.Is(err, pgx.ErrNoRows) { return err } controller.observeScale(sitePlan.Site, sitePlan.CurrentReplicas, sitePlan.CurrentReplicas, "drain_timeout") } return nil } var candidate *store.WorkerInstanceRuntime for index := range instances { instance := &instances[index] if instance.Site != sitePlan.Site || instance.Status != "active" { continue } if candidate == nil || instance.RunningTasks+instance.ActiveLeases < candidate.RunningTasks+candidate.ActiveLeases { candidate = instance } } if candidate == nil { return nil } if err := controller.store.MarkWorkerDraining(ctx, candidate.InstanceID); err != nil && !errors.Is(err, pgx.ErrNoRows) { return err } controller.observeScale(sitePlan.Site, sitePlan.CurrentReplicas, sitePlan.CurrentReplicas, "drain_started") return nil } func (controller *Controller) siteReplicaBounds(site string) (int, int) { switch site { case "ningbo": return controller.cfg.WorkerMinReplicasNingbo, controller.cfg.WorkerMaxReplicasNingbo case "hongkong": return controller.cfg.WorkerMinReplicasHongkong, controller.cfg.WorkerMaxReplicasHongkong default: return 0, 0 } } func (controller *Controller) revisionsMatch( instances []store.WorkerInstanceRuntime, states map[string]KubernetesSiteState, ) bool { if controller.expectedRevision == "" { return true } for _, state := range states { if state.Revision != "" && state.Revision != controller.expectedRevision { return false } } for _, instance := range instances { if instance.Revision != "" && instance.Revision != controller.expectedRevision { return false } } return true } func (controller *Controller) Status() Status { controller.statusMu.RLock() defer controller.statusMu.RUnlock() return controller.status } func (controller *Controller) Handler() http.Handler { mux := http.NewServeMux() mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"ok":true,"service":"easyai-capacity-controller"}`)) }) mux.HandleFunc("GET /readyz", func(w http.ResponseWriter, _ *http.Request) { status := controller.Status() // Followers are intentionally idle while the database advisory lock is // held by the elected leader. They are still ready to take over and must // not make a two-replica Deployment permanently fail its rollout. if status.LastError != "" { http.Error(w, `{"ok":false}`, http.StatusServiceUnavailable) return } w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"ok":true}`)) }) mux.HandleFunc("GET /status", func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(controller.Status()) }) return mux } func (controller *Controller) setLeader(leader bool) { controller.statusMu.Lock() controller.status.Leader = leader controller.statusMu.Unlock() } func (controller *Controller) setError(err error) { controller.statusMu.Lock() controller.status.LastError = err.Error() controller.statusMu.Unlock() } func (controller *Controller) clearError() { controller.statusMu.Lock() controller.status.LastError = "" controller.statusMu.Unlock() } func (controller *Controller) logError(message string, err error) { controller.setError(err) if controller.logger != nil { controller.logger.Error(message, "error", err) } } func (controller *Controller) observeScale(site string, from int, to int, reason string) { controller.statusMu.Lock() controller.status.ScaleActions++ controller.statusMu.Unlock() if controller.logger != nil { controller.logger.Info("worker capacity action", "site", site, "fromReplicas", from, "toReplicas", to, "reason", reason, ) } } func waitContext(ctx context.Context, duration time.Duration) bool { timer := time.NewTimer(duration) defer timer.Stop() select { case <-ctx.Done(): return false case <-timer.C: return true } }