feat(worker): 实现集群限流与自适应负载

保留平台模型 RPM、TPM 和并发策略语义,增加 PostgreSQL 集群级租约、饱和候选重选和多平台自动负载,避免突发任务固定等待首个平台。\n\n新增 Worker 实时负载采样、自适应 active/heavy 容量、心跳与管理端指标,并扩展本地 acceptance runner,覆盖三 Worker、同模型三平台 2/4/6 并发和 48 个带图视频突发任务。\n\n验证:go test ./...、go vet ./...、PostgreSQL 跨 Store 集成测试、gofmt、bash -n、ShellCheck 及本地集群 provider-burst 验收通过;48/48 成功,无越限、重复提交、重复计费、重复回调或终态资源泄漏。
This commit is contained in:
2026-08-03 00:13:46 +08:00
parent 9a01fd4657
commit c28bf74230
52 changed files with 3700 additions and 272 deletions
+43 -7
View File
@@ -19,6 +19,7 @@ import (
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
scriptengine "github.com/easyai/easyai-ai-gateway/apps/api/internal/script"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/workerload"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/riverqueue/river"
@@ -39,6 +40,8 @@ type Service struct {
riverDrainingClients map[asyncExecutionClient]struct{}
riverWorkerCapacity int
workerInstanceID string
workerLoad *workerload.Controller
workerLoadSampler workerLoadSampler
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
admissionWakeMu sync.Mutex
admissionWake chan struct{}
@@ -129,6 +132,9 @@ func NewWithStores(
if cfg.AsyncWorkerRefreshIntervalSeconds == 0 {
cfg.AsyncWorkerRefreshIntervalSeconds = 5
}
if strings.TrimSpace(cfg.AsyncWorkerLoadMode) == "" {
cfg.AsyncWorkerLoadMode = workerload.ModeAdaptive
}
if cfg.MediaMaterializationConcurrency == 0 {
cfg.MediaMaterializationConcurrency = 8
}
@@ -184,8 +190,13 @@ func NewWithStores(
"universal": clients.UniversalClient{HTTPClient: httpClients.none, ScriptExecutor: scriptExecutor},
"simulation": clients.SimulationClient{},
},
httpClients: httpClients,
workerInstanceID: asyncWorkerID(),
httpClients: httpClients,
workerInstanceID: asyncWorkerID(),
workerLoad: workerload.New(workerload.Config{
Mode: cfg.AsyncWorkerLoadMode, HardLimit: cfg.AsyncWorkerInstanceHardLimit,
InitialActive: 4, InitialHeavy: 1, HealthySamples: 3,
}),
workerLoadSampler: workerload.NewSystemSampler(),
admissionWake: make(chan struct{}, 4096),
asyncAdmissionWake: make(chan struct{}, 1),
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
@@ -1266,7 +1277,7 @@ func (s *Service) runCandidate(
}
simulated := isSimulation(task, candidate)
baseAttemptMetrics := mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), parameterPreprocessingMetrics(preprocessing))
reservations := s.rateLimitReservations(ctx, user, candidate, body)
reservations := acceptanceInfrastructureReservations(task, s.rateLimitReservations(ctx, user, candidate, body))
if admittedPlatformModelID == candidate.PlatformModelID && len(admittedLeases) > 0 {
filtered := make([]store.RateLimitReservation, 0, len(reservations))
for _, reservation := range reservations {
@@ -1278,6 +1289,10 @@ func (s *Service) runCandidate(
}
limitResult, err := s.store.ReserveRateLimits(ctx, task.ID, "", reservations)
if err != nil {
var limitErr *store.RateLimitExceededError
if errors.As(err, &limitErr) {
s.observeProviderQuotaWait(limitErr.Metric)
}
retryable := store.RateLimitRetryable(err)
clientErr := &clients.ClientError{Code: "rate_limit", Message: err.Error(), Retryable: retryable}
return clients.Response{}, &localRateLimitError{clientErr: clientErr, cause: err, retryAfter: localRateLimitRetryAfter(err)}
@@ -1438,6 +1453,9 @@ func (s *Service) runCandidate(
); err != nil {
return clients.Response{}, fmt.Errorf("restore upstream submission status: %w", err)
}
if err := enterWorkerWaiting(ctx); err != nil {
return clients.Response{}, err
}
}
setSubmissionStatus := func(status string) error {
if submissionStatus == "response_received" && status != "response_received" {
@@ -1484,7 +1502,10 @@ func (s *Service) runCandidate(
if err := s.persistCompatibilitySubmission(context.WithoutCancel(ctx), task, candidate, remoteTaskID, checkpoint, submissionWire); err != nil {
return err
}
return setSubmissionStatus("response_received")
if err := setSubmissionStatus("response_received"); err != nil {
return err
}
return enterWorkerWaiting(ctx)
},
OnRemoteTaskPolled: func(remoteTaskID string, payload map[string]any) error {
if strings.TrimSpace(remoteTaskID) == "" {
@@ -1496,18 +1517,21 @@ func (s *Service) runCandidate(
}
task.RemoteTaskID = remoteTaskID
task.RemoteTaskPayload = checkpoint
return setSubmissionStatus("response_received")
if err := setSubmissionStatus("response_received"); err != nil {
return err
}
return enterWorkerWaiting(ctx)
},
OnUpstreamSubmissionStarted: func() error {
if err := setSubmissionStatus("submitting"); err != nil {
return err
}
markUpstreamSubmissionStarted(ctx)
return nil
return enterWorkerWaiting(ctx)
},
OnUpstreamResponseReceived: func() error {
submissionStatus = "response_received"
return nil
return enterWorkerFinalizing(ctx)
},
OnUpstreamWireResponse: func(wire *clients.WireResponse) error {
submissionWire = wire
@@ -1521,6 +1545,9 @@ func (s *Service) runCandidate(
UpstreamPreviousResponseID: responseExecution.UpstreamPreviousResponseID,
PreviousResponseTurns: responseExecution.PreviousTurns,
})
if phaseErr := enterWorkerFinalizing(runCtx); err == nil && phaseErr != nil {
err = phaseErr
}
if leaseErr := stopLeaseRenewal(); leaseErr != nil {
err = &clients.ClientError{
Code: "concurrency_lease_lost",
@@ -1717,6 +1744,15 @@ func (s *Service) runCandidate(
return response, nil
}
func (s *Service) observeProviderQuotaWait(metric string) {
observer, ok := s.billingMetrics.(interface {
ObserveProviderQuotaWait(string)
})
if ok {
observer.ObserveProviderQuotaWait(metric)
}
}
func minimalRemoteTaskCheckpoint(provider string, specType string, payload map[string]any) map[string]any {
const maxBytes = 8192
provider = strings.ToLower(strings.TrimSpace(provider))