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:
2026-07-31 05:49:45 +08:00
parent 352e23e099
commit 4af86b22ee
12 changed files with 516 additions and 48 deletions
+143
View File
@@ -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
}
}