package runner import ( "context" "errors" "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 } 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) 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) 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) tryTaskAdmissionWithAdmittedHook( ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string, onAdmitted func(pgx.Tx) error, ) (store.TaskAdmissionResult, error) { mode := "sync" if task.AsyncMode { mode = "async" } input := 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, } 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 (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) { ticker := time.NewTicker(time.Second) defer ticker.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 <-s.currentAdmissionWake(): case <-ticker.C: } } } 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) 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 func() { for ctx.Err() == nil { err := s.store.ListenTaskAdmissionNotifications(ctx, func(string) { s.broadcastAdmissionWake() }) 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) currentAdmissionWake() <-chan struct{} { s.admissionWakeMu.Lock() defer s.admissionWakeMu.Unlock() return s.admissionWake } func (s *Service) broadcastAdmissionWake() { s.admissionWakeMu.Lock() close(s.admissionWake) s.admissionWake = make(chan struct{}) s.admissionWakeMu.Unlock() } 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 } result, err := s.tryTaskAdmissionWithAdmittedHook(ctx, task, plan, "", func(tx pgx.Tx) error { return s.enqueueAsyncTaskTx(ctx, tx, task.ID, asyncTaskInsertOpts(task)) }) 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 !result.Admitted { if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } return nil } if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) { return ctx.Err() } return nil } 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.currentAdmissionWake(): case <-ticker.C: } taskIDs, err := s.store.ListWaitingAsyncAdmissionTaskIDs(ctx, 100) 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.SubmitAsyncTask(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") } } }