Files
easyai-ai-gateway/apps/api/internal/runner/admission.go
T
wangbo fa818e9ebf fix(queue): 重新调度终态 River 任务
等待异步准入的任务只跳过仍处于活跃状态的 River Job;missing 或 terminal Job 交由准入调度器重新创建,避免非空 river_job_id 永久阻断旧任务。\n\n通用 River 恢复排除等待准入任务,消除双恢复循环与误导日志;事务回滚继续使用独立有界上下文。\n\n验证:真实 PostgreSQL 下异步 Worker 准入验收和跨 Store 队列集成测试均通过。
2026-07-30 17:42:33 +08:00

633 lines
18 KiB
Go

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" {
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")
}
}
}