package runner import ( "context" "errors" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/jackc/pgx/v5" ) const ( taskCompletionPollInterval = 250 * time.Millisecond taskCompletionFallbackInterval = 30 * time.Second taskCompletionBatchSize = 2000 ) func (s *Service) StartTaskCompletionWaiter(ctx context.Context) { s.taskCompletionOnce.Do(func() { go s.pollTaskCompletions(ctx) }) } func (s *Service) WaitForTaskCompletion(ctx context.Context, taskID string) (store.GatewayTask, error) { wake, unregister := s.registerTaskCompletionWaiter(taskID) defer unregister() s.signalTaskCompletionPoll() fallback := time.NewTicker(taskCompletionFallbackInterval) defer fallback.Stop() for { select { case <-ctx.Done(): return store.GatewayTask{}, ctx.Err() case <-wake: case <-fallback.C: } task, err := s.store.GetTask(ctx, taskID) if err != nil { if errors.Is(err, pgx.ErrNoRows) { return store.GatewayTask{}, err } if ctx.Err() != nil { return store.GatewayTask{}, ctx.Err() } continue } if terminalTaskStatus(task.Status) { return task, nil } } } func (s *Service) pollTaskCompletions(ctx context.Context) { ticker := time.NewTicker(taskCompletionPollInterval) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-ticker.C: case <-s.taskCompletionPollWake: } taskIDs := s.taskCompletionWaiterIDs(taskCompletionBatchSize) if len(taskIDs) == 0 { continue } statuses, err := s.store.ListTaskStatuses(ctx, taskIDs) if err != nil { if ctx.Err() == nil && s.logger != nil { s.logger.Warn("batch task completion poll failed", "taskCount", len(taskIDs), "error", err) } continue } for taskID, status := range statuses { if terminalTaskStatus(status) { s.signalTaskCompletion(taskID) } } } } func (s *Service) registerTaskCompletionWaiter(taskID string) (<-chan struct{}, func()) { wake := make(chan struct{}, 1) s.taskCompletionMu.Lock() if s.taskCompletionWaiters[taskID] == nil { s.taskCompletionWaiters[taskID] = map[chan struct{}]struct{}{} } s.taskCompletionWaiters[taskID][wake] = struct{}{} s.taskCompletionMu.Unlock() return wake, func() { s.taskCompletionMu.Lock() delete(s.taskCompletionWaiters[taskID], wake) if len(s.taskCompletionWaiters[taskID]) == 0 { delete(s.taskCompletionWaiters, taskID) } s.taskCompletionMu.Unlock() } } func (s *Service) taskCompletionWaiterIDs(limit int) []string { s.taskCompletionMu.Lock() defer s.taskCompletionMu.Unlock() taskIDs := make([]string, 0, min(limit, len(s.taskCompletionWaiters))) for taskID := range s.taskCompletionWaiters { taskIDs = append(taskIDs, taskID) if len(taskIDs) >= limit { break } } return taskIDs } func (s *Service) signalTaskCompletionPoll() { select { case s.taskCompletionPollWake <- struct{}{}: default: } } func (s *Service) signalTaskCompletion(taskID string) { s.taskCompletionMu.Lock() waiters := make([]chan struct{}, 0, len(s.taskCompletionWaiters[taskID])) for wake := range s.taskCompletionWaiters[taskID] { waiters = append(waiters, wake) } s.taskCompletionMu.Unlock() for _, wake := range waiters { select { case wake <- struct{}{}: default: } } } func terminalTaskStatus(status string) bool { switch status { case "succeeded", "failed", "cancelled", "manual_review": return true default: return false } }