package runner import ( "context" "errors" "fmt" "os" "strings" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/riverqueue/river" "github.com/riverqueue/river/riverdriver/riverpgxv5" "github.com/riverqueue/river/rivermigrate" "github.com/riverqueue/river/rivertype" ) const asyncTaskQueueName = "gateway_tasks" type asyncTaskArgs struct { TaskID string `json:"task_id" river:"unique"` } type asyncExecutionClient interface { Start(context.Context) error Stop(context.Context) error StopAndCancel(context.Context) error } func (asyncTaskArgs) Kind() string { return "gateway_task_run" } type asyncTaskWorker struct { river.WorkerDefaults[asyncTaskArgs] service *Service } func (w *asyncTaskWorker) Work(ctx context.Context, job *river.Job[asyncTaskArgs]) error { task, err := w.service.store.GetTask(ctx, job.Args.TaskID) if err != nil { return err } if task.Status == "succeeded" || task.Status == "failed" || task.Status == "cancelled" { return nil } executionToken := uuid.NewString() result, runErr := w.service.executeWithToken(ctx, task, authUserFromTask(task), nil, executionToken) if runErr == nil { w.service.logger.Debug("river async task completed", "taskID", task.ID, "status", result.Task.Status, "riverJobID", job.ID) return nil } if errors.Is(runErr, store.ErrTaskExecutionLeaseUnavailable) { w.service.logger.Debug("river async task execution lease already held", "taskID", task.ID, "riverJobID", job.ID) return nil } if errors.Is(runErr, store.ErrTaskExecutionManualReview) { w.service.logger.Warn("river async task moved to manual review after ambiguous upstream submission", "taskID", task.ID, "riverJobID", job.ID) return nil } var queuedErr *TaskQueuedError if errors.As(runErr, &queuedErr) { return river.JobSnooze(queuedErr.Delay) } if ctx.Err() != nil { task.ExecutionToken = executionToken queued, queueErr := w.service.requeueInterruptedAsyncTask(context.WithoutCancel(ctx), task) if queueErr != nil { return queueErr } w.service.logger.Debug("river async task interrupted and requeued", "taskID", task.ID, "status", queued.Status, "riverJobID", job.ID) return river.JobSnooze(0) } w.service.logger.Warn("river async task completed with failure", "taskID", task.ID, "error", runErr, "riverJobID", job.ID) return nil } func (s *Service) StartAsyncQueueWorker(ctx context.Context) { if err := s.startRiverQueue(ctx); err != nil { s.logger.Error("start river async queue failed", "error", err) panic(err) } } func (s *Service) startRiverQueue(ctx context.Context) error { driver := riverpgxv5.New(s.store.Pool()) migrator, err := rivermigrate.New(driver, nil) if err != nil { return err } if _, err := migrator.Migrate(ctx, rivermigrate.DirectionUp, nil); err != nil { return err } controlClient, err := river.NewClient(driver, &river.Config{ ID: asyncWorkerID() + "-control", Logger: s.logger, TestOnly: s.cfg.AppEnv == "test", }) if err != nil { return err } snapshot, err := s.loadAsyncWorkerCapacity(ctx) if err != nil { return fmt.Errorf("calculate initial async worker capacity: %w", err) } executionClient, err := s.makeAsyncExecutionClient(snapshot.Capacity) if err != nil { return err } if err := executionClient.Start(ctx); err != nil { cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) _ = executionClient.StopAndCancel(cleanupCtx) cancel() return err } s.riverMu.Lock() s.riverControlClient = controlClient s.riverExecutionClient = executionClient s.riverWorkerCapacity = snapshot.Capacity s.riverDrainingClients = make(map[asyncExecutionClient]struct{}) s.riverMu.Unlock() s.observeAsyncWorkerCapacity(snapshot) s.logger.Info("async worker capacity initialized", "capacity", snapshot.Capacity, "desiredCapacity", snapshot.Desired, "hardLimit", snapshot.HardLimit, "enabledModels", snapshot.EnabledModels, "unlimitedModels", snapshot.UnlimitedModels, "modelDesired", snapshot.ModelDesired, "enabledGroups", snapshot.EnabledGroups, "unlimitedGroups", snapshot.UnlimitedGroups, "groupDesired", snapshot.GroupDesired, "capped", snapshot.Capped, ) if err := s.recoverAsyncRiverJobs(ctx); err != nil { cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) _ = executionClient.StopAndCancel(cleanupCtx) cancel() return err } go s.refreshAsyncWorkerCapacity(ctx) go s.stopAsyncWorkersOnShutdown(ctx) return nil } func (s *Service) newRiverAsyncExecutionClient(capacity int) (*river.Client[pgx.Tx], error) { workers := river.NewWorkers() if err := river.AddWorkerSafely(workers, &asyncTaskWorker{service: s}); err != nil { return nil, err } return river.NewClient(riverpgxv5.New(s.store.Pool()), &river.Config{ ID: fmt.Sprintf("%s-exec-%d-%d", asyncWorkerID(), capacity, time.Now().UnixNano()), JobTimeout: -1, Logger: s.logger, CompletedJobRetentionPeriod: 24 * time.Hour, Queues: map[string]river.QueueConfig{ asyncTaskQueueName: {MaxWorkers: capacity}, }, // Provider-backed media jobs commonly poll for 10-20 minutes. River may // execute a still-running job again once this window elapses, so keep the // rescue horizon above the longest configured provider poll timeout. RescueStuckJobsAfter: time.Hour, TestOnly: s.cfg.AppEnv == "test", Workers: workers, }) } func (s *Service) makeAsyncExecutionClient(capacity int) (asyncExecutionClient, error) { if s.asyncClientFactory != nil { return s.asyncClientFactory(capacity) } return s.newRiverAsyncExecutionClient(capacity) } func (s *Service) loadAsyncWorkerCapacity(ctx context.Context) (store.AsyncWorkerCapacitySnapshot, error) { if s.asyncCapacityLoader != nil { return s.asyncCapacityLoader(ctx, s.cfg.AsyncWorkerHardLimit) } return s.store.AsyncWorkerCapacity(ctx, s.cfg.AsyncWorkerHardLimit) } func (s *Service) refreshAsyncWorkerCapacity(ctx context.Context) { ticker := time.NewTicker(time.Duration(s.cfg.AsyncWorkerRefreshIntervalSeconds) * time.Second) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-ticker.C: s.resizeAsyncWorkerCapacity(ctx) } } } func (s *Service) resizeAsyncWorkerCapacity(ctx context.Context) { snapshot, err := s.loadAsyncWorkerCapacity(ctx) if err != nil { s.observeAsyncWorkerResize("refresh_failed") s.logger.Warn("refresh async worker capacity failed; keeping current client", "error", err) return } s.riverMu.RLock() currentCapacity := s.riverWorkerCapacity s.riverMu.RUnlock() s.observeAsyncWorkerCapacity(snapshot) if snapshot.Capacity == currentCapacity { return } newClient, err := s.makeAsyncExecutionClient(snapshot.Capacity) if err != nil { s.observeAsyncWorkerResize("create_failed") s.logger.Warn("create replacement async worker client failed; keeping current client", "error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity) return } if err := newClient.Start(ctx); err != nil { cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) _ = newClient.StopAndCancel(cleanupCtx) cancel() s.observeAsyncWorkerResize("start_failed") s.logger.Warn("start replacement async worker client failed; keeping current client", "error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity) return } if ctx.Err() != nil { cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) _ = newClient.StopAndCancel(cleanupCtx) cancel() return } s.riverMu.Lock() oldClient := s.riverExecutionClient s.riverExecutionClient = newClient s.riverWorkerCapacity = snapshot.Capacity if oldClient != nil { if s.riverDrainingClients == nil { s.riverDrainingClients = make(map[asyncExecutionClient]struct{}) } s.riverDrainingClients[oldClient] = struct{}{} } s.riverMu.Unlock() s.observeAsyncWorkerResize("success") s.logger.Info("async worker capacity resized", "previousCapacity", currentCapacity, "capacity", snapshot.Capacity, "desiredCapacity", snapshot.Desired, "hardLimit", snapshot.HardLimit, "modelDesired", snapshot.ModelDesired, "groupDesired", snapshot.GroupDesired, "capped", snapshot.Capped, ) if oldClient != nil { go s.drainAsyncWorkerClient(oldClient) } } func (s *Service) drainAsyncWorkerClient(client asyncExecutionClient) { stopCtx, cancel := context.WithTimeout(context.Background(), time.Hour) defer cancel() if err := client.Stop(stopCtx); err != nil { s.logger.Warn("gracefully drain previous async worker client failed", "error", err) } s.riverMu.Lock() delete(s.riverDrainingClients, client) s.riverMu.Unlock() } func (s *Service) stopAsyncWorkersOnShutdown(ctx context.Context) { <-ctx.Done() s.riverMu.Lock() clients := make([]asyncExecutionClient, 0, 1+len(s.riverDrainingClients)) if s.riverExecutionClient != nil { clients = append(clients, s.riverExecutionClient) } for client := range s.riverDrainingClients { clients = append(clients, client) } s.riverExecutionClient = nil s.riverDrainingClients = nil s.riverMu.Unlock() for _, client := range clients { stopCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) if err := client.StopAndCancel(stopCtx); err != nil { s.logger.Warn("stop river async queue failed", "error", err) } cancel() } } func (s *Service) asyncControlClient() *river.Client[pgx.Tx] { s.riverMu.RLock() defer s.riverMu.RUnlock() return s.riverControlClient } func (s *Service) observeAsyncWorkerCapacity(snapshot store.AsyncWorkerCapacitySnapshot) { observer, ok := s.billingMetrics.(interface { SetAsyncWorkerCapacity(current, desired, hardLimit int, capped bool) }) if ok { observer.SetAsyncWorkerCapacity(snapshot.Capacity, snapshot.Desired, snapshot.HardLimit, snapshot.Capped) } } func (s *Service) observeAsyncWorkerResize(outcome string) { observer, ok := s.billingMetrics.(interface { ObserveAsyncWorkerResize(string) }) if ok { observer.ObserveAsyncWorkerResize(outcome) } } func (s *Service) EnqueueAsyncTask(ctx context.Context, task store.GatewayTask) error { riverClient := s.asyncControlClient() if riverClient == nil { return errors.New("river async queue is not started") } result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task)) if err != nil { return err } if result.Job != nil { return s.store.SetTaskRiverJobID(ctx, task.ID, result.Job.ID) } return nil } func (s *Service) WakeAsyncQueueAfter(ctx context.Context, delay time.Duration) { } func (s *Service) RunAsyncTask(ctx context.Context, task store.GatewayTask, user *auth.User) { if err := s.EnqueueAsyncTask(ctx, task); err != nil { s.logger.Warn("enqueue river async task failed", "taskID", task.ID, "error", err) } } func (s *Service) recoverAsyncRiverJobs(ctx context.Context) error { riverClient := s.asyncControlClient() if riverClient == nil { return errors.New("river async queue is not started") } items, err := s.store.ListRecoverableAsyncTasks(ctx, 1000) if err != nil { return err } for _, item := range items { result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item, time.Now())) if err != nil { return err } if result.Job != nil { if err := s.store.SetTaskRiverJobID(ctx, item.ID, result.Job.ID); err != nil { return err } } } if len(items) > 0 { s.logger.Info("river async queue recovered persisted tasks", "count", len(items)) } return nil } func asyncTaskInsertOpts(task store.GatewayTask) *river.InsertOpts { priority := 2 if task.ID == "" { priority = 3 } return &river.InsertOpts{ MaxAttempts: 1000, Priority: priority, Queue: asyncTaskQueueName, Tags: []string{"gateway-task"}, UniqueOpts: river.UniqueOpts{ ByArgs: true, ByQueue: true, ByState: []rivertype.JobState{ rivertype.JobStateAvailable, rivertype.JobStatePending, rivertype.JobStateRetryable, rivertype.JobStateRunning, rivertype.JobStateScheduled, }, }, } } func asyncTaskRecoveryInsertOpts(item store.AsyncTaskQueueItem, now time.Time) *river.InsertOpts { opts := asyncTaskInsertOpts(store.GatewayTask{ID: item.ID}) if item.NextRunAt.After(now) { opts.ScheduledAt = item.NextRunAt } // A replacement process must not be blocked by a River row that the dead // process left in running state. PostgreSQL execution leases still ensure // that only one recovery job can call the upstream provider. opts.UniqueOpts = river.UniqueOpts{} return opts } func authUserFromTask(task store.GatewayTask) *auth.User { roles := []string{"user"} if strings.TrimSpace(task.UserID) == "" { roles = nil } return &auth.User{ ID: firstNonEmptyString(task.GatewayUserID, task.UserID), Roles: roles, TenantID: task.TenantID, GatewayTenantID: task.GatewayTenantID, TenantKey: task.TenantKey, Source: firstNonEmptyString(task.UserSource, "gateway"), GatewayUserID: task.GatewayUserID, UserGroupID: task.UserGroupID, UserGroupKey: task.UserGroupKey, APIKeyID: task.APIKeyID, APIKeyName: task.APIKeyName, APIKeyPrefix: task.APIKeyPrefix, } } func asyncWorkerID() string { host, _ := os.Hostname() host = strings.TrimSpace(host) if host == "" { host = "localhost" } return fmt.Sprintf("%s:%d:%d", host, os.Getpid(), time.Now().UnixNano()) } var _ river.Worker[asyncTaskArgs] = (*asyncTaskWorker)(nil)