fix(worker): 将 Gemini 同步请求交由异步队列执行
将非流式 Gemini generateContent 改为 River Worker 执行,API 通过批量状态查询等待完成,避免高并发同步请求耗尽 API 数据库连接池。\n\n生成图片在持久化时保存带哈希的资产引用,响应恢复时下载并校验后重建 Base64,数据库不保存媒体原文。请求体物化增加前置内存门禁,并扩展双 API/Worker PostgreSQL 压力测试。\n\n验证:go test ./... -count=1;go vet ./...;pnpm openapi;迁移安全检查;64 与 256 请求双角色异步 Gemini 压力测试。
This commit is contained in:
@@ -0,0 +1,143 @@
|
||||
package runner
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
const (
|
||||
taskCompletionPollInterval = 250 * time.Millisecond
|
||||
taskCompletionFallbackInterval = 30 * time.Second
|
||||
taskCompletionBatchSize = 2000
|
||||
)
|
||||
|
||||
func (s *Service) StartTaskCompletionWaiter(ctx context.Context) {
|
||||
s.taskCompletionOnce.Do(func() {
|
||||
go s.pollTaskCompletions(ctx)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) WaitForTaskCompletion(ctx context.Context, taskID string) (store.GatewayTask, error) {
|
||||
wake, unregister := s.registerTaskCompletionWaiter(taskID)
|
||||
defer unregister()
|
||||
s.signalTaskCompletionPoll()
|
||||
|
||||
fallback := time.NewTicker(taskCompletionFallbackInterval)
|
||||
defer fallback.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return store.GatewayTask{}, ctx.Err()
|
||||
case <-wake:
|
||||
case <-fallback.C:
|
||||
}
|
||||
task, err := s.store.GetTask(ctx, taskID)
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return store.GatewayTask{}, err
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return store.GatewayTask{}, ctx.Err()
|
||||
}
|
||||
continue
|
||||
}
|
||||
if terminalTaskStatus(task.Status) {
|
||||
return task, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) pollTaskCompletions(ctx context.Context) {
|
||||
ticker := time.NewTicker(taskCompletionPollInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
case <-s.taskCompletionPollWake:
|
||||
}
|
||||
taskIDs := s.taskCompletionWaiterIDs(taskCompletionBatchSize)
|
||||
if len(taskIDs) == 0 {
|
||||
continue
|
||||
}
|
||||
statuses, err := s.store.ListTaskStatuses(ctx, taskIDs)
|
||||
if err != nil {
|
||||
if ctx.Err() == nil && s.logger != nil {
|
||||
s.logger.Warn("batch task completion poll failed", "taskCount", len(taskIDs), "error", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
for taskID, status := range statuses {
|
||||
if terminalTaskStatus(status) {
|
||||
s.signalTaskCompletion(taskID)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) registerTaskCompletionWaiter(taskID string) (<-chan struct{}, func()) {
|
||||
wake := make(chan struct{}, 1)
|
||||
s.taskCompletionMu.Lock()
|
||||
if s.taskCompletionWaiters[taskID] == nil {
|
||||
s.taskCompletionWaiters[taskID] = map[chan struct{}]struct{}{}
|
||||
}
|
||||
s.taskCompletionWaiters[taskID][wake] = struct{}{}
|
||||
s.taskCompletionMu.Unlock()
|
||||
return wake, func() {
|
||||
s.taskCompletionMu.Lock()
|
||||
delete(s.taskCompletionWaiters[taskID], wake)
|
||||
if len(s.taskCompletionWaiters[taskID]) == 0 {
|
||||
delete(s.taskCompletionWaiters, taskID)
|
||||
}
|
||||
s.taskCompletionMu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) taskCompletionWaiterIDs(limit int) []string {
|
||||
s.taskCompletionMu.Lock()
|
||||
defer s.taskCompletionMu.Unlock()
|
||||
taskIDs := make([]string, 0, min(limit, len(s.taskCompletionWaiters)))
|
||||
for taskID := range s.taskCompletionWaiters {
|
||||
taskIDs = append(taskIDs, taskID)
|
||||
if len(taskIDs) >= limit {
|
||||
break
|
||||
}
|
||||
}
|
||||
return taskIDs
|
||||
}
|
||||
|
||||
func (s *Service) signalTaskCompletionPoll() {
|
||||
select {
|
||||
case s.taskCompletionPollWake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) signalTaskCompletion(taskID string) {
|
||||
s.taskCompletionMu.Lock()
|
||||
waiters := make([]chan struct{}, 0, len(s.taskCompletionWaiters[taskID]))
|
||||
for wake := range s.taskCompletionWaiters[taskID] {
|
||||
waiters = append(waiters, wake)
|
||||
}
|
||||
s.taskCompletionMu.Unlock()
|
||||
for _, wake := range waiters {
|
||||
select {
|
||||
case wake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func terminalTaskStatus(status string) bool {
|
||||
switch status {
|
||||
case "succeeded", "failed", "cancelled", "manual_review":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user