将 River 执行容量按平台模型与用户组的有效并发策略动态调整,并为长任务续租、并发租约原子抢占和限流退避补充保护。\n\n统一平台模型限流继承语义,兼容历史 platformLimits/modelLimits,并为三个迁移图像模型建立独立并发租约。补充管理端显式继承/覆盖配置、指标、单元测试及隔离 PostgreSQL 验收。\n\n验证:go test ./... -count=1;pnpm lint;pnpm test;pnpm build;隔离 PostgreSQL 并发原子性/续租测试;128 任务与三分钟长任务动态 Worker 验收。
434 lines
13 KiB
Go
434 lines
13 KiB
Go
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)
|