package runner import ( "context" "errors" "hash/fnv" "strings" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" "github.com/easyai/easyai-ai-gateway/apps/api/internal/clients" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/google/uuid" "github.com/jackc/pgx/v5" ) type taskAdmissionPlan struct { Candidate store.RuntimeModelCandidate Body map[string]any ModelType string GroupID string Scopes []store.AdmissionScope Eligible bool } type admissionTaskWaiter struct { wake chan struct{} waiterID string } const ( asyncWorkerCapacityScopeKey = "global" asyncWorkerQueueLimit = 10000 asyncWorkerMaxWaitSeconds = 24 * 60 * 60 ) func distributedAdmissionModelType(modelType string) bool { switch strings.ToLower(strings.TrimSpace(modelType)) { case "image_generate", "image_edit", "image_analysis", "image_vectorize", "video_generate", "video_enhance", "image_to_video", "text_to_video", "video_edit", "video_reference", "video_first_last_frame", "video_understanding", "omni_video", "omni", "audio_generate", "audio_understanding", "text_to_speech", "voice_clone": return true default: return false } } func (s *Service) buildTaskAdmissionPlan(ctx context.Context, task store.GatewayTask, user *auth.User) (taskAdmissionPlan, error) { restoredRequest, err := s.restoreTaskRequestReferences(ctx, task) if err != nil { return taskAdmissionPlan{}, err } body := normalizeRequest(task.Kind, restoredRequest) modelType := modelTypeFromKind(task.Kind, body) if !distributedAdmissionModelType(modelType) { return taskAdmissionPlan{Body: body, ModelType: modelType}, nil } if err := validateRequest(task.Kind, body); err != nil { return taskAdmissionPlan{}, parameterPreprocessClientError(err) } body, clonedVoice, err := s.resolveClonedVoiceBinding(ctx, user, task.Kind, body) if err != nil { return taskAdmissionPlan{}, err } runnerPolicy, err := s.store.GetActiveRunnerPolicy(ctx) if err != nil { return taskAdmissionPlan{}, err } cacheAffinityKeys := buildCacheAffinityKeys(task.Kind, modelType, body) candidates, err := s.store.ListModelCandidates(ctx, task.Model, modelType, user, store.ListModelCandidatesOptions{ CacheAffinityKey: cacheAffinityKeys.Primary, CacheAffinityKeys: cacheAffinityKeys.Lookup, CacheAffinityPolicy: runnerPolicy.CacheAffinityPolicy, }) if err == nil { candidates, err = filterCandidatesByRequestedPlatform(candidates, body) } if err == nil { candidates, err = filterCandidatesByClonedVoiceBinding(candidates, clonedVoice) } if err == nil { candidates, _, err = filterRuntimeCandidatesByRequest(task.Kind, task.Model, modelType, body, candidates) } if err == nil { candidates, _, err = filterRuntimeCandidatesByOutputTokens(task.Kind, task.Model, modelType, body, candidates) } if err != nil { return taskAdmissionPlan{}, err } for _, candidate := range candidates { available, availabilityErr := s.store.RuntimeCandidateAvailable(ctx, candidate.PlatformID, candidate.PlatformModelID) if availabilityErr != nil { return taskAdmissionPlan{}, availabilityErr } if !available { continue } scopes, groupID := s.admissionScopes(ctx, user, candidate) if task.AsyncMode { scopes, err = s.withAsyncWorkerCapacityScope(ctx, scopes) if err != nil { return taskAdmissionPlan{}, err } } hasConcurrentLimit := false for _, scope := range scopes { if scope.ConcurrentLimit > 0 { hasConcurrentLimit = true break } } if !hasConcurrentLimit { return taskAdmissionPlan{ Candidate: candidate, Body: body, ModelType: modelType, }, nil } if err := s.store.CheckRateLimits(ctx, s.rateLimitReservations(ctx, user, candidate, body)); err != nil { return taskAdmissionPlan{}, err } return taskAdmissionPlan{ Candidate: candidate, Body: body, ModelType: modelType, GroupID: groupID, Scopes: scopes, Eligible: true, }, nil } return taskAdmissionPlan{}, store.ErrNoModelCandidate } func (s *Service) withAsyncWorkerCapacityScope(ctx context.Context, scopes []store.AdmissionScope) ([]store.AdmissionScope, error) { capacity, err := s.store.ActiveWorkerCapacity(ctx) if err != nil { return nil, err } if capacity <= 0 { return scopes, nil } out := append([]store.AdmissionScope{}, scopes...) out = append(out, store.AdmissionScope{ ScopeType: "worker_capacity", ScopeKey: asyncWorkerCapacityScopeKey, ScopeName: "async worker capacity", ConcurrentLimit: float64(capacity), Amount: 1, LeaseTTLSeconds: 120, QueueLimit: asyncWorkerQueueLimit, MaxWaitSeconds: asyncWorkerMaxWaitSeconds, }) return out, nil } func (s *Service) tryTaskAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string) (store.TaskAdmissionResult, error) { return s.tryTaskAdmissionWithAdmittedHook(ctx, task, plan, waiterID, nil) } func (s *Service) activeAsyncTaskAdmission( ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, ) (store.TaskAdmissionResult, bool, error) { if !task.AsyncMode { return store.TaskAdmissionResult{}, false, nil } admission, err := s.store.GetTaskAdmission(ctx, task.ID) if errors.Is(err, pgx.ErrNoRows) { return store.TaskAdmissionResult{}, false, nil } if err != nil { return store.TaskAdmissionResult{}, false, err } if admission.Status != "admitted" || admission.PlatformID != plan.Candidate.PlatformID || admission.PlatformModelID != plan.Candidate.PlatformModelID || admission.UserGroupID != plan.GroupID { return store.TaskAdmissionResult{}, false, nil } leases, err := s.store.ActiveTaskAdmissionLeases(ctx, task.ID) if err != nil { return store.TaskAdmissionResult{}, false, err } if len(leases) == 0 { return store.TaskAdmissionResult{}, false, nil } return store.TaskAdmissionResult{ Admission: admission, Admitted: true, Leases: leases, }, true, nil } func (s *Service) tryTaskAdmissionWithAdmittedHook( ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string, onAdmitted func(pgx.Tx) error, ) (store.TaskAdmissionResult, error) { input := taskAdmissionInput(task, plan, waiterID) current, currentErr := s.store.GetTaskAdmission(ctx, task.ID) if currentErr != nil && !errors.Is(currentErr, pgx.ErrNoRows) { return store.TaskAdmissionResult{}, currentErr } if currentErr == nil && current.Status == "waiting" && (current.PlatformID != input.PlatformID || current.PlatformModelID != input.PlatformModelID || current.UserGroupID != input.UserGroupID) { if _, rebindErr := s.store.RebindWaitingTaskAdmission(ctx, input); rebindErr != nil { return store.TaskAdmissionResult{}, rebindErr } s.observeTaskAdmission("candidate_migrated") } var result store.TaskAdmissionResult var err error if onAdmitted == nil { result, err = s.store.TryTaskAdmission(ctx, input) } else { result, err = s.store.TryTaskAdmissionWithAdmittedHook(ctx, input, onAdmitted) } if result.NewlyAdmitted { s.observeTaskAdmission("admitted") s.observeTaskAdmissionWait(time.Since(result.Admission.EnqueuedAt)) } var limitErr *store.RateLimitExceededError switch { case errors.As(err, &limitErr) && limitErr.Reason == "queue_full": s.observeTaskAdmission("queue_full") case errors.Is(err, store.ErrQueueTimeout): s.observeTaskAdmission("timeout") } return result, err } func taskAdmissionInput(task store.GatewayTask, plan taskAdmissionPlan, waiterID string) store.TaskAdmissionInput { mode := "sync" if task.AsyncMode { mode = "async" } return store.TaskAdmissionInput{ TaskID: task.ID, PlatformID: plan.Candidate.PlatformID, PlatformModelID: plan.Candidate.PlatformModelID, UserGroupID: plan.GroupID, QueueKey: plan.Candidate.QueueKey, Mode: mode, Priority: task.Priority, WaiterID: waiterID, Scopes: plan.Scopes, } } func (s *Service) waitForSynchronousAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan) (store.TaskAdmissionResult, error) { waiterID := uuid.NewString() return s.waitForTaskAdmission(ctx, task, plan, waiterID) } func (s *Service) waitForTaskAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string) (store.TaskAdmissionResult, error) { taskWake, unregister := s.registerAdmissionTaskWaiter(task.ID, waiterID) defer unregister() pollTimer := time.NewTimer(admissionPollInterval(task.ID)) defer pollTimer.Stop() for { result, err := s.tryTaskAdmission(ctx, task, plan, waiterID) if err != nil || result.Admitted { return result, err } select { case <-ctx.Done(): _ = s.store.DeleteTaskAdmission(context.WithoutCancel(ctx), task.ID) s.observeTaskAdmission("cancelled") return store.TaskAdmissionResult{}, ctx.Err() case <-taskWake: case <-pollTimer.C: pollTimer.Reset(admissionPollInterval(task.ID)) } } } func admissionPollInterval(taskID string) time.Duration { hasher := fnv.New32a() _, _ = hasher.Write([]byte(taskID)) return 30*time.Second + time.Duration(hasher.Sum32()%5000)*time.Millisecond } func (s *Service) ensureCandidateAdmission( ctx context.Context, task store.GatewayTask, user *auth.User, body map[string]any, candidate store.RuntimeModelCandidate, ) (store.TaskAdmissionResult, bool, error) { scopes, groupID := s.admissionScopes(ctx, user, candidate) if task.AsyncMode { var err error scopes, err = s.withAsyncWorkerCapacityScope(ctx, scopes) if err != nil { return store.TaskAdmissionResult{}, true, err } } hasConcurrentLimit := false for _, scope := range scopes { if scope.ConcurrentLimit > 0 { hasConcurrentLimit = true break } } if !hasConcurrentLimit { if err := s.store.DeleteTaskAdmission(context.WithoutCancel(ctx), task.ID); err != nil && !errors.Is(err, pgx.ErrNoRows) { return store.TaskAdmissionResult{}, false, err } return store.TaskAdmissionResult{}, false, nil } if err := s.store.CheckRateLimits(ctx, s.rateLimitReservations(ctx, user, candidate, body)); err != nil { return store.TaskAdmissionResult{}, true, err } plan := taskAdmissionPlan{ Candidate: candidate, Body: body, ModelType: candidate.ModelType, GroupID: groupID, Scopes: scopes, Eligible: true, } waiterID := "" if !task.AsyncMode { waiterID = uuid.NewString() } result, err := s.waitForTaskAdmission(ctx, task, plan, waiterID) return result, true, err } func (s *Service) observeTaskAdmission(event string) { observer, ok := s.billingMetrics.(interface { ObserveTaskAdmission(string) }) if ok { observer.ObserveTaskAdmission(event) } } func (s *Service) observeTaskAdmissionWait(wait time.Duration) { observer, ok := s.billingMetrics.(interface { ObserveTaskAdmissionWait(time.Duration) }) if ok { observer.ObserveTaskAdmissionWait(wait) } } func (s *Service) StartAdmissionNotifier(ctx context.Context) { s.admissionListener.Do(func() { go s.dispatchWaitingSynchronousAdmissions(ctx) go s.renewSynchronousAdmissionWaiters(ctx) go func() { for ctx.Err() == nil { err := s.store.ListenTaskAdmissionNotifications(ctx, func(taskID string) { s.signalAdmissionWake(taskID) }) if ctx.Err() != nil { return } if s.logger != nil { s.logger.Warn("task admission LISTEN connection interrupted; periodic polling remains active", "error", err) } timer := time.NewTimer(time.Second) select { case <-ctx.Done(): timer.Stop() return case <-timer.C: } } }() }) } func (s *Service) dispatchWaitingSynchronousAdmissions(ctx context.Context) { for { select { case <-ctx.Done(): return case <-s.admissionWake: } taskIDs, err := s.store.ListWaitingTaskAdmissionIDs(ctx, 1) if err != nil { if s.logger != nil { s.logger.Warn("list waiting synchronous admissions failed", "error", err) } continue } s.admissionWakeMu.Lock() waiters := make([]chan struct{}, 0, len(taskIDs)) for _, taskID := range taskIDs { if waiter := s.admissionTaskWaiters[taskID]; waiter != nil { waiters = append(waiters, waiter.wake) } } s.admissionWakeMu.Unlock() for _, taskWake := range waiters { select { case taskWake <- struct{}{}: default: } } } } func (s *Service) renewSynchronousAdmissionWaiters(ctx context.Context) { ticker := time.NewTicker(5 * time.Second) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-ticker.C: } s.admissionWakeMu.Lock() taskIDs := make([]string, 0, len(s.admissionTaskWaiters)) waiterIDs := make([]string, 0, len(s.admissionTaskWaiters)) for taskID, waiter := range s.admissionTaskWaiters { if waiter == nil || strings.TrimSpace(waiter.waiterID) == "" { continue } taskIDs = append(taskIDs, taskID) waiterIDs = append(waiterIDs, waiter.waiterID) } s.admissionWakeMu.Unlock() if err := s.store.RenewTaskAdmissionWaiters(ctx, taskIDs, waiterIDs); err != nil && s.logger != nil { s.logger.Warn("renew synchronous admission waiters failed", "waiterCount", len(taskIDs), "error", err) } } } func (s *Service) registerAdmissionTaskWaiter(taskID string, waiterID string) (<-chan struct{}, func()) { s.admissionWakeMu.Lock() waiter := &admissionTaskWaiter{wake: make(chan struct{}, 1), waiterID: waiterID} s.admissionTaskWaiters[taskID] = waiter s.admissionWakeMu.Unlock() return waiter.wake, func() { s.admissionWakeMu.Lock() if s.admissionTaskWaiters[taskID] == waiter { delete(s.admissionTaskWaiters, taskID) } s.admissionWakeMu.Unlock() } } func (s *Service) signalAdmissionWake(taskID string) { taskID = strings.TrimSpace(taskID) if taskID == "" || taskID == "*" { select { case s.admissionWake <- struct{}{}: default: } } else { s.admissionWakeMu.Lock() waiter := s.admissionTaskWaiters[taskID] s.admissionWakeMu.Unlock() if waiter != nil { select { case waiter.wake <- struct{}{}: default: } } } select { case s.asyncAdmissionWake <- struct{}{}: default: } } func (s *Service) SubmitAsyncTask(ctx context.Context, task store.GatewayTask) error { if !task.AsyncMode { return errors.New("only asynchronous tasks can be submitted to the async queue") } if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } user := authUserFromTask(task) plan, err := s.buildTaskAdmissionPlan(ctx, task, user) if err != nil { if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } _, _ = s.store.FailQueuedTask(context.WithoutCancel(ctx), task.ID, clients.ErrorCode(err), err.Error()) return err } if !plan.Eligible { if err := s.EnqueueAsyncTask(ctx, task); err != nil { if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } _, _ = s.store.FailQueuedTask(context.WithoutCancel(ctx), task.ID, "enqueue_failed", err.Error()) return err } if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } return nil } input := taskAdmissionInput(task, plan, "") current, currentErr := s.store.GetTaskAdmission(ctx, task.ID) if currentErr != nil && !errors.Is(currentErr, pgx.ErrNoRows) { return currentErr } if currentErr == nil && current.Status == "waiting" && (current.PlatformID != input.PlatformID || current.PlatformModelID != input.PlatformModelID || current.UserGroupID != input.UserGroupID) { if _, err := s.store.RebindWaitingTaskAdmission(ctx, input); err != nil { return err } s.observeTaskAdmission("candidate_migrated") } // Register FIFO order without creating a River job. A Worker dispatcher // first reserves both business concurrency and a live execution slot, then // creates the unique River job in the same transaction. if _, err := s.store.QueueTaskAdmissionWithHook(ctx, input, nil); err != nil { if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } _, _ = s.store.FailQueuedTask(context.WithoutCancel(ctx), task.ID, clients.ErrorCode(err), err.Error()) return err } if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } return nil } func (s *Service) dispatchWaitingAsyncTask(ctx context.Context, task store.GatewayTask) error { user := authUserFromTask(task) plan, err := s.buildTaskAdmissionPlan(ctx, task, user) if err != nil { return err } if !plan.Eligible { return s.EnqueueAsyncTask(ctx, task) } _, err = s.tryTaskAdmissionWithAdmittedHook(ctx, task, plan, "", func(tx pgx.Tx) error { return s.enqueueAsyncTaskTx(ctx, tx, task.ID, asyncTaskInsertOpts(task)) }) return err } func (s *Service) cancelAsyncSubmissionIfDisconnected(ctx context.Context, taskID string) bool { if ctx.Err() == nil { return false } cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second) defer cancel() if _, changed, err := s.store.CancelQueuedTask(cleanupCtx, taskID, "client disconnected before upstream submission"); err != nil { if s.logger != nil { s.logger.Warn("cancel disconnected asynchronous task failed", "taskID", taskID, "error", err) } } else if changed { s.observeTaskAdmission("cancelled") } return true } func (s *Service) dispatchWaitingAsyncAdmissions(ctx context.Context) { ticker := time.NewTicker(time.Second) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-s.asyncAdmissionWake: case <-ticker.C: } taskIDs, err := s.store.ListWaitingAsyncAdmissionTaskIDs(ctx, 1000) if err != nil { s.logger.Warn("list waiting async admissions failed", "error", err) continue } for _, taskID := range taskIDs { task, err := s.store.GetTask(ctx, taskID) if err != nil { if !errors.Is(err, pgx.ErrNoRows) { s.logger.Warn("load waiting async task failed", "taskID", taskID, "error", err) } continue } if task.Status != "queued" || task.RiverJobID > 0 { continue } if err := s.dispatchWaitingAsyncTask(ctx, task); err != nil { s.logger.Warn("dispatch waiting async admission failed", "taskID", taskID, "error", err) } } } } func (s *Service) reapExpiredTaskAdmissions(ctx context.Context) { ticker := time.NewTicker(5 * time.Second) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-ticker.C: } result, err := s.store.ReapExpiredTaskAdmissions(ctx, 500) if err != nil { if s.logger != nil { s.logger.Warn("reap expired task admissions failed", "error", err) } continue } for range result.ExpiredWaiters { s.observeTaskAdmission("expired") } for range result.ExpiredDeadlines { s.observeTaskAdmission("timeout") } } }