diff --git a/apps/api/internal/httpapi/gemini_base64_stress_integration_test.go b/apps/api/internal/httpapi/gemini_base64_stress_integration_test.go index 13919d7..188590a 100644 --- a/apps/api/internal/httpapi/gemini_base64_stress_integration_test.go +++ b/apps/api/internal/httpapi/gemini_base64_stress_integration_test.go @@ -43,7 +43,7 @@ func TestGeminiBase64ImageEditThousandSynchronousRequests(t *testing.T) { mediaConcurrent = 8 ) 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) 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)) @@ -59,6 +59,11 @@ func TestGeminiBase64ImageEditThousandSynchronousRequests(t *testing.T) { t.Fatalf("connect API store: %v", err) } 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) if err != nil { t.Fatalf("parse observation pool config: %v", err) @@ -152,6 +157,25 @@ WHERE status = 'active'`); err != nil { serverCtx, cancelServer := context.WithCancel(ctx) defer cancelServer() 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{ AppEnv: "test", HTTPAddr: ":0", @@ -369,10 +393,11 @@ LIMIT 1`).Scan(&gatewayUserID); err != nil { } elapsed := time.Since(startedAt) - var tasks, succeeded, attempts, duplicateAttemptTasks int + var tasks, asyncTasks, succeeded, attempts, duplicateAttemptTasks int if err := pool.QueryRow(ctx, ` SELECT COUNT(*)::int, + COUNT(*) FILTER (WHERE task.async_mode)::int, COUNT(*) FILTER (WHERE task.status = 'succeeded')::int, COALESCE(SUM(attempts.attempt_count), 0)::int, COUNT(*) FILTER (WHERE attempts.attempt_count <> 1)::int @@ -383,11 +408,11 @@ LEFT JOIN LATERAL ( WHERE attempt.task_id = task.id ) attempts ON true 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) } - if tasks != 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) + if tasks != taskCount || asyncTasks != taskCount || succeeded != taskCount || attempts != taskCount || duplicateAttemptTasks != 0 { + 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 { 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 apiKey string retainedObjects map[string][]byte + generatedObject []byte requestAssetUploads 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) return } + s.mu.Lock() + s.generatedObject = filePayload + s.mu.Unlock() s.generatedUploads.Add(1) default: s.invalidUploads.Add(1) @@ -573,6 +602,9 @@ func (s *geminiStressUploadService) handleObject(w http.ResponseWriter, r *http. hash := strings.TrimPrefix(r.URL.Path, "/objects/") s.mu.RLock() payload, ok := s.retainedObjects[hash] + if !ok && hash == s.outputHash && len(s.generatedObject) > 0 { + payload, ok = s.generatedObject, true + } s.mu.RUnlock() if !ok { http.NotFound(w, r) diff --git a/apps/api/internal/httpapi/gemini_compat.go b/apps/api/internal/httpapi/gemini_compat.go index f21d79a..d0f8ef2 100644 --- a/apps/api/internal/httpapi/gemini_compat.go +++ b/apps/api/internal/httpapi/gemini_compat.go @@ -49,7 +49,7 @@ var geminiGenerateContentRoutePrefixes = []string{ func (s *Server) registerGeminiGenerateContentRoutes(mux *http.ServeMux) { handler := s.requireProtocolUser(clients.ProtocolGeminiGenerateContent, http.HandlerFunc(s.geminiGenerateContent)) 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) return } - releaseRequestBody, err := s.acquireMediaRequestBodySlot(r.Context()) + releaseRequestBody, err := s.requestMediaBodyRelease(r.Context()) if err != nil { return } @@ -181,7 +181,7 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) { Model: mapping.Model, RunMode: runMode, AcceptanceRunID: admission.AcceptanceRunID, - Async: false, + Async: !streamMode, Request: prepared.Body, ConversationID: prepared.ConversationID, 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) return } - result, runErr := s.runner.Execute(runCtx, task, user) - if runErr != nil { + if err := s.runner.SubmitAsyncTask(runCtx, task); err != nil { if !requestStillConnected(r) { return } - applyRunErrorHeaders(w, runErr) - if wire := clients.ErrorWireResponse(runErr); wireResponseMatches(wire, clients.ProtocolGeminiGenerateContent) { - writeWireResponse(w, wire) - return - } - writeGeminiTaskError(statusFromRunError(runErr), runErrorMessage(runErr), runErrorDetails(runErr), runErrorCode(runErr)) + applyRunErrorHeaders(w, err) + writeGeminiTaskError(statusFromRunError(err), runErrorMessage(err), runErrorDetails(err), runErrorCode(err)) + return + } + task, err = s.runner.WaitForTaskCompletion(runCtx, task.ID) + 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 } if !requestStillConnected(r) { return } - if wireResponseMatches(result.Wire, clients.ProtocolGeminiGenerateContent) { - writeWireResponse(w, result.Wire) + if _, native := task.Result["candidates"]; native { + writeJSON(w, http.StatusOK, task.Result) 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)) { diff --git a/apps/api/internal/httpapi/request_preparation.go b/apps/api/internal/httpapi/request_preparation.go index c0c2d7b..d2e3852 100644 --- a/apps/api/internal/httpapi/request_preparation.go +++ b/apps/api/internal/httpapi/request_preparation.go @@ -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) { text, ok := value.(string) if !ok { diff --git a/apps/api/internal/httpapi/request_preparation_test.go b/apps/api/internal/httpapi/request_preparation_test.go index 638d1a9..9f56148 100644 --- a/apps/api/internal/httpapi/request_preparation_test.go +++ b/apps/api/internal/httpapi/request_preparation_test.go @@ -12,6 +12,8 @@ import ( "os" "path/filepath" "strings" + "sync" + "sync/atomic" "testing" "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) { got := requestModelName(map[string]any{ "model": map[string]any{ diff --git a/apps/api/internal/httpapi/server.go b/apps/api/internal/httpapi/server.go index e69a392..43ca3f0 100644 --- a/apps/api/internal/httpapi/server.go +++ b/apps/api/internal/httpapi/server.go @@ -99,6 +99,9 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor } server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey server.runner.StartAdmissionNotifier(ctx) + if cfg.RunsPublicHTTP() { + server.runner.StartTaskCompletionWaiter(ctx) + } if cfg.RunsAsyncExecutionWorker() { server.runner.StartAsyncQueueWorker(ctx) } else { diff --git a/apps/api/internal/runner/binary_results.go b/apps/api/internal/runner/binary_results.go index c83b48a..57f9d73 100644 --- a/apps/api/internal/runner/binary_results.go +++ b/apps/api/internal/runner/binary_results.go @@ -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) diff --git a/apps/api/internal/runner/binary_results_test.go b/apps/api/internal/runner/binary_results_test.go index 0b8d699..45d69bf 100644 --- a/apps/api/internal/runner/binary_results_test.go +++ b/apps/api/internal/runner/binary_results_test.go @@ -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 diff --git a/apps/api/internal/runner/service.go b/apps/api/internal/runner/service.go index 331b400..a5eb450 100644 --- a/apps/api/internal/runner/service.go +++ b/apps/api/internal/runner/service.go @@ -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] diff --git a/apps/api/internal/runner/task_completion.go b/apps/api/internal/runner/task_completion.go new file mode 100644 index 0000000..f41ef4d --- /dev/null +++ b/apps/api/internal/runner/task_completion.go @@ -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 + } +} diff --git a/apps/api/internal/runner/task_completion_test.go b/apps/api/internal/runner/task_completion_test.go new file mode 100644 index 0000000..889a19c --- /dev/null +++ b/apps/api/internal/runner/task_completion_test.go @@ -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) + } + } +} diff --git a/apps/api/internal/runner/upload.go b/apps/api/internal/runner/upload.go index a1b23dc..a5e97bf 100644 --- a/apps/api/internal/runner/upload.go +++ b/apps/api/internal/runner/upload.go @@ -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) } diff --git a/apps/api/internal/store/postgres.go b/apps/api/internal/store/postgres.go index e219f78..051028b 100644 --- a/apps/api/internal/store/postgres.go +++ b/apps/api/internal/store/postgres.go @@ -2195,6 +2195,30 @@ SELECT `+gatewayTaskColumns+` 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 { Scan(dest ...any) error }