perf(queue): 按策略动态扩缩异步 Worker
将 River 执行容量按平台模型与用户组的有效并发策略动态调整,并为长任务续租、并发租约原子抢占和限流退避补充保护。\n\n统一平台模型限流继承语义,兼容历史 platformLimits/modelLimits,并为三个迁移图像模型建立独立并发租约。补充管理端显式继承/覆盖配置、指标、单元测试及隔离 PostgreSQL 验收。\n\n验证:go test ./... -count=1;pnpm lint;pnpm test;pnpm build;隔离 PostgreSQL 并发原子性/续租测试;128 任务与三分钟长任务动态 Worker 验收。
This commit is contained in:
@@ -74,7 +74,7 @@ func (s *Service) rateLimitReservations(ctx context.Context, user *auth.User, ca
|
||||
"groupKey": group.GroupKey,
|
||||
"name": group.Name,
|
||||
},
|
||||
group.RateLimitPolicy,
|
||||
store.NormalizeRateLimitPolicy(group.RateLimitPolicy),
|
||||
body,
|
||||
)...)
|
||||
}
|
||||
@@ -82,27 +82,15 @@ func (s *Service) rateLimitReservations(ctx context.Context, user *auth.User, ca
|
||||
}
|
||||
|
||||
func effectiveRateLimitPolicy(candidate store.RuntimeModelCandidate) map[string]any {
|
||||
policy := candidate.PlatformRateLimitPolicy
|
||||
if strings.TrimSpace(candidate.RuntimePolicySetID) != "" {
|
||||
policy = candidate.RuntimeRateLimitPolicy
|
||||
} else if hasRules(candidate.RuntimeRateLimitPolicy) {
|
||||
policy = mergeMap(policy, candidate.RuntimeRateLimitPolicy)
|
||||
}
|
||||
if _, hasOverride := candidate.RuntimePolicyOverride["rateLimitPolicy"]; hasOverride {
|
||||
nested, _ := candidate.RuntimePolicyOverride["rateLimitPolicy"].(map[string]any)
|
||||
if len(nested) == 0 {
|
||||
policy = nil
|
||||
} else {
|
||||
policy = mergeMap(policy, nested)
|
||||
}
|
||||
}
|
||||
if hasRules(candidate.ModelRateLimitPolicy) {
|
||||
policy = mergeMap(policy, candidate.ModelRateLimitPolicy)
|
||||
}
|
||||
if hasRules(policy) {
|
||||
return policy
|
||||
}
|
||||
return nil
|
||||
return store.EffectiveRateLimitPolicy(store.EffectiveRateLimitPolicyInput{
|
||||
BasePolicy: candidate.BaseRateLimitPolicy,
|
||||
PlatformPolicy: candidate.PlatformRateLimitPolicy,
|
||||
RuntimePolicy: candidate.RuntimeRateLimitPolicy,
|
||||
RuntimePolicyExplicit: candidate.RuntimePolicyExplicit,
|
||||
RuntimePolicyOverride: candidate.RateLimitRuntimeOverride,
|
||||
ModelPolicy: candidate.ModelRateLimitPolicy,
|
||||
ModelPolicyMode: candidate.ModelRateLimitPolicyMode,
|
||||
})
|
||||
}
|
||||
|
||||
func effectiveRetryPolicy(candidate store.RuntimeModelCandidate) map[string]any {
|
||||
|
||||
@@ -102,8 +102,10 @@ func TestEffectiveRateLimitPolicyTreatsEmptyRuntimePolicyAsUnlimited(t *testing.
|
||||
PlatformRateLimitPolicy: map[string]any{"rules": []any{
|
||||
map[string]any{"metric": "rpm", "limit": 500},
|
||||
}},
|
||||
RuntimePolicySetID: "runtime-policy-1",
|
||||
RuntimeRateLimitPolicy: map[string]any{"rules": []any{}},
|
||||
RuntimePolicySetID: "runtime-policy-1",
|
||||
RuntimePolicyExplicit: true,
|
||||
RuntimeRateLimitPolicy: map[string]any{"rules": []any{}},
|
||||
ModelRateLimitPolicyMode: "inherit",
|
||||
})
|
||||
|
||||
if hasRules(policy) {
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"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"
|
||||
@@ -23,6 +24,12 @@ 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 {
|
||||
@@ -87,21 +94,70 @@ func (s *Service) startRiverQueue(ctx context.Context) error {
|
||||
return err
|
||||
}
|
||||
|
||||
workers := river.NewWorkers()
|
||||
if err := river.AddWorkerSafely(workers, &asyncTaskWorker{service: s}); err != nil {
|
||||
controlClient, err := river.NewClient(driver, &river.Config{
|
||||
ID: asyncWorkerID() + "-control",
|
||||
Logger: s.logger,
|
||||
TestOnly: s.cfg.AppEnv == "test",
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
riverClient, err := river.NewClient(driver, &river.Config{
|
||||
ID: asyncWorkerID(),
|
||||
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{
|
||||
// Image providers may hold a worker while polling for several
|
||||
// minutes. Keep enough workers available for production bursts so
|
||||
// unrelated models do not remain queued behind long-running media
|
||||
// tasks.
|
||||
asyncTaskQueueName: {MaxWorkers: 96},
|
||||
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
|
||||
@@ -110,32 +166,160 @@ func (s *Service) startRiverQueue(ctx context.Context) error {
|
||||
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 {
|
||||
return err
|
||||
s.observeAsyncWorkerResize("refresh_failed")
|
||||
s.logger.Warn("refresh async worker capacity failed; keeping current client", "error", err)
|
||||
return
|
||||
}
|
||||
s.riverClient = riverClient
|
||||
if err := riverClient.Start(ctx); err != nil {
|
||||
return err
|
||||
s.riverMu.RLock()
|
||||
currentCapacity := s.riverWorkerCapacity
|
||||
s.riverMu.RUnlock()
|
||||
s.observeAsyncWorkerCapacity(snapshot)
|
||||
if snapshot.Capacity == currentCapacity {
|
||||
return
|
||||
}
|
||||
if err := s.recoverAsyncRiverJobs(ctx); err != nil {
|
||||
return err
|
||||
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
|
||||
}
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
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)
|
||||
defer cancel()
|
||||
if err := riverClient.StopAndCancel(stopCtx); err != nil {
|
||||
if err := client.StopAndCancel(stopCtx); err != nil {
|
||||
s.logger.Warn("stop river async queue failed", "error", err)
|
||||
}
|
||||
}()
|
||||
return nil
|
||||
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 {
|
||||
if s.riverClient == nil {
|
||||
riverClient := s.asyncControlClient()
|
||||
if riverClient == nil {
|
||||
return errors.New("river async queue is not started")
|
||||
}
|
||||
result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task))
|
||||
result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -155,12 +339,16 @@ func (s *Service) RunAsyncTask(ctx context.Context, task store.GatewayTask, user
|
||||
}
|
||||
|
||||
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 := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item, time.Now()))
|
||||
result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item, time.Now()))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
package runner
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||
)
|
||||
|
||||
type fakeAsyncExecutionClient struct {
|
||||
startErr error
|
||||
stopGate <-chan struct{}
|
||||
started atomic.Int64
|
||||
stopped atomic.Int64
|
||||
stopCancelled atomic.Int64
|
||||
}
|
||||
|
||||
func (c *fakeAsyncExecutionClient) Start(context.Context) error {
|
||||
c.started.Add(1)
|
||||
return c.startErr
|
||||
}
|
||||
|
||||
func (c *fakeAsyncExecutionClient) Stop(ctx context.Context) error {
|
||||
c.stopped.Add(1)
|
||||
if c.stopGate == nil {
|
||||
return nil
|
||||
}
|
||||
select {
|
||||
case <-c.stopGate:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *fakeAsyncExecutionClient) StopAndCancel(context.Context) error {
|
||||
c.stopCancelled.Add(1)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestResizeAsyncWorkerCapacityStartsReplacementBeforeGracefulDrain(t *testing.T) {
|
||||
oldClient := &fakeAsyncExecutionClient{}
|
||||
drainGate := make(chan struct{})
|
||||
oldClient.stopGate = drainGate
|
||||
newClient := &fakeAsyncExecutionClient{}
|
||||
service := asyncWorkerManagerTestService(1)
|
||||
service.riverExecutionClient = oldClient
|
||||
service.asyncCapacityLoader = fixedAsyncCapacity(3)
|
||||
service.asyncClientFactory = func(capacity int) (asyncExecutionClient, error) {
|
||||
if capacity != 3 {
|
||||
t.Fatalf("factory capacity=%d, want=3", capacity)
|
||||
}
|
||||
return newClient, nil
|
||||
}
|
||||
|
||||
service.resizeAsyncWorkerCapacity(context.Background())
|
||||
if newClient.started.Load() != 1 {
|
||||
t.Fatalf("replacement start count=%d, want=1", newClient.started.Load())
|
||||
}
|
||||
if service.riverExecutionClient != newClient || service.riverWorkerCapacity != 3 {
|
||||
t.Fatalf("replacement was not installed: capacity=%d client=%T", service.riverWorkerCapacity, service.riverExecutionClient)
|
||||
}
|
||||
waitForAtomicValue(t, &oldClient.stopped, 1)
|
||||
if oldClient.stopCancelled.Load() != 0 {
|
||||
t.Fatal("graceful drain cancelled an already running old task")
|
||||
}
|
||||
close(drainGate)
|
||||
waitForDrainingClients(t, service, 0)
|
||||
}
|
||||
|
||||
func TestResizeAsyncWorkerCapacityFailuresKeepCurrentClient(t *testing.T) {
|
||||
t.Run("refresh failure", func(t *testing.T) {
|
||||
oldClient := &fakeAsyncExecutionClient{}
|
||||
service := asyncWorkerManagerTestService(4)
|
||||
service.riverExecutionClient = oldClient
|
||||
service.asyncCapacityLoader = func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||
return store.AsyncWorkerCapacitySnapshot{}, errors.New("database unavailable")
|
||||
}
|
||||
service.asyncClientFactory = func(int) (asyncExecutionClient, error) {
|
||||
t.Fatal("factory must not run after refresh failure")
|
||||
return nil, nil
|
||||
}
|
||||
service.resizeAsyncWorkerCapacity(context.Background())
|
||||
if service.riverExecutionClient != oldClient || service.riverWorkerCapacity != 4 {
|
||||
t.Fatal("refresh failure replaced the current client")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("create failure", func(t *testing.T) {
|
||||
oldClient := &fakeAsyncExecutionClient{}
|
||||
service := asyncWorkerManagerTestService(4)
|
||||
service.riverExecutionClient = oldClient
|
||||
service.asyncCapacityLoader = fixedAsyncCapacity(8)
|
||||
service.asyncClientFactory = func(int) (asyncExecutionClient, error) {
|
||||
return nil, errors.New("factory failed")
|
||||
}
|
||||
service.resizeAsyncWorkerCapacity(context.Background())
|
||||
if service.riverExecutionClient != oldClient || service.riverWorkerCapacity != 4 {
|
||||
t.Fatal("create failure replaced the current client")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("start failure", func(t *testing.T) {
|
||||
oldClient := &fakeAsyncExecutionClient{}
|
||||
failedClient := &fakeAsyncExecutionClient{startErr: errors.New("start failed")}
|
||||
service := asyncWorkerManagerTestService(4)
|
||||
service.riverExecutionClient = oldClient
|
||||
service.asyncCapacityLoader = fixedAsyncCapacity(8)
|
||||
service.asyncClientFactory = func(int) (asyncExecutionClient, error) {
|
||||
return failedClient, nil
|
||||
}
|
||||
service.resizeAsyncWorkerCapacity(context.Background())
|
||||
if service.riverExecutionClient != oldClient || service.riverWorkerCapacity != 4 {
|
||||
t.Fatal("start failure replaced the current client")
|
||||
}
|
||||
if failedClient.stopCancelled.Load() != 1 {
|
||||
t.Fatal("failed replacement client was not cleaned up")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestStopAsyncWorkersIncludesCurrentAndDrainingClients(t *testing.T) {
|
||||
current := &fakeAsyncExecutionClient{}
|
||||
draining := &fakeAsyncExecutionClient{}
|
||||
service := asyncWorkerManagerTestService(2)
|
||||
service.riverExecutionClient = current
|
||||
service.riverDrainingClients[draining] = struct{}{}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
service.stopAsyncWorkersOnShutdown(ctx)
|
||||
if current.stopCancelled.Load() != 1 || draining.stopCancelled.Load() != 1 {
|
||||
t.Fatalf("shutdown cancellations current=%d draining=%d, want 1/1", current.stopCancelled.Load(), draining.stopCancelled.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func asyncWorkerManagerTestService(capacity int) *Service {
|
||||
return &Service{
|
||||
cfg: config.Config{
|
||||
AsyncWorkerHardLimit: 2048,
|
||||
AsyncWorkerRefreshIntervalSeconds: 5,
|
||||
},
|
||||
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
||||
riverDrainingClients: make(map[asyncExecutionClient]struct{}),
|
||||
riverWorkerCapacity: capacity,
|
||||
riverExecutionClient: nil,
|
||||
}
|
||||
}
|
||||
|
||||
func fixedAsyncCapacity(capacity int) func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||
return func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||
return store.AsyncWorkerCapacitySnapshot{Capacity: capacity, Desired: capacity, HardLimit: 2048}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func waitForAtomicValue(t *testing.T, value *atomic.Int64, want int64) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if value.Load() == want {
|
||||
return
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
t.Fatalf("atomic value=%d, want=%d", value.Load(), want)
|
||||
}
|
||||
|
||||
func waitForDrainingClients(t *testing.T, service *Service, want int) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
service.riverMu.RLock()
|
||||
count := len(service.riverDrainingClients)
|
||||
service.riverMu.RUnlock()
|
||||
if count == want {
|
||||
return
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
service.riverMu.RLock()
|
||||
count := len(service.riverDrainingClients)
|
||||
service.riverMu.RUnlock()
|
||||
t.Fatalf("draining client count=%d, want=%d", count, want)
|
||||
}
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
@@ -23,14 +24,20 @@ import (
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
cfg config.Config
|
||||
store *store.Store
|
||||
logger *slog.Logger
|
||||
clients map[string]clients.Client
|
||||
scriptExecutor *scriptengine.Executor
|
||||
httpClients *httpClientCache
|
||||
riverClient *river.Client[pgx.Tx]
|
||||
billingMetrics billingMetricsObserver
|
||||
cfg config.Config
|
||||
store *store.Store
|
||||
logger *slog.Logger
|
||||
clients map[string]clients.Client
|
||||
scriptExecutor *scriptengine.Executor
|
||||
httpClients *httpClientCache
|
||||
riverMu sync.RWMutex
|
||||
riverControlClient *river.Client[pgx.Tx]
|
||||
riverExecutionClient asyncExecutionClient
|
||||
riverDrainingClients map[asyncExecutionClient]struct{}
|
||||
riverWorkerCapacity int
|
||||
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
|
||||
asyncClientFactory func(int) (asyncExecutionClient, error)
|
||||
billingMetrics billingMetricsObserver
|
||||
}
|
||||
|
||||
type billingMetricsObserver interface {
|
||||
@@ -75,6 +82,12 @@ func (e *TaskQueuedError) Is(target error) bool {
|
||||
}
|
||||
|
||||
func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...billingMetricsObserver) *Service {
|
||||
if cfg.AsyncWorkerHardLimit == 0 {
|
||||
cfg.AsyncWorkerHardLimit = 2048
|
||||
}
|
||||
if cfg.AsyncWorkerRefreshIntervalSeconds == 0 {
|
||||
cfg.AsyncWorkerRefreshIntervalSeconds = 5
|
||||
}
|
||||
httpClients := newHTTPClientCache()
|
||||
scriptExecutor := &scriptengine.Executor{Logger: logger}
|
||||
service := &Service{
|
||||
@@ -891,7 +904,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
_ = s.store.ReleaseRateLimitReservations(context.WithoutCancel(ctx), limitResult.Reservations, "attempt_failed")
|
||||
}
|
||||
}()
|
||||
defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.LeaseIDs)
|
||||
defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.Leases)
|
||||
|
||||
attemptID, err := s.store.CreateTaskAttempt(ctx, store.CreateTaskAttemptInput{
|
||||
TaskID: task.ID,
|
||||
@@ -1031,7 +1044,8 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
return nil
|
||||
}
|
||||
var submissionWire *clients.WireResponse
|
||||
response, err := client.Run(ctx, clients.Request{
|
||||
runCtx, stopLeaseRenewal := s.startConcurrencyLeaseRenewal(ctx, task.ID, limitResult.Leases)
|
||||
response, err := client.Run(runCtx, clients.Request{
|
||||
Kind: task.Kind,
|
||||
ModelType: candidate.ModelType,
|
||||
Model: task.Model,
|
||||
@@ -1083,6 +1097,13 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
||||
UpstreamPreviousResponseID: responseExecution.UpstreamPreviousResponseID,
|
||||
PreviousResponseTurns: responseExecution.PreviousTurns,
|
||||
})
|
||||
if leaseErr := stopLeaseRenewal(); leaseErr != nil {
|
||||
err = &clients.ClientError{
|
||||
Code: "concurrency_lease_lost",
|
||||
Message: leaseErr.Error(),
|
||||
Retryable: true,
|
||||
}
|
||||
}
|
||||
callFinishedAt := time.Now()
|
||||
if err == nil {
|
||||
if markErr := setSubmissionStatus("response_received"); markErr != nil {
|
||||
@@ -1500,6 +1521,7 @@ func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.Gateway
|
||||
if delay <= 0 {
|
||||
delay = 5 * time.Second
|
||||
}
|
||||
delay += taskRetryJitter(task.ID)
|
||||
queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, delay, candidate.QueueKey)
|
||||
if err != nil {
|
||||
return store.GatewayTask{}, 0, err
|
||||
@@ -1515,6 +1537,81 @@ func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.Gateway
|
||||
return queued, delay, nil
|
||||
}
|
||||
|
||||
func taskRetryJitter(taskID string) time.Duration {
|
||||
var sum uint32
|
||||
for _, value := range []byte(taskID) {
|
||||
sum = sum*33 + uint32(value)
|
||||
}
|
||||
return time.Duration(sum%251) * time.Millisecond
|
||||
}
|
||||
|
||||
func (s *Service) startConcurrencyLeaseRenewal(ctx context.Context, taskID string, leases []store.ConcurrencyLease) (context.Context, func() error) {
|
||||
if len(leases) == 0 {
|
||||
return ctx, func() error { return nil }
|
||||
}
|
||||
interval := 30 * time.Second
|
||||
for _, lease := range leases {
|
||||
ttl := lease.TTL
|
||||
if ttl <= 0 {
|
||||
ttl = 120 * time.Second
|
||||
}
|
||||
candidate := ttl / 3
|
||||
if candidate < time.Second {
|
||||
candidate = time.Second
|
||||
}
|
||||
if candidate < interval {
|
||||
interval = candidate
|
||||
}
|
||||
}
|
||||
runCtx, cancelRun := context.WithCancel(ctx)
|
||||
renewCtx, cancelRenew := context.WithCancel(ctx)
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-renewCtx.Done():
|
||||
done <- nil
|
||||
return
|
||||
case <-ticker.C:
|
||||
if err := s.store.RenewConcurrencyLeases(renewCtx, leases); err != nil {
|
||||
if renewCtx.Err() != nil {
|
||||
done <- nil
|
||||
return
|
||||
}
|
||||
outcome := "failure"
|
||||
if errors.Is(err, store.ErrConcurrencyLeaseLost) {
|
||||
outcome = "lost"
|
||||
}
|
||||
s.observeConcurrencyLeaseRenewal(outcome)
|
||||
s.logger.Error("concurrency lease renewal failed; cancelling upstream execution",
|
||||
"taskID", taskID, "leaseCount", len(leases), "outcome", outcome, "error", err)
|
||||
done <- err
|
||||
cancelRun()
|
||||
return
|
||||
}
|
||||
s.observeConcurrencyLeaseRenewal("success")
|
||||
}
|
||||
}
|
||||
}()
|
||||
return runCtx, func() error {
|
||||
cancelRenew()
|
||||
renewalErr := <-done
|
||||
cancelRun()
|
||||
return renewalErr
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) observeConcurrencyLeaseRenewal(outcome string) {
|
||||
observer, ok := s.billingMetrics.(interface {
|
||||
ObserveConcurrencyLeaseRenewal(string)
|
||||
})
|
||||
if ok {
|
||||
observer.ObserveConcurrencyLeaseRenewal(outcome)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) requeueInterruptedAsyncTask(ctx context.Context, task store.GatewayTask) (store.GatewayTask, error) {
|
||||
queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, 0, "")
|
||||
if err != nil {
|
||||
|
||||
@@ -58,10 +58,11 @@ func (s *Service) CancelTask(ctx context.Context, taskID string, user *auth.User
|
||||
return taskCancelUnavailable(task, "任务已开始执行,当前阶段不可取消,请继续查询结果"), nil
|
||||
}
|
||||
if task.RiverJobID > 0 {
|
||||
if s.riverClient == nil {
|
||||
riverClient := s.asyncControlClient()
|
||||
if riverClient == nil {
|
||||
return taskCancelUnavailable(task, "任务取消队列未就绪,请继续查询结果"), nil
|
||||
}
|
||||
job, err := s.riverClient.JobGet(ctx, task.RiverJobID)
|
||||
job, err := riverClient.JobGet(ctx, task.RiverJobID)
|
||||
if errors.Is(err, rivertype.ErrNotFound) {
|
||||
return taskCancelUnavailable(task, "任务已不在本地排队队列,可能已提交上游,当前不可取消,请继续查询结果"), nil
|
||||
}
|
||||
@@ -71,7 +72,7 @@ func (s *Service) CancelTask(ctx context.Context, taskID string, user *auth.User
|
||||
if job == nil || !riverJobStateCancellable(job.State) {
|
||||
return taskCancelUnavailable(task, "任务已不在可取消队列状态,请继续查询结果"), nil
|
||||
}
|
||||
if _, err := s.riverClient.JobDelete(ctx, task.RiverJobID); err != nil {
|
||||
if _, err := riverClient.JobDelete(ctx, task.RiverJobID); err != nil {
|
||||
if errors.Is(err, rivertype.ErrJobRunning) || errors.Is(err, rivertype.ErrNotFound) {
|
||||
return taskCancelUnavailable(task, "任务已被工作进程领取,当前不可取消,请继续查询结果"), nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user