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:
@@ -43,7 +43,7 @@ func TestGeminiBase64ImageEditThousandSynchronousRequests(t *testing.T) {
|
|||||||
mediaConcurrent = 8
|
mediaConcurrent = 8
|
||||||
)
|
)
|
||||||
baselineProfile := acceptanceworkload.GeminiBaseline
|
baselineProfile := acceptanceworkload.GeminiBaseline
|
||||||
taskCount := baselineProfile.Requests
|
taskCount := stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_REQUESTS", baselineProfile.Requests)
|
||||||
imageBytes := stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_IMAGE_BYTES", baselineProfile.InputBytes)
|
imageBytes := stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_IMAGE_BYTES", baselineProfile.InputBytes)
|
||||||
upstreamDelay := time.Duration(stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_UPSTREAM_DELAY_MS", int(baselineProfile.DelayMin/time.Millisecond))) * time.Millisecond
|
upstreamDelay := time.Duration(stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_UPSTREAM_DELAY_MS", int(baselineProfile.DelayMin/time.Millisecond))) * time.Millisecond
|
||||||
maxHeapGrowth := uint64(stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_MAX_HEAP_BYTES", 1<<30))
|
maxHeapGrowth := uint64(stressEnvInt(t, "AI_GATEWAY_GEMINI_STRESS_MAX_HEAP_BYTES", 1<<30))
|
||||||
@@ -59,6 +59,11 @@ func TestGeminiBase64ImageEditThousandSynchronousRequests(t *testing.T) {
|
|||||||
t.Fatalf("connect API store: %v", err)
|
t.Fatalf("connect API store: %v", err)
|
||||||
}
|
}
|
||||||
defer apiDB.Close()
|
defer apiDB.Close()
|
||||||
|
workerDB, err := store.ConnectWithMaxConns(ctx, databaseURL, 16)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect Worker store: %v", err)
|
||||||
|
}
|
||||||
|
defer workerDB.Close()
|
||||||
poolConfig, err := pgxpool.ParseConfig(databaseURL)
|
poolConfig, err := pgxpool.ParseConfig(databaseURL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("parse observation pool config: %v", err)
|
t.Fatalf("parse observation pool config: %v", err)
|
||||||
@@ -152,6 +157,25 @@ WHERE status = 'active'`); err != nil {
|
|||||||
serverCtx, cancelServer := context.WithCancel(ctx)
|
serverCtx, cancelServer := context.WithCancel(ctx)
|
||||||
defer cancelServer()
|
defer cancelServer()
|
||||||
storageRoot := t.TempDir()
|
storageRoot := t.TempDir()
|
||||||
|
workerConfig := config.Config{
|
||||||
|
AppEnv: "test",
|
||||||
|
HTTPAddr: ":0",
|
||||||
|
DatabaseURL: databaseURL,
|
||||||
|
DatabaseMaxConns: 16,
|
||||||
|
IdentityMode: "hybrid",
|
||||||
|
JWTSecret: "gemini-base64-stress-secret",
|
||||||
|
BillingEngineMode: "observe",
|
||||||
|
ProcessRole: "worker",
|
||||||
|
LocalUploadedStorageDir: filepath.Join(storageRoot, "worker-uploaded"),
|
||||||
|
LocalGeneratedStorageDir: filepath.Join(storageRoot, "worker-generated"),
|
||||||
|
MediaRequestConcurrency: 16,
|
||||||
|
MediaMaterializationConcurrency: mediaConcurrent,
|
||||||
|
AsyncQueueWorkerEnabled: true,
|
||||||
|
AsyncWorkerHardLimit: upstreamConcurrent,
|
||||||
|
AsyncWorkerInstanceHardLimit: upstreamConcurrent,
|
||||||
|
AsyncWorkerRefreshIntervalSeconds: 1,
|
||||||
|
}
|
||||||
|
_ = NewServerWithContext(serverCtx, workerConfig, workerDB, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||||
server := httptest.NewServer(NewServerWithContext(serverCtx, config.Config{
|
server := httptest.NewServer(NewServerWithContext(serverCtx, config.Config{
|
||||||
AppEnv: "test",
|
AppEnv: "test",
|
||||||
HTTPAddr: ":0",
|
HTTPAddr: ":0",
|
||||||
@@ -369,10 +393,11 @@ LIMIT 1`).Scan(&gatewayUserID); err != nil {
|
|||||||
}
|
}
|
||||||
elapsed := time.Since(startedAt)
|
elapsed := time.Since(startedAt)
|
||||||
|
|
||||||
var tasks, succeeded, attempts, duplicateAttemptTasks int
|
var tasks, asyncTasks, succeeded, attempts, duplicateAttemptTasks int
|
||||||
if err := pool.QueryRow(ctx, `
|
if err := pool.QueryRow(ctx, `
|
||||||
SELECT
|
SELECT
|
||||||
COUNT(*)::int,
|
COUNT(*)::int,
|
||||||
|
COUNT(*) FILTER (WHERE task.async_mode)::int,
|
||||||
COUNT(*) FILTER (WHERE task.status = 'succeeded')::int,
|
COUNT(*) FILTER (WHERE task.status = 'succeeded')::int,
|
||||||
COALESCE(SUM(attempts.attempt_count), 0)::int,
|
COALESCE(SUM(attempts.attempt_count), 0)::int,
|
||||||
COUNT(*) FILTER (WHERE attempts.attempt_count <> 1)::int
|
COUNT(*) FILTER (WHERE attempts.attempt_count <> 1)::int
|
||||||
@@ -383,11 +408,11 @@ LEFT JOIN LATERAL (
|
|||||||
WHERE attempt.task_id = task.id
|
WHERE attempt.task_id = task.id
|
||||||
) attempts ON true
|
) attempts ON true
|
||||||
WHERE task.model = $1
|
WHERE task.model = $1
|
||||||
AND task.created_at >= $2`, model, startedAt).Scan(&tasks, &succeeded, &attempts, &duplicateAttemptTasks); err != nil {
|
AND task.created_at >= $2`, model, startedAt).Scan(&tasks, &asyncTasks, &succeeded, &attempts, &duplicateAttemptTasks); err != nil {
|
||||||
t.Fatalf("read Gemini stress task results: %v", err)
|
t.Fatalf("read Gemini stress task results: %v", err)
|
||||||
}
|
}
|
||||||
if tasks != taskCount || succeeded != taskCount || attempts != taskCount || duplicateAttemptTasks != 0 {
|
if tasks != taskCount || asyncTasks != taskCount || succeeded != taskCount || attempts != taskCount || duplicateAttemptTasks != 0 {
|
||||||
t.Fatalf("task results tasks=%d succeeded=%d attempts=%d duplicate_attempt_tasks=%d", tasks, succeeded, attempts, duplicateAttemptTasks)
|
t.Fatalf("task results tasks=%d async=%d succeeded=%d attempts=%d duplicate_attempt_tasks=%d", tasks, asyncTasks, succeeded, attempts, duplicateAttemptTasks)
|
||||||
}
|
}
|
||||||
if upstreamCalls.Load() != int64(taskCount) || invalidInputs.Load() != 0 {
|
if upstreamCalls.Load() != int64(taskCount) || invalidInputs.Load() != 0 {
|
||||||
t.Fatalf("upstream calls=%d invalid_inputs=%d, want %d/0", upstreamCalls.Load(), invalidInputs.Load(), taskCount)
|
t.Fatalf("upstream calls=%d invalid_inputs=%d, want %d/0", upstreamCalls.Load(), invalidInputs.Load(), taskCount)
|
||||||
@@ -454,6 +479,7 @@ type geminiStressUploadService struct {
|
|||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
apiKey string
|
apiKey string
|
||||||
retainedObjects map[string][]byte
|
retainedObjects map[string][]byte
|
||||||
|
generatedObject []byte
|
||||||
|
|
||||||
requestAssetUploads atomic.Int64
|
requestAssetUploads atomic.Int64
|
||||||
generatedUploads atomic.Int64
|
generatedUploads atomic.Int64
|
||||||
@@ -557,6 +583,9 @@ func (s *geminiStressUploadService) handleUpload(w http.ResponseWriter, r *http.
|
|||||||
http.Error(w, "unexpected generated result hash", http.StatusBadRequest)
|
http.Error(w, "unexpected generated result hash", http.StatusBadRequest)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
s.generatedObject = filePayload
|
||||||
|
s.mu.Unlock()
|
||||||
s.generatedUploads.Add(1)
|
s.generatedUploads.Add(1)
|
||||||
default:
|
default:
|
||||||
s.invalidUploads.Add(1)
|
s.invalidUploads.Add(1)
|
||||||
@@ -573,6 +602,9 @@ func (s *geminiStressUploadService) handleObject(w http.ResponseWriter, r *http.
|
|||||||
hash := strings.TrimPrefix(r.URL.Path, "/objects/")
|
hash := strings.TrimPrefix(r.URL.Path, "/objects/")
|
||||||
s.mu.RLock()
|
s.mu.RLock()
|
||||||
payload, ok := s.retainedObjects[hash]
|
payload, ok := s.retainedObjects[hash]
|
||||||
|
if !ok && hash == s.outputHash && len(s.generatedObject) > 0 {
|
||||||
|
payload, ok = s.generatedObject, true
|
||||||
|
}
|
||||||
s.mu.RUnlock()
|
s.mu.RUnlock()
|
||||||
if !ok {
|
if !ok {
|
||||||
http.NotFound(w, r)
|
http.NotFound(w, r)
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ var geminiGenerateContentRoutePrefixes = []string{
|
|||||||
func (s *Server) registerGeminiGenerateContentRoutes(mux *http.ServeMux) {
|
func (s *Server) registerGeminiGenerateContentRoutes(mux *http.ServeMux) {
|
||||||
handler := s.requireProtocolUser(clients.ProtocolGeminiGenerateContent, http.HandlerFunc(s.geminiGenerateContent))
|
handler := s.requireProtocolUser(clients.ProtocolGeminiGenerateContent, http.HandlerFunc(s.geminiGenerateContent))
|
||||||
for _, prefix := range geminiGenerateContentRoutePrefixes {
|
for _, prefix := range geminiGenerateContentRoutePrefixes {
|
||||||
mux.Handle("POST "+prefix, geminiGenerateContentRouteHandler(prefix, handler))
|
mux.Handle("POST "+prefix, geminiGenerateContentRouteHandler(prefix, s.withMediaRequestBodySlot(handler)))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -122,7 +122,7 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
|
|||||||
writeTaskTrafficError(w, err, clients.ProtocolGeminiGenerateContent)
|
writeTaskTrafficError(w, err, clients.ProtocolGeminiGenerateContent)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
releaseRequestBody, err := s.acquireMediaRequestBodySlot(r.Context())
|
releaseRequestBody, err := s.requestMediaBodyRelease(r.Context())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -181,7 +181,7 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
|
|||||||
Model: mapping.Model,
|
Model: mapping.Model,
|
||||||
RunMode: runMode,
|
RunMode: runMode,
|
||||||
AcceptanceRunID: admission.AcceptanceRunID,
|
AcceptanceRunID: admission.AcceptanceRunID,
|
||||||
Async: false,
|
Async: !streamMode,
|
||||||
Request: prepared.Body,
|
Request: prepared.Body,
|
||||||
ConversationID: prepared.ConversationID,
|
ConversationID: prepared.ConversationID,
|
||||||
NewMessageCount: prepared.NewMessageCount,
|
NewMessageCount: prepared.NewMessageCount,
|
||||||
@@ -240,27 +240,36 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
|
|||||||
s.writeGeminiGenerateContentStream(runCtx, w, r, task, user, mapping.Model, writeGeminiTaskError)
|
s.writeGeminiGenerateContentStream(runCtx, w, r, task, user, mapping.Model, writeGeminiTaskError)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
result, runErr := s.runner.Execute(runCtx, task, user)
|
if err := s.runner.SubmitAsyncTask(runCtx, task); err != nil {
|
||||||
if runErr != nil {
|
|
||||||
if !requestStillConnected(r) {
|
if !requestStillConnected(r) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
applyRunErrorHeaders(w, runErr)
|
applyRunErrorHeaders(w, err)
|
||||||
if wire := clients.ErrorWireResponse(runErr); wireResponseMatches(wire, clients.ProtocolGeminiGenerateContent) {
|
writeGeminiTaskError(statusFromRunError(err), runErrorMessage(err), runErrorDetails(err), runErrorCode(err))
|
||||||
writeWireResponse(w, wire)
|
return
|
||||||
return
|
}
|
||||||
}
|
task, err = s.runner.WaitForTaskCompletion(runCtx, task.ID)
|
||||||
writeGeminiTaskError(statusFromRunError(runErr), runErrorMessage(runErr), runErrorDetails(runErr), runErrorCode(runErr))
|
if err != nil || !requestStillConnected(r) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if task.Status != "succeeded" {
|
||||||
|
status := storedTaskErrorStatus(task.ErrorCode)
|
||||||
|
writeGeminiTaskError(status, firstNonEmpty(task.ErrorMessage, task.Error, task.Message, "task failed"), nil, firstNonEmpty(task.ErrorCode, task.Status))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
task, err = s.hydrateTaskResult(runCtx, task)
|
||||||
|
if err != nil {
|
||||||
|
writeGeminiTaskError(statusFromRunError(err), err.Error(), nil, clients.ErrorCode(err))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !requestStillConnected(r) {
|
if !requestStillConnected(r) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if wireResponseMatches(result.Wire, clients.ProtocolGeminiGenerateContent) {
|
if _, native := task.Result["candidates"]; native {
|
||||||
writeWireResponse(w, result.Wire)
|
writeJSON(w, http.StatusOK, task.Result)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
writeJSON(w, http.StatusOK, geminiGenerateContentResponse(result.Output, mapping.Model))
|
writeJSON(w, http.StatusOK, geminiGenerateContentResponse(task.Result, mapping.Model))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) writeGeminiGenerateContentStream(runCtx context.Context, w http.ResponseWriter, r *http.Request, task store.GatewayTask, user *auth.User, model string, writeGeminiTaskError func(int, string, map[string]any, string)) {
|
func (s *Server) writeGeminiGenerateContentStream(runCtx context.Context, w http.ResponseWriter, r *http.Request, task store.GatewayTask, user *auth.User, model string, writeGeminiTaskError func(int, string, map[string]any, string)) {
|
||||||
|
|||||||
@@ -217,6 +217,38 @@ func (s *Server) acquireMediaRequestBodySlot(ctx context.Context) (func(), error
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type mediaRequestBodyLeaseContextKey struct{}
|
||||||
|
|
||||||
|
type mediaRequestBodyLease struct {
|
||||||
|
release func()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *mediaRequestBodyLease) Release() {
|
||||||
|
if l != nil && l.release != nil {
|
||||||
|
l.release()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) withMediaRequestBodySlot(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
release, err := s.acquireMediaRequestBodySlot(r.Context())
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
lease := &mediaRequestBodyLease{release: release}
|
||||||
|
defer lease.Release()
|
||||||
|
ctx := context.WithValue(r.Context(), mediaRequestBodyLeaseContextKey{}, lease)
|
||||||
|
next.ServeHTTP(w, r.WithContext(ctx))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) requestMediaBodyRelease(ctx context.Context) (func(), error) {
|
||||||
|
if lease, ok := ctx.Value(mediaRequestBodyLeaseContextKey{}).(*mediaRequestBodyLease); ok && lease != nil {
|
||||||
|
return lease.Release, nil
|
||||||
|
}
|
||||||
|
return s.acquireMediaRequestBodySlot(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
func requestAssetFromValue(key string, path []string, value any, siblings map[string]any) (decodedRequestAsset, bool, error) {
|
func requestAssetFromValue(key string, path []string, value any, siblings map[string]any) (decodedRequestAsset, bool, error) {
|
||||||
text, ok := value.(string)
|
text, ok := value.(string)
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -39,6 +41,59 @@ func TestRequestAssetFromValueDetectsDataURLAndRawBase64(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMediaRequestBodySlotLimitsPreAuthWorkAndCanReleaseEarly(t *testing.T) {
|
||||||
|
server := &Server{mediaRequestBodySlots: make(chan struct{}, 2)}
|
||||||
|
var critical atomic.Int64
|
||||||
|
var maxCritical atomic.Int64
|
||||||
|
var afterRelease atomic.Int64
|
||||||
|
var maxAfterRelease atomic.Int64
|
||||||
|
updateMax := func(target *atomic.Int64, value int64) {
|
||||||
|
for {
|
||||||
|
current := target.Load()
|
||||||
|
if value <= current || target.CompareAndSwap(current, value) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
handler := server.withMediaRequestBodySlot(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
release, err := server.requestMediaBodyRelease(r.Context())
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("request media body release: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
current := critical.Add(1)
|
||||||
|
updateMax(&maxCritical, current)
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
critical.Add(-1)
|
||||||
|
release()
|
||||||
|
current = afterRelease.Add(1)
|
||||||
|
updateMax(&maxAfterRelease, current)
|
||||||
|
time.Sleep(40 * time.Millisecond)
|
||||||
|
afterRelease.Add(-1)
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
}))
|
||||||
|
|
||||||
|
var wait sync.WaitGroup
|
||||||
|
for index := 0; index < 10; index++ {
|
||||||
|
wait.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wait.Done()
|
||||||
|
request := httptest.NewRequest(http.MethodPost, "/v1beta/models/test:generateContent", nil)
|
||||||
|
handler.ServeHTTP(httptest.NewRecorder(), request)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wait.Wait()
|
||||||
|
if maxCritical.Load() != 2 {
|
||||||
|
t.Fatalf("pre-auth critical concurrency=%d, want 2", maxCritical.Load())
|
||||||
|
}
|
||||||
|
if maxAfterRelease.Load() <= 2 {
|
||||||
|
t.Fatalf("early release did not admit later requests, post-release concurrency=%d", maxAfterRelease.Load())
|
||||||
|
}
|
||||||
|
if len(server.mediaRequestBodySlots) != 0 {
|
||||||
|
t.Fatalf("media request body slots leaked: %d", len(server.mediaRequestBodySlots))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestRequestModelNameSupportsObjectModelReference(t *testing.T) {
|
func TestRequestModelNameSupportsObjectModelReference(t *testing.T) {
|
||||||
got := requestModelName(map[string]any{
|
got := requestModelName(map[string]any{
|
||||||
"model": map[string]any{
|
"model": map[string]any{
|
||||||
|
|||||||
@@ -99,6 +99,9 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
|
|||||||
}
|
}
|
||||||
server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey
|
server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey
|
||||||
server.runner.StartAdmissionNotifier(ctx)
|
server.runner.StartAdmissionNotifier(ctx)
|
||||||
|
if cfg.RunsPublicHTTP() {
|
||||||
|
server.runner.StartTaskCompletionWaiter(ctx)
|
||||||
|
}
|
||||||
if cfg.RunsAsyncExecutionWorker() {
|
if cfg.RunsAsyncExecutionWorker() {
|
||||||
server.runner.StartAsyncQueueWorker(ctx)
|
server.runner.StartAsyncQueueWorker(ctx)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
"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/config"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -405,6 +406,17 @@ func (s *Service) hydrateLocalBinaryValue(ctx context.Context, taskID string, va
|
|||||||
}
|
}
|
||||||
switch typed := value.(type) {
|
switch typed := value.(type) {
|
||||||
case map[string]any:
|
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))
|
next := make(map[string]any, len(typed))
|
||||||
changed := false
|
changed := false
|
||||||
for key, childValue := range typed {
|
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) {
|
func (s *Service) readLocalBinaryResult(taskID string, descriptor localBinaryDescriptor) ([]byte, error) {
|
||||||
path := filepath.Join(s.localBinaryResultRoot(), safeLocalBinaryTaskDir(taskID), descriptor.SHA256+".bin")
|
path := filepath.Join(s.localBinaryResultRoot(), safeLocalBinaryTaskDir(taskID), descriptor.SHA256+".bin")
|
||||||
info, err := os.Stat(path)
|
info, err := os.Stat(path)
|
||||||
|
|||||||
@@ -2,9 +2,13 @@ package runner
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
|
"encoding/hex"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"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) {
|
func TestHydrateLocalBinaryResultReturnsExpiredAndCorruptedErrors(t *testing.T) {
|
||||||
service := newLocalBinaryTestService(t)
|
service := newLocalBinaryTestService(t)
|
||||||
service.cfg.LocalResultTTLHours = 1
|
service.cfg.LocalResultTTLHours = 1
|
||||||
|
|||||||
@@ -25,28 +25,32 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type Service struct {
|
type Service struct {
|
||||||
cfg config.Config
|
cfg config.Config
|
||||||
store *store.Store
|
store *store.Store
|
||||||
logger *slog.Logger
|
logger *slog.Logger
|
||||||
clients map[string]clients.Client
|
clients map[string]clients.Client
|
||||||
scriptExecutor *scriptengine.Executor
|
scriptExecutor *scriptengine.Executor
|
||||||
httpClients *httpClientCache
|
httpClients *httpClientCache
|
||||||
riverMu sync.RWMutex
|
riverMu sync.RWMutex
|
||||||
riverControlClient *river.Client[pgx.Tx]
|
riverControlClient *river.Client[pgx.Tx]
|
||||||
riverExecutionClient asyncExecutionClient
|
riverExecutionClient asyncExecutionClient
|
||||||
riverDrainingClients map[asyncExecutionClient]struct{}
|
riverDrainingClients map[asyncExecutionClient]struct{}
|
||||||
riverWorkerCapacity int
|
riverWorkerCapacity int
|
||||||
workerInstanceID string
|
workerInstanceID string
|
||||||
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
|
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
|
||||||
admissionWakeMu sync.Mutex
|
admissionWakeMu sync.Mutex
|
||||||
admissionWake chan struct{}
|
admissionWake chan struct{}
|
||||||
asyncAdmissionWake chan struct{}
|
asyncAdmissionWake chan struct{}
|
||||||
admissionTaskWaiters map[string]*admissionTaskWaiter
|
admissionTaskWaiters map[string]*admissionTaskWaiter
|
||||||
admissionListener sync.Once
|
admissionListener sync.Once
|
||||||
asyncClientFactory func(int) (asyncExecutionClient, error)
|
asyncClientFactory func(int) (asyncExecutionClient, error)
|
||||||
mediaResultSlots chan struct{}
|
mediaResultSlots chan struct{}
|
||||||
directOSS *directOSSUploader
|
directOSS *directOSSUploader
|
||||||
billingMetrics billingMetricsObserver
|
taskCompletionOnce sync.Once
|
||||||
|
taskCompletionMu sync.Mutex
|
||||||
|
taskCompletionWaiters map[string]map[chan struct{}]struct{}
|
||||||
|
taskCompletionPollWake chan struct{}
|
||||||
|
billingMetrics billingMetricsObserver
|
||||||
}
|
}
|
||||||
|
|
||||||
type billingMetricsObserver interface {
|
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},
|
"universal": clients.UniversalClient{HTTPClient: httpClients.none, ScriptExecutor: scriptExecutor},
|
||||||
"simulation": clients.SimulationClient{},
|
"simulation": clients.SimulationClient{},
|
||||||
},
|
},
|
||||||
httpClients: httpClients,
|
httpClients: httpClients,
|
||||||
workerInstanceID: asyncWorkerID(),
|
workerInstanceID: asyncWorkerID(),
|
||||||
admissionWake: make(chan struct{}, 4096),
|
admissionWake: make(chan struct{}, 4096),
|
||||||
asyncAdmissionWake: make(chan struct{}, 1),
|
asyncAdmissionWake: make(chan struct{}, 1),
|
||||||
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
|
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
|
||||||
mediaResultSlots: make(chan struct{}, cfg.MediaMaterializationConcurrency),
|
mediaResultSlots: make(chan struct{}, cfg.MediaMaterializationConcurrency),
|
||||||
directOSS: newDirectOSSUploader(cfg),
|
directOSS: newDirectOSSUploader(cfg),
|
||||||
|
taskCompletionWaiters: map[string]map[chan struct{}]struct{}{},
|
||||||
|
taskCompletionPollWake: make(chan struct{}, 1),
|
||||||
}
|
}
|
||||||
if len(observers) > 0 {
|
if len(observers) > 0 {
|
||||||
service.billingMetrics = observers[0]
|
service.billingMetrics = observers[0]
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -236,6 +236,9 @@ func (s *Service) uploadGeneratedAssets(ctx context.Context, taskID string, task
|
|||||||
if contentType != "" && stringFromAny(merged["mime_type"]) == "" {
|
if contentType != "" && stringFromAny(merged["mime_type"]) == "" {
|
||||||
merged["mime_type"] = contentType
|
merged["mime_type"] = contentType
|
||||||
}
|
}
|
||||||
|
if decision.Inline != nil && strings.TrimSpace(sourceKey) != "" {
|
||||||
|
merged[sourceKey] = generatedRawMediaReference(decision.Inline, upload, contentType, kind, strategy)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
nextData = append(nextData, merged)
|
nextData = append(nextData, merged)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2195,6 +2195,30 @@ SELECT `+gatewayTaskColumns+`
|
|||||||
return task, nil
|
return task, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Store) ListTaskStatuses(ctx context.Context, taskIDs []string) (map[string]string, error) {
|
||||||
|
statuses := make(map[string]string, len(taskIDs))
|
||||||
|
if len(taskIDs) == 0 {
|
||||||
|
return statuses, nil
|
||||||
|
}
|
||||||
|
rows, err := s.pool.Query(ctx, `
|
||||||
|
SELECT id::text, status
|
||||||
|
FROM gateway_tasks
|
||||||
|
WHERE id = ANY($1::uuid[])`, taskIDs)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
for rows.Next() {
|
||||||
|
var taskID string
|
||||||
|
var status string
|
||||||
|
if err := rows.Scan(&taskID, &status); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
statuses[taskID] = status
|
||||||
|
}
|
||||||
|
return statuses, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
type taskScanner interface {
|
type taskScanner interface {
|
||||||
Scan(dest ...any) error
|
Scan(dest ...any) error
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user