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
@@ -17,6 +17,7 @@ import (
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
const (
@@ -405,6 +406,17 @@ func (s *Service) hydrateLocalBinaryValue(ctx context.Context, taskID string, va
}
switch typed := value.(type) {
case map[string]any:
if ref, ok := generatedResultAssetReference(typed); ok {
payload, contentType, err := s.readGeneratedResultAsset(ctx, ref)
if err != nil {
return nil, false, err
}
encoded := base64.StdEncoding.EncodeToString(payload)
if generatedResultAssetUsesDataURI(typed) {
return "data:" + contentType + ";base64," + encoded, true, nil
}
return encoded, true, nil
}
next := make(map[string]any, len(typed))
changed := false
for key, childValue := range typed {
@@ -453,6 +465,62 @@ func (s *Service) hydrateLocalBinaryValue(ctx context.Context, taskID string, va
}
}
func generatedResultAssetReference(value map[string]any) (store.RequestAsset, bool) {
ref, ok := value["assetRef"].(map[string]any)
if !ok {
return store.RequestAsset{}, false
}
storage, _ := value["assetStorage"].(map[string]any)
if stringFromAny(storage["scene"]) != store.FileStorageSceneImageResult {
return store.RequestAsset{}, false
}
asset := store.RequestAsset{
SHA256: strings.ToLower(strings.TrimSpace(stringFromAny(ref["sha256"]))),
ContentType: firstNonEmptyString(stringFromAny(ref["contentType"]), stringFromAny(storage["contentType"])),
URL: firstNonEmptyString(stringFromAny(ref["url"]), stringFromAny(value["url"])),
StorageProvider: stringFromAny(ref["storageProvider"]),
}
if size := floatFromAny(ref["size"]); size > 0 {
asset.ByteSize = int64(size)
}
if expiresAt := stringFromAny(ref["expiresAt"]); expiresAt != "" {
if parsed, err := time.Parse(time.RFC3339, expiresAt); err == nil {
asset.ExpiresAt = &parsed
}
}
if asset.URL == "" || asset.SHA256 == "" || asset.ByteSize <= 0 {
return store.RequestAsset{}, false
}
return asset, true
}
func generatedResultAssetUsesDataURI(value map[string]any) bool {
storage, _ := value["assetStorage"].(map[string]any)
source := normalizeLocalBinaryKey(stringFromAny(storage["source"]))
return source == "datauri"
}
func (s *Service) readGeneratedResultAsset(ctx context.Context, asset store.RequestAsset) ([]byte, string, error) {
payload, err := s.readRequestAssetBytes(ctx, asset)
if err != nil {
return nil, "", err
}
digest := sha256.Sum256(payload)
if int64(len(payload)) != asset.ByteSize || hex.EncodeToString(digest[:]) != asset.SHA256 {
return nil, "", &clients.ClientError{
Code: "binary_result_corrupted",
Message: "stored result asset failed size or hash verification",
StatusCode: 500,
Retryable: false,
}
}
contentType := strings.TrimSpace(asset.ContentType)
if contentType == "" {
contentType = "application/octet-stream"
}
return payload, contentType, nil
}
func (s *Service) readLocalBinaryResult(taskID string, descriptor localBinaryDescriptor) ([]byte, error) {
path := filepath.Join(s.localBinaryResultRoot(), safeLocalBinaryTaskDir(taskID), descriptor.SHA256+".bin")
info, err := os.Stat(path)
@@ -2,9 +2,13 @@ package runner
import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
@@ -112,6 +116,45 @@ func TestMaterializeAndHydrateLocalBinaryResult(t *testing.T) {
}
}
func TestHydrateGeneratedResultAssetReference(t *testing.T) {
payload := []byte("verified generated image bytes")
digest := sha256.Sum256(payload)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "image/png")
_, _ = w.Write(payload)
}))
defer server.Close()
reference := func(hash string) map[string]any {
return map[string]any{
"assetRef": map[string]any{
"sha256": hash,
"contentType": "image/png",
"size": len(payload),
"url": server.URL + "/generated.png",
},
"assetStorage": map[string]any{
"scene": store.FileStorageSceneImageResult,
"source": "b64_json",
},
}
}
service := &Service{}
result := map[string]any{"data": []any{map[string]any{"b64_json": reference(hex.EncodeToString(digest[:]))}}}
hydrated, err := service.HydrateTaskResult(context.Background(), "task-remote", result)
if err != nil {
t.Fatalf("hydrate generated result asset: %v", err)
}
item := hydrated["data"].([]any)[0].(map[string]any)
if got, want := item["b64_json"], base64.StdEncoding.EncodeToString(payload); got != want {
t.Fatalf("hydrated Base64=%v, want %v", got, want)
}
result = map[string]any{"data": []any{map[string]any{"b64_json": reference(strings.Repeat("0", 64))}}}
_, err = service.HydrateTaskResult(context.Background(), "task-corrupted", result)
assertClientErrorCode(t, err, "binary_result_corrupted")
}
func TestHydrateLocalBinaryResultReturnsExpiredAndCorruptedErrors(t *testing.T) {
service := newLocalBinaryTestService(t)
service.cfg.LocalResultTTLHours = 1
+35 -29
View File
@@ -25,28 +25,32 @@ import (
)
type Service struct {
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
workerInstanceID string
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{}
directOSS *directOSSUploader
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
workerInstanceID string
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{}
directOSS *directOSSUploader
taskCompletionOnce sync.Once
taskCompletionMu sync.Mutex
taskCompletionWaiters map[string]map[chan struct{}]struct{}
taskCompletionPollWake chan struct{}
billingMetrics billingMetricsObserver
}
type billingMetricsObserver interface {
@@ -149,13 +153,15 @@ 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{}, 4096),
asyncAdmissionWake: make(chan struct{}, 1),
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
mediaResultSlots: make(chan struct{}, cfg.MediaMaterializationConcurrency),
directOSS: newDirectOSSUploader(cfg),
httpClients: httpClients,
workerInstanceID: asyncWorkerID(),
admissionWake: make(chan struct{}, 4096),
asyncAdmissionWake: make(chan struct{}, 1),
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
mediaResultSlots: make(chan struct{}, cfg.MediaMaterializationConcurrency),
directOSS: newDirectOSSUploader(cfg),
taskCompletionWaiters: map[string]map[chan struct{}]struct{}{},
taskCompletionPollWake: make(chan struct{}, 1),
}
if len(observers) > 0 {
service.billingMetrics = observers[0]
+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
}
}
@@ -0,0 +1,50 @@
package runner
import (
"testing"
"time"
)
func TestTaskCompletionWaitersAreBatchedAndSignaled(t *testing.T) {
service := &Service{
taskCompletionWaiters: map[string]map[chan struct{}]struct{}{},
taskCompletionPollWake: make(chan struct{}, 1),
}
first, unregisterFirst := service.registerTaskCompletionWaiter("task-1")
defer unregisterFirst()
second, unregisterSecond := service.registerTaskCompletionWaiter("task-2")
taskIDs := service.taskCompletionWaiterIDs(10)
if len(taskIDs) != 2 {
t.Fatalf("batched task IDs=%v", taskIDs)
}
service.signalTaskCompletion("task-1")
select {
case <-first:
case <-time.After(time.Second):
t.Fatal("task completion waiter was not signaled")
}
select {
case <-second:
t.Fatal("unrelated task waiter was signaled")
default:
}
unregisterSecond()
if got := service.taskCompletionWaiterIDs(10); len(got) != 1 || got[0] != "task-1" {
t.Fatalf("waiter unregister left unexpected IDs=%v", got)
}
}
func TestTerminalTaskStatus(t *testing.T) {
for _, status := range []string{"succeeded", "failed", "cancelled", "manual_review"} {
if !terminalTaskStatus(status) {
t.Fatalf("status %q should be terminal", status)
}
}
for _, status := range []string{"queued", "running", "pending"} {
if terminalTaskStatus(status) {
t.Fatalf("status %q should not be terminal", status)
}
}
}
+3
View File
@@ -236,6 +236,9 @@ func (s *Service) uploadGeneratedAssets(ctx context.Context, taskID string, task
if contentType != "" && stringFromAny(merged["mime_type"]) == "" {
merged["mime_type"] = contentType
}
if decision.Inline != nil && strings.TrimSpace(sourceKey) != "" {
merged[sourceKey] = generatedRawMediaReference(decision.Inline, upload, contentType, kind, strategy)
}
}
nextData = append(nextData, merged)
}