perf(queue): 限制大媒体任务内存并消除准入惊群
将同步与异步非文本任务的准入唤醒改为按任务和全局队首推进,批量续租等待者,避免大量请求同时争抢 advisory lock 和数据库连接。 对 Base64 请求解析、素材解码、上游媒体执行和结果物化增加分层并发限制,复用请求素材并释放重复 Gemini wire 数据;生产 API 默认入口解析并发 16,媒体物化并发 8。 新增 1000 个同步 Gemini 图像编辑请求的模拟上游压力验收,10 MiB 输入和输出下全部成功,Heap 峰值增长约 2.58 GiB,并验证共享上传哈希、单次 attempt 与本地零落盘。 验证:go test ./... -count=1;go vet ./...;gofmt -l 无输出;kubectl kustomize deploy/kubernetes/production;10 MiB Gemini Base64 千任务压力测试通过。
This commit is contained in:
@@ -3,6 +3,7 @@ package runner
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"hash/fnv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -22,6 +23,17 @@ type taskAdmissionPlan struct {
|
||||
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",
|
||||
@@ -99,6 +111,12 @@ func (s *Service) buildTaskAdmissionPlan(ctx context.Context, task store.Gateway
|
||||
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 {
|
||||
@@ -128,10 +146,67 @@ func (s *Service) buildTaskAdmissionPlan(ctx context.Context, task store.Gateway
|
||||
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,
|
||||
@@ -139,21 +214,7 @@ func (s *Service) tryTaskAdmissionWithAdmittedHook(
|
||||
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,
|
||||
}
|
||||
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
|
||||
@@ -186,14 +247,34 @@ func (s *Service) tryTaskAdmissionWithAdmittedHook(
|
||||
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) {
|
||||
ticker := time.NewTicker(time.Second)
|
||||
defer ticker.Stop()
|
||||
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 {
|
||||
@@ -204,12 +285,19 @@ func (s *Service) waitForTaskAdmission(ctx context.Context, task store.GatewayTa
|
||||
_ = s.store.DeleteTaskAdmission(context.WithoutCancel(ctx), task.ID)
|
||||
s.observeTaskAdmission("cancelled")
|
||||
return store.TaskAdmissionResult{}, ctx.Err()
|
||||
case <-s.currentAdmissionWake():
|
||||
case <-ticker.C:
|
||||
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,
|
||||
@@ -218,6 +306,13 @@ func (s *Service) ensureCandidateAdmission(
|
||||
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 {
|
||||
@@ -270,10 +365,12 @@ func (s *Service) observeTaskAdmissionWait(wait time.Duration) {
|
||||
|
||||
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(string) {
|
||||
s.broadcastAdmissionWake()
|
||||
err := s.store.ListenTaskAdmissionNotifications(ctx, func(taskID string) {
|
||||
s.signalAdmissionWake(taskID)
|
||||
})
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
@@ -293,17 +390,99 @@ func (s *Service) StartAdmissionNotifier(ctx context.Context) {
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) currentAdmissionWake() <-chan struct{} {
|
||||
s.admissionWakeMu.Lock()
|
||||
defer s.admissionWakeMu.Unlock()
|
||||
return s.admissionWake
|
||||
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) broadcastAdmissionWake() {
|
||||
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()
|
||||
close(s.admissionWake)
|
||||
s.admissionWake = make(chan struct{})
|
||||
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 {
|
||||
@@ -335,28 +514,49 @@ func (s *Service) SubmitAsyncTask(ctx context.Context, task store.GatewayTask) e
|
||||
}
|
||||
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 {
|
||||
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 !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) 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
|
||||
@@ -380,10 +580,10 @@ func (s *Service) dispatchWaitingAsyncAdmissions(ctx context.Context) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-s.currentAdmissionWake():
|
||||
case <-s.asyncAdmissionWake:
|
||||
case <-ticker.C:
|
||||
}
|
||||
taskIDs, err := s.store.ListWaitingAsyncAdmissionTaskIDs(ctx, 100)
|
||||
taskIDs, err := s.store.ListWaitingAsyncAdmissionTaskIDs(ctx, 1000)
|
||||
if err != nil {
|
||||
s.logger.Warn("list waiting async admissions failed", "error", err)
|
||||
continue
|
||||
@@ -399,7 +599,7 @@ func (s *Service) dispatchWaitingAsyncAdmissions(ctx context.Context) {
|
||||
if task.Status != "queued" || task.RiverJobID > 0 {
|
||||
continue
|
||||
}
|
||||
if err := s.SubmitAsyncTask(ctx, task); err != nil {
|
||||
if err := s.dispatchWaitingAsyncTask(ctx, task); err != nil {
|
||||
s.logger.Warn("dispatch waiting async admission failed", "taskID", taskID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,11 @@ type httpClientCache struct {
|
||||
|
||||
const providerHTTPClientTimeout = 10 * time.Minute
|
||||
|
||||
const (
|
||||
providerHTTPMaxIdleConnections = 2048
|
||||
providerHTTPMaxIdleConnectionsPerHost = 1024
|
||||
)
|
||||
|
||||
func newHTTPClientCache() *httpClientCache {
|
||||
return &httpClientCache{
|
||||
none: newHTTPClient(nil),
|
||||
@@ -68,6 +73,12 @@ func (c *httpClientCache) customClient(rawProxy string) (*http.Client, error) {
|
||||
func newHTTPClient(proxy func(*http.Request) (*url.URL, error)) *http.Client {
|
||||
transport := http.DefaultTransport.(*http.Transport).Clone()
|
||||
transport.Proxy = proxy
|
||||
// Media workers may keep hundreds of slow upstream tasks in flight and
|
||||
// poll the same provider repeatedly. The standard library only retains two
|
||||
// idle connections per host, which turns a high-concurrency polling load
|
||||
// into avoidable TCP/TLS handshakes.
|
||||
transport.MaxIdleConns = providerHTTPMaxIdleConnections
|
||||
transport.MaxIdleConnsPerHost = providerHTTPMaxIdleConnectionsPerHost
|
||||
return &http.Client{
|
||||
Timeout: providerHTTPClientTimeout,
|
||||
Transport: transport,
|
||||
|
||||
@@ -16,6 +16,20 @@ func TestProviderHTTPClientTimeoutAllowsLongRunningMediaRequests(t *testing.T) {
|
||||
if client.Timeout != 10*time.Minute {
|
||||
t.Fatalf("unexpected provider HTTP timeout: got %s want %s", client.Timeout, 10*time.Minute)
|
||||
}
|
||||
transport, ok := client.Transport.(*http.Transport)
|
||||
if !ok {
|
||||
t.Fatalf("provider transport type = %T, want *http.Transport", client.Transport)
|
||||
}
|
||||
if transport.MaxIdleConns != providerHTTPMaxIdleConnections ||
|
||||
transport.MaxIdleConnsPerHost != providerHTTPMaxIdleConnectionsPerHost {
|
||||
t.Fatalf(
|
||||
"provider idle connection pool = %d/%d, want %d/%d",
|
||||
transport.MaxIdleConns,
|
||||
transport.MaxIdleConnsPerHost,
|
||||
providerHTTPMaxIdleConnections,
|
||||
providerHTTPMaxIdleConnectionsPerHost,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlatformProxyModeNoneIgnoresEnvironmentProxy(t *testing.T) {
|
||||
|
||||
@@ -40,8 +40,11 @@ type Service struct {
|
||||
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
|
||||
admissionWakeMu sync.Mutex
|
||||
admissionWake chan struct{}
|
||||
asyncAdmissionWake chan struct{}
|
||||
admissionTaskWaiters map[string]*admissionTaskWaiter
|
||||
admissionListener sync.Once
|
||||
asyncClientFactory func(int) (asyncExecutionClient, error)
|
||||
mediaResultSlots chan struct{}
|
||||
billingMetrics billingMetricsObserver
|
||||
}
|
||||
|
||||
@@ -93,6 +96,9 @@ func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...b
|
||||
if cfg.AsyncWorkerRefreshIntervalSeconds == 0 {
|
||||
cfg.AsyncWorkerRefreshIntervalSeconds = 5
|
||||
}
|
||||
if cfg.MediaMaterializationConcurrency == 0 {
|
||||
cfg.MediaMaterializationConcurrency = 8
|
||||
}
|
||||
if cfg.TaskProgressCallbackTimeoutMS == 0 {
|
||||
cfg.TaskProgressCallbackTimeoutMS = 5000
|
||||
}
|
||||
@@ -139,9 +145,12 @@ func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...b
|
||||
"universal": clients.UniversalClient{HTTPClient: httpClients.none, ScriptExecutor: scriptExecutor},
|
||||
"simulation": clients.SimulationClient{},
|
||||
},
|
||||
httpClients: httpClients,
|
||||
workerInstanceID: asyncWorkerID(),
|
||||
admissionWake: make(chan struct{}),
|
||||
httpClients: httpClients,
|
||||
workerInstanceID: asyncWorkerID(),
|
||||
admissionWake: make(chan struct{}, 4096),
|
||||
asyncAdmissionWake: make(chan struct{}, 1),
|
||||
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
|
||||
mediaResultSlots: make(chan struct{}, cfg.MediaMaterializationConcurrency),
|
||||
}
|
||||
if len(observers) > 0 {
|
||||
service.billingMetrics = observers[0]
|
||||
@@ -539,7 +548,11 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask,
|
||||
var admissionResult store.TaskAdmissionResult
|
||||
var admissionErr error
|
||||
if task.AsyncMode {
|
||||
admissionResult, admissionErr = s.tryTaskAdmission(ctx, task, plan, "")
|
||||
var alreadyAdmitted bool
|
||||
admissionResult, alreadyAdmitted, admissionErr = s.activeAsyncTaskAdmission(ctx, task, plan)
|
||||
if admissionErr == nil && !alreadyAdmitted {
|
||||
admissionResult, admissionErr = s.tryTaskAdmission(ctx, task, plan, "")
|
||||
}
|
||||
if admissionErr == nil && !admissionResult.Admitted {
|
||||
_ = s.store.ReleaseTaskPreparation(context.WithoutCancel(ctx), task.ID, task.ExecutionToken)
|
||||
return Result{Task: task, Output: task.Result}, &TaskQueuedError{Delay: time.Second}
|
||||
@@ -1272,6 +1285,15 @@ func (s *Service) runCandidate(
|
||||
})
|
||||
return clients.Response{}, err
|
||||
}
|
||||
mediaSlotPreacquired := mediaTaskNeedsPreUpstreamMaterializationSlot(candidate.ModelType)
|
||||
releaseMediaSlot := func() {}
|
||||
if mediaSlotPreacquired {
|
||||
releaseMediaSlot, err = s.acquireMediaMaterializationSlot(ctx)
|
||||
if err != nil {
|
||||
return clients.Response{}, err
|
||||
}
|
||||
defer releaseMediaSlot()
|
||||
}
|
||||
providerBody, err = s.hydrateProviderRequestAssets(ctx, providerBody, candidate)
|
||||
if err != nil {
|
||||
_ = s.store.FinishTaskAttempt(ctx, store.FinishTaskAttemptInput{
|
||||
@@ -1465,6 +1487,21 @@ func (s *Service) runCandidate(
|
||||
}
|
||||
return clients.Response{}, err
|
||||
}
|
||||
if task.AsyncMode {
|
||||
// Async callers consume the canonical task result through the task API,
|
||||
// never the original provider wire response. Release the raw JSON before
|
||||
// decoding/uploading inline media so a large Base64 response is not held
|
||||
// twice during the most memory-intensive phase.
|
||||
response.Wire = nil
|
||||
submissionWire = nil
|
||||
}
|
||||
if !mediaSlotPreacquired {
|
||||
releaseMediaSlot, err = s.acquireMediaResultSlot(ctx, response.Result)
|
||||
if err != nil {
|
||||
return clients.Response{}, err
|
||||
}
|
||||
defer releaseMediaSlot()
|
||||
}
|
||||
uploadedResult, err := s.uploadGeneratedAssets(ctx, task.ID, task.Kind, response.Result)
|
||||
if err != nil {
|
||||
metrics := mergeMetrics(taskMetrics(task, user, body, candidate, response, simulated), parameterPreprocessingMetrics(preprocessing), map[string]any{
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
||||
@@ -77,6 +78,40 @@ func defaultGeneratedAssetUploadPolicy() generatedAssetUploadPolicy {
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) acquireMediaMaterializationSlot(ctx context.Context) (func(), error) {
|
||||
if s.mediaResultSlots == nil {
|
||||
return func() {}, nil
|
||||
}
|
||||
select {
|
||||
case s.mediaResultSlots <- struct{}{}:
|
||||
var once sync.Once
|
||||
return func() {
|
||||
once.Do(func() {
|
||||
<-s.mediaResultSlots
|
||||
})
|
||||
}, nil
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) acquireMediaResultSlot(ctx context.Context, result map[string]any) (func(), error) {
|
||||
if s.mediaResultSlots == nil ||
|
||||
(!TaskResultHasInlineBinary(result) && !generatedRawValueHasInlineMedia(result["raw"], "", nil)) {
|
||||
return func() {}, nil
|
||||
}
|
||||
return s.acquireMediaMaterializationSlot(ctx)
|
||||
}
|
||||
|
||||
func mediaTaskNeedsPreUpstreamMaterializationSlot(modelType string) bool {
|
||||
switch canonicalModelType(modelType) {
|
||||
case "image_generate", "image_edit", "audio_generate", "text_to_speech", "voice_clone", "omni":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) uploadGeneratedAssets(ctx context.Context, taskID string, taskKind string, result map[string]any) (map[string]any, error) {
|
||||
data, _ := result["data"].([]any)
|
||||
rawNeedsUpload := generatedRawValueHasInlineMedia(result["raw"], "", nil)
|
||||
|
||||
@@ -53,6 +53,60 @@ func TestGeneratedAssetDecisionUploadsInlineImageBase64(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMediaResultMaterializationConcurrencyIsBounded(t *testing.T) {
|
||||
service := New(config.Config{MediaMaterializationConcurrency: 1}, nil, nil)
|
||||
result := map[string]any{
|
||||
"data": []any{map[string]any{
|
||||
"b64_json": base64.StdEncoding.EncodeToString([]byte("inline image")),
|
||||
"mime_type": "image/png",
|
||||
}},
|
||||
}
|
||||
release, err := service.acquireMediaResultSlot(context.Background(), result)
|
||||
if err != nil {
|
||||
t.Fatalf("acquire first media result slot: %v", err)
|
||||
}
|
||||
|
||||
waitCtx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond)
|
||||
defer cancel()
|
||||
if _, err := service.acquireMediaResultSlot(waitCtx, result); err == nil {
|
||||
t.Fatal("second media result materialization bypassed configured concurrency")
|
||||
}
|
||||
release()
|
||||
|
||||
release, err = service.acquireMediaResultSlot(context.Background(), result)
|
||||
if err != nil {
|
||||
t.Fatalf("acquire released media result slot: %v", err)
|
||||
}
|
||||
release()
|
||||
}
|
||||
|
||||
func TestMediaTaskPreacquiresMaterializationSlotForInlineResultTypes(t *testing.T) {
|
||||
t.Parallel()
|
||||
for _, modelType := range []string{
|
||||
"image_generate",
|
||||
"image_edit",
|
||||
"audio_generate",
|
||||
"text_to_speech",
|
||||
"voice_clone",
|
||||
"omni",
|
||||
} {
|
||||
if !mediaTaskNeedsPreUpstreamMaterializationSlot(modelType) {
|
||||
t.Fatalf("%s should preacquire the media materialization slot", modelType)
|
||||
}
|
||||
}
|
||||
for _, modelType := range []string{
|
||||
"text_generate",
|
||||
"text_embedding",
|
||||
"text_rerank",
|
||||
"video_generate",
|
||||
"image_to_video",
|
||||
} {
|
||||
if mediaTaskNeedsPreUpstreamMaterializationSlot(modelType) {
|
||||
t.Fatalf("%s should not hold a media materialization slot across the upstream call", modelType)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeneratedAssetDecisionUploadsVectorDocumentWithoutChangingType(t *testing.T) {
|
||||
item := map[string]any{
|
||||
"type": "file",
|
||||
|
||||
Reference in New Issue
Block a user