From d7951cfdd2745c2c55ded8509230fd7d49b0165b Mon Sep 17 00:00:00 2001 From: wangbo Date: Tue, 21 Jul 2026 23:54:04 +0800 Subject: [PATCH] =?UTF-8?q?feat(volces):=20=E6=8E=A5=E5=85=A5=E4=BB=BB?= =?UTF-8?q?=E5=8A=A1=E5=85=BC=E5=AE=B9=E6=9F=A5=E8=AF=A2=E4=B8=8E=E4=B8=8A?= =?UTF-8?q?=E6=B8=B8=E5=8F=96=E6=B6=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/api/internal/clients/clients_test.go | 71 ++++ apps/api/internal/clients/types.go | 1 + apps/api/internal/clients/volces.go | 186 +++++++--- apps/api/internal/httpapi/server.go | 10 + .../httpapi/volces_compat_handlers.go | 330 ++++++++++++++++++ .../httpapi/volces_compat_handlers_test.go | 29 ++ apps/api/internal/runner/service.go | 34 ++ apps/api/internal/runner/task_cancel.go | 56 +++ .../internal/store/remote_task_candidate.go | 43 +++ apps/api/internal/store/tasks_runtime.go | 69 +++- .../internal/store/volces_compatible_tasks.go | 111 ++++++ 11 files changed, 885 insertions(+), 55 deletions(-) create mode 100644 apps/api/internal/httpapi/volces_compat_handlers.go create mode 100644 apps/api/internal/httpapi/volces_compat_handlers_test.go create mode 100644 apps/api/internal/store/remote_task_candidate.go create mode 100644 apps/api/internal/store/volces_compatible_tasks.go diff --git a/apps/api/internal/clients/clients_test.go b/apps/api/internal/clients/clients_test.go index f941145..700f6d9 100644 --- a/apps/api/internal/clients/clients_test.go +++ b/apps/api/internal/clients/clients_test.go @@ -1934,6 +1934,77 @@ func TestVolcesClientVideoSubmitsAndPollsTask(t *testing.T) { } } +func TestVolcesClientVideoRetriesTransientPollAndKeepsOfficialResult(t *testing.T) { + polls := 0 + persisted := make([]string, 0) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case "POST /contents/generations/tasks": + _ = json.NewEncoder(w).Encode(map[string]any{"id": "cgt-retry"}) + case "GET /contents/generations/tasks/cgt-retry": + polls++ + if polls == 1 { + w.WriteHeader(http.StatusServiceUnavailable) + _, _ = w.Write([]byte(`{"error":{"message":"try later"}}`)) + return + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "id": "cgt-retry", "model": "doubao-seedance-2-0-mini-260615", "status": "succeeded", + "created_at": 123, "updated_at": 124, "content": map[string]any{"video_url": "https://example.com/retry.mp4"}, + "usage": map[string]any{"total_tokens": 8}, "seed": 7, + }) + default: + t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path) + } + })) + defer server.Close() + + response, err := (VolcesClient{HTTPClient: server.Client()}).Run(context.Background(), Request{ + Kind: "videos.generations", Model: "seedance", Body: map[string]any{"model": "seedance", "prompt": "retry"}, + Candidate: store.RuntimeModelCandidate{BaseURL: server.URL, ProviderModelName: "doubao-seedance-2-0-mini-260615", Credentials: map[string]any{"apiKey": "key"}, PlatformConfig: map[string]any{"volcesPollIntervalMs": 100, "volcesPollRetryMaxMs": 100, "volcesPollTimeoutSeconds": 2}}, + OnRemoteTaskPolled: func(remoteTaskID string, payload map[string]any) error { + persisted = append(persisted, remoteTaskID+":"+stringFromAny(payload["status"])) + return nil + }, + }) + if err != nil { + t.Fatalf("run retrying Volces video: %v", err) + } + if polls != 2 || len(persisted) != 1 || persisted[0] != "cgt-retry:succeeded" { + t.Fatalf("unexpected poll state polls=%d persisted=%+v", polls, persisted) + } + if response.Result["updated_at"] != float64(124) || response.Result["seed"] != float64(7) || response.Result["raw"] == nil { + t.Fatalf("official result fields lost: %+v", response.Result) + } +} + +func TestVolcesClientDeletesOfficialVideoTask(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodDelete || r.URL.Path != "/contents/generations/tasks/cgt-delete" { + t.Fatalf("unexpected delete request %s %s", r.Method, r.URL.Path) + } + if r.Header.Get("Authorization") != "Bearer delete-key" { + t.Fatalf("unexpected delete authorization: %q", r.Header.Get("Authorization")) + } + _ = json.NewEncoder(w).Encode(map[string]any{"id": "cgt-delete", "status": "cancelled"}) + })) + defer server.Close() + + result, _, err := (VolcesClient{HTTPClient: server.Client()}).DeleteVideoTask(context.Background(), Request{ + Candidate: store.RuntimeModelCandidate{BaseURL: server.URL, Credentials: map[string]any{"apiKey": "delete-key"}}, + RemoteTaskID: "cgt-delete", + }) + if err != nil || result["status"] != "cancelled" { + t.Fatalf("unexpected delete response result=%+v err=%v", result, err) + } +} + +func TestVolcesCancelledTaskUsesDedicatedCancellationCode(t *testing.T) { + if got := volcesTaskErrorCode(map[string]any{"status": "cancelled"}); got != "volces_task_cancelled" { + t.Fatalf("cancelled task error code = %q", got) + } +} + func TestVolcesClientVideoRejectsDuplicateFirstFrameBeforeSubmit(t *testing.T) { var submitted bool server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { diff --git a/apps/api/internal/clients/types.go b/apps/api/internal/clients/types.go index 6f1db88..9b2d924 100644 --- a/apps/api/internal/clients/types.go +++ b/apps/api/internal/clients/types.go @@ -20,6 +20,7 @@ type Request struct { RemoteTaskID string RemoteTaskPayload map[string]any OnRemoteTaskSubmitted func(remoteTaskID string, payload map[string]any) error + OnRemoteTaskPolled func(remoteTaskID string, payload map[string]any) error Stream bool StreamDelta StreamDelta UpstreamProtocol string diff --git a/apps/api/internal/clients/volces.go b/apps/api/internal/clients/volces.go index 9eefa56..297105a 100644 --- a/apps/api/internal/clients/volces.go +++ b/apps/api/internal/clients/volces.go @@ -100,66 +100,105 @@ func (c VolcesClient) runVideo(ctx context.Context, request Request, apiKey stri timeout := volcesPollTimeout(request) deadline := time.NewTimer(timeout) defer deadline.Stop() - - ticker := time.NewTicker(interval) - defer ticker.Stop() + nextPoll := time.NewTimer(0) + defer nextPoll.Stop() var lastResult map[string]any + lastRequestID := firstNonEmpty(submitRequestID, upstreamTaskID) + transientFailures := 0 for { select { case <-ctx.Done(): - return Response{}, &ClientError{Code: "cancelled", Message: ctx.Err().Error(), RequestID: submitRequestID, Retryable: true} - default: - } - - pollStartedAt := time.Now() - pollResult, pollRequestID, err := c.getJSON(ctx, request, request.Candidate.BaseURL, taskPath+"/"+upstreamTaskID, apiKey) - pollFinishedAt := time.Now() - requestID := firstNonEmpty(pollRequestID, submitRequestID, upstreamTaskID) - if err != nil { - return Response{}, annotateResponseError(err, requestID, pollStartedAt, pollFinishedAt) - } - lastResult = pollResult - - switch volcesTaskStatus(pollResult) { - case "succeeded": - result := volcesVideoSuccessResult(request, upstreamTaskID, pollResult) - return Response{ - Result: result, - RequestID: requestID, - Usage: volcesVideoUsage(pollResult), - Progress: volcesVideoProgress(request, upstreamTaskID), - ResponseStartedAt: submitStartedAt, - ResponseFinishedAt: pollFinishedAt, - ResponseDurationMS: responseDurationMS(submitStartedAt, pollFinishedAt), - }, nil - case "failed", "cancelled": - return Response{}, &ClientError{ - Code: volcesTaskErrorCode(pollResult), - Message: volcesTaskErrorMessage(pollResult), - RequestID: requestID, - ResponseStartedAt: submitStartedAt, - ResponseFinishedAt: pollFinishedAt, - ResponseDurationMS: responseDurationMS(submitStartedAt, pollFinishedAt), - Retryable: false, - } - } - - select { - case <-ctx.Done(): - return Response{}, &ClientError{Code: "cancelled", Message: ctx.Err().Error(), RequestID: requestID, Retryable: true} + return Response{}, &ClientError{Code: "cancelled", Message: ctx.Err().Error(), RequestID: lastRequestID, Retryable: true} case <-deadline.C: return Response{}, &ClientError{ Code: "timeout", Message: fmt.Sprintf("volces video task %s did not finish before timeout; last status: %s", upstreamTaskID, volcesTaskStatus(lastResult)), - RequestID: requestID, + RequestID: lastRequestID, Retryable: true, } - case <-ticker.C: + case <-nextPoll.C: + pollStartedAt := time.Now() + pollResult, pollRequestID, err := c.getJSON(ctx, request, request.Candidate.BaseURL, taskPath+"/"+upstreamTaskID, apiKey) + pollFinishedAt := time.Now() + requestID := firstNonEmpty(pollRequestID, submitRequestID, upstreamTaskID) + lastRequestID = requestID + if err != nil { + err = annotateResponseError(err, requestID, pollStartedAt, pollFinishedAt) + if !IsRetryable(err) { + return Response{}, err + } + transientFailures++ + resetVolcesPollTimer(nextPoll, volcesRetryPollInterval(request, interval, transientFailures)) + continue + } + transientFailures = 0 + lastResult = pollResult + if request.OnRemoteTaskPolled != nil { + if err := request.OnRemoteTaskPolled(upstreamTaskID, pollResult); err != nil { + return Response{}, err + } + } + + switch volcesTaskStatus(pollResult) { + case "succeeded": + result := volcesVideoSuccessResult(request, upstreamTaskID, pollResult) + return Response{ + Result: result, + RequestID: requestID, + Usage: volcesVideoUsage(pollResult), + Progress: volcesVideoProgress(request, upstreamTaskID), + ResponseStartedAt: submitStartedAt, + ResponseFinishedAt: pollFinishedAt, + ResponseDurationMS: responseDurationMS(submitStartedAt, pollFinishedAt), + }, nil + case "failed", "cancelled": + return Response{}, &ClientError{ + Code: volcesTaskErrorCode(pollResult), + Message: volcesTaskErrorMessage(pollResult), + RequestID: requestID, + ResponseStartedAt: submitStartedAt, + ResponseFinishedAt: pollFinishedAt, + ResponseDurationMS: responseDurationMS(submitStartedAt, pollFinishedAt), + Retryable: false, + } + } + resetVolcesPollTimer(nextPoll, interval) } } } +// DeleteVideoTask calls the official contents-generations cancellation endpoint. +// It is intentionally separate from Run so task cancellation can use the same +// provider credentials that submitted the remote task. +func (c VolcesClient) DeleteVideoTask(ctx context.Context, request Request) (map[string]any, string, error) { + apiKey := credential(request.Candidate.Credentials, "apiKey", "api_key", "key", "token") + if apiKey == "" { + return nil, "", &ClientError{Code: "missing_credentials", Message: "volces api key is required", Retryable: false} + } + remoteTaskID := strings.TrimSpace(request.RemoteTaskID) + if remoteTaskID == "" { + return nil, "", &ClientError{Code: "invalid_parameter", Message: "volces remote task id is required", Retryable: false} + } + taskPath := volcesVideoTaskPath(request) + "/" + remoteTaskID + req, err := http.NewRequestWithContext(ctx, http.MethodDelete, joinURL(request.Candidate.BaseURL, taskPath), nil) + if err != nil { + return nil, "", err + } + req.Header.Set("Authorization", "Bearer "+apiKey) + response, err := httpClient(request.HTTPClient, c.HTTPClient).Do(req) + if err != nil { + return nil, "", &ClientError{Code: "network", Message: err.Error(), Retryable: true} + } + requestID := requestIDFromHTTPResponse(response) + result, err := decodeHTTPResponse(response) + if err != nil { + return result, requestID, annotateResponseError(err, requestID, time.Now(), time.Now()) + } + result, envelopeRequestID, err := normalizeVolcesCompatibleResult(result) + return result, firstNonEmpty(requestID, envelopeRequestID), err +} + func volcesVideoTaskPath(request Request) string { path := firstNonEmptyStringValue( request.Candidate.PlatformConfig, @@ -997,6 +1036,9 @@ func volcesTaskErrorCode(result map[string]any) string { return code } status := volcesTaskStatus(result) + if status == "cancelled" { + return "volces_task_cancelled" + } if status != "" { return status } @@ -1015,6 +1057,10 @@ func volcesTaskErrorMessage(result map[string]any) string { } func volcesVideoSuccessResult(request Request, upstreamTaskID string, raw map[string]any) map[string]any { + result := cloneMapAny(raw) + if result == nil { + result = map[string]any{} + } content, _ := raw["content"].(map[string]any) videoURL := strings.TrimSpace(stringFromAny(content["video_url"])) created := intFromAny(raw["created_at"]) @@ -1025,16 +1071,17 @@ func volcesVideoSuccessResult(request Request, upstreamTaskID string, raw map[st if videoURL != "" { data = append(data, map[string]any{"url": videoURL, "type": "video"}) } - return map[string]any{ - "id": upstreamTaskID, - "object": "video.generation", - "created": created, - "model": upstreamModelName(request.Candidate), - "status": "succeeded", - "upstream_task_id": upstreamTaskID, - "data": data, - "raw": raw, + result["id"] = firstNonEmpty(stringFromAny(raw["id"]), upstreamTaskID) + if strings.TrimSpace(stringFromAny(result["model"])) == "" { + result["model"] = upstreamModelName(request.Candidate) } + result["status"] = "succeeded" + result["object"] = "video.generation" + result["created"] = created + result["upstream_task_id"] = upstreamTaskID + result["data"] = data + result["raw"] = cloneMapAny(raw) + return result } func volcesVideoUsage(raw map[string]any) Usage { @@ -1074,6 +1121,37 @@ func volcesPollTimeout(request Request) time.Duration { return time.Duration(seconds) * time.Second } +func volcesRetryPollInterval(request Request, normal time.Duration, failures int) time.Duration { + if failures < 1 { + return normal + } + max := time.Duration(numericValue(firstPresent(request.Candidate.PlatformConfig["volcesPollRetryMaxMs"], request.Body["pollRetryMaxMs"], request.Body["poll_retry_max_ms"]), 30000)) * time.Millisecond + if max < normal { + max = normal + } + delay := normal + for attempt := 1; attempt < failures && delay < max; attempt++ { + delay *= 2 + } + if delay > max { + return max + } + return delay +} + +func resetVolcesPollTimer(timer *time.Timer, delay time.Duration) { + if delay < 100*time.Millisecond { + delay = 100 * time.Millisecond + } + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(delay) +} + func firstNonEmpty(values ...string) string { for _, value := range values { if strings.TrimSpace(value) != "" { diff --git a/apps/api/internal/httpapi/server.go b/apps/api/internal/httpapi/server.go index 04454e8..3802515 100644 --- a/apps/api/internal/httpapi/server.go +++ b/apps/api/internal/httpapi/server.go @@ -258,6 +258,8 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor mux.Handle("POST /api/v1/images/generations", server.requireUser(auth.PermissionBasic, server.createTask("images.generations", false))) mux.Handle("POST /api/v1/images/edits", server.requireUser(auth.PermissionBasic, server.createTask("images.edits", false))) mux.Handle("POST /api/v1/videos/generations", server.requireUser(auth.PermissionBasic, server.createTask("videos.generations", false))) + mux.Handle("POST /api/v1/video/generations", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.createLegacyVolcesVideoGeneration))) + mux.Handle("GET /api/v1/ai/result/{taskID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getLegacyVolcesVideoResult))) mux.Handle("POST /api/v1/song/generations", server.requireUser(auth.PermissionBasic, server.createTask("song.generations", true))) mux.Handle("POST /api/v1/music/generations", server.requireUser(auth.PermissionBasic, server.createTask("music.generations", true))) mux.Handle("POST /api/v1/speech/generations", server.requireUser(auth.PermissionBasic, server.createTask("speech.generations", true))) @@ -265,6 +267,14 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor mux.Handle("GET /api/v1/voice_clone/voices", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices))) mux.Handle("DELETE /api/v1/voice_clone/voices/{voiceID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice))) mux.Handle("POST /api/v1/files/upload", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.uploadFile))) + mux.Handle("GET /api/v1/resource/material/seedance-portrait-assets/capability", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getSeedancePortraitAssetCapability))) + mux.Handle("GET /api/v1/resource/material/user/materials", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listSeedancePortraitAssets))) + mux.Handle("POST /api/v1/resource/material", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.createSeedancePortraitAsset))) + mux.Handle("POST /api/v1/resource/material/seedance-portrait-assets/sync", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.syncSeedancePortraitAssets))) + mux.Handle("POST /api/v3/contents/generations/tasks", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.createVolcesContentsGenerationTask))) + mux.Handle("GET /api/v3/contents/generations/tasks", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listVolcesContentsGenerationTasks))) + mux.Handle("GET /api/v3/contents/generations/tasks/{taskID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getVolcesContentsGenerationTask))) + mux.Handle("DELETE /api/v3/contents/generations/tasks/{taskID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.deleteVolcesContentsGenerationTask))) server.registerGeminiGenerateContentRoutes(mux) server.registerKlingCompatibilityRoutes(mux) mux.Handle("POST /upload/{version}/files", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.geminiFilesUpload))) diff --git a/apps/api/internal/httpapi/volces_compat_handlers.go b/apps/api/internal/httpapi/volces_compat_handlers.go new file mode 100644 index 0000000..f089c77 --- /dev/null +++ b/apps/api/internal/httpapi/volces_compat_handlers.go @@ -0,0 +1,330 @@ +package httpapi + +import ( + "encoding/json" + "errors" + "net/http" + "strconv" + "strings" + + "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" + "github.com/easyai/easyai-ai-gateway/apps/api/internal/clients" + "github.com/easyai/easyai-ai-gateway/apps/api/internal/runner" + "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" +) + +const volcesContentsCompatibilityMarker = "volces_contents_generations_v3" + +// createVolcesContentsGenerationTask godoc +// @Summary 创建火山内容生成任务 +// @Description 兼容火山方舟 POST /api/v3/contents/generations/tasks。网关 task id 是查询与取消用的公开 id;上游 id 另以 upstream_task_id 保留。 +// @Tags volces-compatible +// @Accept json +// @Produce json +// @Security BearerAuth +// @Success 200 {object} map[string]any +// @Router /api/v3/contents/generations/tasks [post] +func (s *Server) createVolcesContentsGenerationTask(w http.ResponseWriter, r *http.Request) { + user, ok := auth.UserFromContext(r.Context()) + if !ok || user == nil { + writeError(w, http.StatusUnauthorized, "unauthorized") + return + } + body, err := s.decodeTaskRequestBody(r.Context(), w, r, "videos.generations") + if err != nil { + writeError(w, http.StatusBadRequest, err.Error(), clients.ErrorCode(err)) + return + } + task, err := s.createVolcesCompatibleTask(r, user, body) + if err != nil { + writeVolcesCompatibleTaskError(w, err) + return + } + writeJSON(w, http.StatusOK, volcesCompatibleTask(task)) +} + +// getVolcesContentsGenerationTask godoc +// @Summary 查询火山内容生成任务 +// @Tags volces-compatible +// @Produce json +// @Security BearerAuth +// @Success 200 {object} map[string]any +// @Router /api/v3/contents/generations/tasks/{taskID} [get] +func (s *Server) getVolcesContentsGenerationTask(w http.ResponseWriter, r *http.Request) { + task, ok := s.volcesCompatibleTaskForUser(w, r) + if !ok { + return + } + writeJSON(w, http.StatusOK, volcesCompatibleTask(task)) +} + +// listVolcesContentsGenerationTasks godoc +// @Summary 列出火山内容生成任务 +// @Tags volces-compatible +// @Produce json +// @Security BearerAuth +// @Success 200 {object} map[string]any +// @Router /api/v3/contents/generations/tasks [get] +func (s *Server) listVolcesContentsGenerationTasks(w http.ResponseWriter, r *http.Request) { + user, ok := auth.UserFromContext(r.Context()) + if !ok || user == nil { + writeError(w, http.StatusUnauthorized, "unauthorized") + return + } + page := portraitAssetQueryInt(r, "page_num", "pageNumber", "page") + pageSize := portraitAssetQueryInt(r, "page_size", "pageSize") + tasks, err := s.store.ListVolcesCompatibleTasks(r.Context(), user, store.VolcesCompatibleTaskListFilter{ + CompatibilityMarker: volcesContentsCompatibilityMarker, + Status: r.URL.Query().Get("filter.status"), + Model: r.URL.Query().Get("filter.model"), + TaskIDs: r.URL.Query()["filter.task_ids"], + Page: page, + PageSize: pageSize, + }) + if err != nil { + s.logger.Error("list Volces-compatible tasks failed", "error", err) + writeError(w, http.StatusInternalServerError, "list tasks failed") + return + } + items := make([]any, 0) + for _, task := range tasks.Items { + items = append(items, volcesCompatibleTask(task)) + } + writeJSON(w, http.StatusOK, map[string]any{ + "items": items, "total": tasks.Total, + "page_num": tasks.Page, "page_size": tasks.PageSize, + // data/page are retained as additive gateway fields for existing callers. + "data": items, "page": tasks.Page, + }) +} + +// deleteVolcesContentsGenerationTask godoc +// @Summary 取消火山内容生成任务 +// @Description 取消网关任务;对于已提交且保存了上游任务标识的 Volces 视频任务,同时调用火山 DELETE 接口并持久化取消状态。 +// @Tags volces-compatible +// @Produce json +// @Security BearerAuth +// @Success 200 {object} map[string]any +// @Router /api/v3/contents/generations/tasks/{taskID} [delete] +func (s *Server) deleteVolcesContentsGenerationTask(w http.ResponseWriter, r *http.Request) { + task, ok := s.volcesCompatibleTaskForUser(w, r) + if !ok { + return + } + user, _ := auth.UserFromContext(r.Context()) + result, err := s.runner.CancelVolcesVideoTask(r.Context(), task, user) + if err != nil { + if errors.Is(err, runner.ErrTaskAccessDenied) { + writeError(w, http.StatusNotFound, "task not found") + return + } + s.logger.Error("cancel Volces-compatible task failed", "error", err) + writeError(w, http.StatusInternalServerError, "cancel task failed") + return + } + updated, err := s.store.GetTask(r.Context(), task.ID) + if err != nil { + writeError(w, http.StatusInternalServerError, "get cancelled task failed") + return + } + response := volcesCompatibleTask(updated) + response["cancelled"] = result.Cancelled + response["cancellable"] = result.Cancellable + response["message"] = result.Message + writeJSON(w, http.StatusOK, response) +} + +// createLegacyVolcesVideoGeneration godoc +// @Summary 创建 server-main 兼容视频任务 +// @Description 兼容 server-main 的 /api/v1/video/generations,返回 submitted 和 task_id;额外保留火山任务字段。 +// @Tags volces-compatible +// @Accept json +// @Produce json +// @Security BearerAuth +// @Success 200 {object} map[string]any +// @Router /api/v1/video/generations [post] +func (s *Server) createLegacyVolcesVideoGeneration(w http.ResponseWriter, r *http.Request) { + user, ok := auth.UserFromContext(r.Context()) + if !ok || user == nil { + writeError(w, http.StatusUnauthorized, "unauthorized") + return + } + body, err := s.decodeTaskRequestBody(r.Context(), w, r, "videos.generations") + if err != nil { + writeError(w, http.StatusBadRequest, err.Error(), clients.ErrorCode(err)) + return + } + task, err := s.createVolcesCompatibleTask(r, user, body) + if err != nil { + writeVolcesCompatibleTaskError(w, err) + return + } + response := volcesCompatibleTask(task) + response["status"] = "submitted" + response["task_id"] = task.ID + writeJSON(w, http.StatusOK, response) +} + +// getLegacyVolcesVideoResult godoc +// @Summary 查询 server-main 兼容视频结果 +// @Tags volces-compatible +// @Produce json +// @Security BearerAuth +// @Success 200 {object} map[string]any +// @Router /api/v1/ai/result/{taskID} [get] +func (s *Server) getLegacyVolcesVideoResult(w http.ResponseWriter, r *http.Request) { + task, ok := s.volcesCompatibleTaskForUser(w, r) + if !ok { + return + } + compat := volcesCompatibleTask(task) + legacyStatus := "process" + switch compat["status"] { + case "succeeded": + legacyStatus = "success" + case "failed", "cancelled": + legacyStatus = "failed" + } + writeJSON(w, http.StatusOK, map[string]any{ + "status": legacyStatus, "task_id": task.ID, "data": compat["content"], "result": compat, + }) +} + +func (s *Server) createVolcesCompatibleTask(r *http.Request, user *auth.User, body map[string]any) (store.GatewayTask, error) { + model := strings.TrimSpace(volcesCompatString(body["model"])) + if model == "" { + return store.GatewayTask{}, &clients.ClientError{Code: "invalid_parameter", Message: "model is required", StatusCode: http.StatusBadRequest, Retryable: false} + } + if !apiKeyScopeAllowed(user, "videos.generations") { + return store.GatewayTask{}, &clients.ClientError{Code: "forbidden", Message: "api key scope does not allow video generation", StatusCode: http.StatusForbidden, Retryable: false} + } + body["_gateway_compatibility"] = volcesContentsCompatibilityMarker + task, err := s.prepareAndCreateGatewayTask(r.Context(), r, user, "videos.generations", model, body, true) + if err != nil { + return store.GatewayTask{}, err + } + if err := s.runner.EnqueueAsyncTask(r.Context(), task); err != nil { + return store.GatewayTask{}, &clients.ClientError{Code: "enqueue_failed", Message: err.Error(), StatusCode: http.StatusServiceUnavailable, Retryable: true} + } + return task, nil +} + +func (s *Server) volcesCompatibleTaskForUser(w http.ResponseWriter, r *http.Request) (store.GatewayTask, bool) { + user, ok := auth.UserFromContext(r.Context()) + if !ok || user == nil { + writeError(w, http.StatusUnauthorized, "unauthorized") + return store.GatewayTask{}, false + } + task, err := s.store.GetTask(r.Context(), strings.TrimSpace(r.PathValue("taskID"))) + if err != nil { + if store.IsNotFound(err) { + writeError(w, http.StatusNotFound, "task not found") + return store.GatewayTask{}, false + } + s.logger.Error("get Volces-compatible task failed", "error", err) + writeError(w, http.StatusInternalServerError, "get task failed") + return store.GatewayTask{}, false + } + if !isVolcesCompatibleTask(task) || !kelingCompatTaskOwnedBy(task, user) { + writeError(w, http.StatusNotFound, "task not found") + return store.GatewayTask{}, false + } + return task, true +} + +func isVolcesCompatibleTask(task store.GatewayTask) bool { + return task.Kind == "videos.generations" && strings.TrimSpace(volcesCompatString(task.Request["_gateway_compatibility"])) == volcesContentsCompatibilityMarker +} + +func volcesCompatibleTask(task store.GatewayTask) map[string]any { + response := cloneVolcesCompatibleMap(task.Result) + if len(response) == 0 { + response = cloneVolcesCompatibleMap(task.RemoteTaskPayload) + } + if response == nil { + response = map[string]any{} + } + response["id"] = task.ID + response["model"] = firstNonEmpty(volcesCompatString(response["model"]), task.Model) + response["status"] = volcesCompatibleTaskStatus(task.Status) + response["created_at"] = task.CreatedAt.Unix() + response["updated_at"] = task.UpdatedAt.Unix() + if task.RemoteTaskID != "" { + response["upstream_task_id"] = task.RemoteTaskID + } + for _, key := range []string{"content", "seed", "resolution", "ratio", "duration", "frames", "framespersecond"} { + if response[key] == nil && task.Request[key] != nil { + response[key] = task.Request[key] + } + } + if len(task.Usage) > 0 && response["usage"] == nil { + response["usage"] = task.Usage + } + if task.Status == "failed" || task.Status == "cancelled" { + response["error"] = map[string]any{"code": firstNonEmpty(task.ErrorCode, strings.ToUpper(task.Status)), "message": firstNonEmpty(task.ErrorMessage, task.Error, task.Message)} + } + response["gateway_task_id"] = task.ID + response["gateway_status"] = task.Status + response["billings"] = task.Billings + response["billing_summary"] = task.BillingSummary + response["final_charge_amount"] = task.FinalChargeAmount + return response +} + +func volcesCompatibleTaskStatus(status string) string { + switch strings.ToLower(strings.TrimSpace(status)) { + case "succeeded", "success", "completed": + return "succeeded" + case "failed": + return "failed" + case "cancelled", "canceled": + return "cancelled" + case "running", "processing": + return "running" + default: + return "queued" + } +} + +func cloneVolcesCompatibleMap(source map[string]any) map[string]any { + if len(source) == 0 { + return nil + } + raw, err := json.Marshal(source) + if err != nil { + return map[string]any{} + } + var out map[string]any + if err := json.Unmarshal(raw, &out); err != nil { + return map[string]any{} + } + return out +} + +func writeVolcesCompatibleTaskError(w http.ResponseWriter, err error) { + status := http.StatusInternalServerError + var staged *gatewayTaskCreationError + if errors.As(err, &staged) { + err = staged.Err + } + var clientErr *clients.ClientError + if errors.As(err, &clientErr) && clientErr.StatusCode > 0 { + status = clientErr.StatusCode + } else if errors.As(err, &clientErr) { + status = http.StatusBadRequest + } + writeError(w, status, err.Error(), clients.ErrorCode(err)) +} + +func volcesCompatString(value any) string { + switch typed := value.(type) { + case string: + return strings.TrimSpace(typed) + case json.Number: + return typed.String() + case float64: + return strconv.FormatFloat(typed, 'f', -1, 64) + default: + return "" + } +} diff --git a/apps/api/internal/httpapi/volces_compat_handlers_test.go b/apps/api/internal/httpapi/volces_compat_handlers_test.go new file mode 100644 index 0000000..2069843 --- /dev/null +++ b/apps/api/internal/httpapi/volces_compat_handlers_test.go @@ -0,0 +1,29 @@ +package httpapi + +import ( + "testing" + "time" + + "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" +) + +func TestVolcesCompatibleTaskPreservesOfficialFieldsAndGatewayBilling(t *testing.T) { + now := time.Date(2026, 7, 18, 8, 0, 0, 0, time.UTC) + task := store.GatewayTask{ + ID: "gateway-task-1", Kind: "videos.generations", Status: "succeeded", Model: "doubao-seedance-2-0-mini-260615", + RemoteTaskID: "cgt-upstream-1", CreatedAt: now, UpdatedAt: now.Add(time.Second), + Result: map[string]any{ + "id": "cgt-upstream-1", "model": "doubao-seedance-2-0-mini-260615", "status": "succeeded", + "content": map[string]any{"video_url": "https://example.com/out.mp4"}, "usage": map[string]any{"total_tokens": 9}, + }, + Billings: []any{map[string]any{"amount": 3}}, BillingSummary: map[string]any{"currency": "resource"}, FinalChargeAmount: 3, + } + got := volcesCompatibleTask(task) + if got["id"] != task.ID || got["upstream_task_id"] != task.RemoteTaskID || got["status"] != "succeeded" { + t.Fatalf("unexpected compatibility identity/status: %+v", got) + } + content, _ := got["content"].(map[string]any) + if content["video_url"] != "https://example.com/out.mp4" || got["usage"] == nil || got["billings"] == nil { + t.Fatalf("official or billing fields were lost: %+v", got) + } +} diff --git a/apps/api/internal/runner/service.go b/apps/api/internal/runner/service.go index d6fafe4..2211ccf 100644 --- a/apps/api/internal/runner/service.go +++ b/apps/api/internal/runner/service.go @@ -526,6 +526,20 @@ candidatesLoop: candidateBody := preprocessing.Body candidatePricing := pricingByCandidate[pricingCandidateKey(candidate)] response, err := s.runCandidate(ctx, task, user, candidateBody, preprocessing.Log, candidate, candidatePricing, nextAttemptNo, onDelta, responseExecution, singleSourceProtected, runnerPolicy.CacheAffinityPolicy, cacheAffinityKeys.Record) + if err != nil && isVolcesRemoteTaskCancellation(candidate, err) { + cancelled, changed, cancelErr := s.store.CancelSubmittedTask(ctx, task.ID, task.ExecutionToken, "任务已由火山引擎取消") + if cancelErr != nil { + return Result{}, cancelErr + } + if changed { + // CancelSubmittedTask atomically transfers any reservation to the release Outbox. + walletReservationFinalized = true + if emitErr := s.emit(ctx, task.ID, "task.cancelled", "cancelled", "cancelled", 1, "任务已由火山引擎取消", map[string]any{"taskId": task.ID, "reason": "upstream_cancelled"}, isSimulation(task, candidate)); emitErr != nil { + return Result{}, emitErr + } + return Result{Task: cancelled, Output: cancelled.Result}, nil + } + } if err == nil { attemptNo = nextAttemptNo var billings []any @@ -592,6 +606,13 @@ candidatesLoop: ResponseDurationMS: record.ResponseDurationMS, }) if finishErr != nil { + if errors.Is(finishErr, store.ErrTaskExecutionLeaseLost) { + latest, latestErr := s.store.GetTask(ctx, task.ID) + if latestErr == nil && latest.Status == "cancelled" { + walletReservationFinalized = true + return Result{Task: latest, Output: latest.Result}, nil + } + } return Result{}, finishErr } walletReservationFinalized = true @@ -965,6 +986,12 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user } return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload) }, + OnRemoteTaskPolled: func(remoteTaskID string, payload map[string]any) error { + if strings.TrimSpace(remoteTaskID) == "" { + return nil + } + return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload) + }, Stream: boolFromMap(providerBody, "stream"), StreamDelta: onDelta, UpstreamProtocol: candidate.ResponseProtocol, @@ -1199,12 +1226,19 @@ func (s *Service) failTask(ctx context.Context, taskID string, executionToken st if err != nil { return store.GatewayTask{}, err } + if failed.Status == "cancelled" { + return failed, nil + } if eventErr := s.emit(ctx, taskID, "task.failed", "failed", "failed", 1, message, map[string]any{"code": code, "requestId": requestID, "metrics": metrics}, simulated); eventErr != nil { return store.GatewayTask{}, eventErr } return failed, nil } +func isVolcesRemoteTaskCancellation(candidate store.RuntimeModelCandidate, err error) bool { + return isVolcesCancellationCandidate(candidate) && strings.EqualFold(clients.ErrorCode(err), "volces_task_cancelled") +} + type failedAttemptRecord struct { Task store.GatewayTask Body map[string]any diff --git a/apps/api/internal/runner/task_cancel.go b/apps/api/internal/runner/task_cancel.go index 600f778..4f18d37 100644 --- a/apps/api/internal/runner/task_cancel.go +++ b/apps/api/internal/runner/task_cancel.go @@ -6,6 +6,7 @@ import ( "strings" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" + "github.com/easyai/easyai-ai-gateway/apps/api/internal/clients" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/riverqueue/river/rivertype" ) @@ -104,6 +105,61 @@ func (s *Service) CancelTask(ctx context.Context, taskID string, user *auth.User }, nil } +// CancelVolcesVideoTask extends local queue cancellation with the official +// Volces DELETE call once a video task has a persisted remote task id. +func (s *Service) CancelVolcesVideoTask(ctx context.Context, task store.GatewayTask, user *auth.User) (TaskCancelResult, error) { + local, err := s.CancelTask(ctx, task.ID, user) + if err != nil || local.Cancelled || strings.TrimSpace(task.RemoteTaskID) == "" { + return local, err + } + if taskCancelTerminalStatus(task.Status) { + return local, nil + } + var latest store.TaskAttempt + for _, attempt := range task.Attempts { + if attempt.PlatformModelID != "" && (latest.AttemptNo == 0 || attempt.AttemptNo >= latest.AttemptNo) { + latest = attempt + } + } + candidate, found, err := s.store.GetRuntimeModelCandidateForRemoteTask(ctx, latest.PlatformModelID, latest.PlatformID) + if err != nil { + return TaskCancelResult{}, err + } + if !found || !isVolcesCancellationCandidate(candidate) { + return local, nil + } + httpClient, err := s.httpClientForCandidate(candidate, false) + if err != nil { + return TaskCancelResult{}, err + } + _, _, err = (clients.VolcesClient{HTTPClient: httpClient}).DeleteVideoTask(ctx, clients.Request{ + Kind: "videos.generations", Candidate: candidate, HTTPClient: httpClient, RemoteTaskID: task.RemoteTaskID, + }) + if err != nil { + return TaskCancelResult{}, err + } + cancelledTask, cancelled, err := s.store.CancelSubmittedTask(ctx, task.ID, task.ExecutionToken, "任务已由火山引擎取消") + if err != nil { + return TaskCancelResult{}, err + } + if !cancelled { + latestTask, latestErr := s.store.GetTask(ctx, task.ID) + if latestErr == nil { + return taskCancelUnavailable(latestTask, "任务状态已变化,未覆盖本地最终状态"), nil + } + return local, nil + } + if err := s.emit(ctx, cancelledTask.ID, "task.cancelled", "cancelled", "cancelled", 1, "任务已由火山引擎取消", map[string]any{"taskId": cancelledTask.ID, "reason": "upstream_cancel"}, cancelledTask.RunMode == "simulation"); err != nil { + return TaskCancelResult{}, err + } + return TaskCancelResult{TaskID: cancelledTask.ID, Cancelled: true, Cancellable: true, Submitted: true, Message: "任务已由火山引擎取消"}, nil +} + +func isVolcesCancellationCandidate(candidate store.RuntimeModelCandidate) bool { + provider := strings.ToLower(strings.TrimSpace(candidate.Provider)) + return provider == "volces" || provider == "volces-openai" +} + func taskCancelUnavailable(task store.GatewayTask, message string) TaskCancelResult { return TaskCancelResult{ TaskID: task.ID, diff --git a/apps/api/internal/store/remote_task_candidate.go b/apps/api/internal/store/remote_task_candidate.go new file mode 100644 index 0000000..8be9ea9 --- /dev/null +++ b/apps/api/internal/store/remote_task_candidate.go @@ -0,0 +1,43 @@ +package store + +import ( + "context" + "strings" +) + +// GetRuntimeModelCandidateForRemoteTask restores the exact platform model used +// to submit an asynchronous provider task. It deliberately ignores enabled +// state so a task can still be cancelled after its platform is disabled. +func (s *Store) GetRuntimeModelCandidateForRemoteTask(ctx context.Context, platformModelID string, platformID string) (RuntimeModelCandidate, bool, error) { + platformModelID = strings.TrimSpace(platformModelID) + platformID = strings.TrimSpace(platformID) + if platformModelID == "" || platformID == "" { + return RuntimeModelCandidate{}, false, nil + } + var candidate RuntimeModelCandidate + var credentials, config []byte + err := s.pool.QueryRow(ctx, ` +SELECT p.id::text, p.platform_key, p.name, p.provider, + COALESCE(NULLIF(p.config->>'specType', ''), p.provider), COALESCE(p.base_url, ''), p.auth_type, + p.credentials, p.config, m.id::text, COALESCE(NULLIF(m.provider_model_name, ''), m.model_name), + m.model_name, COALESCE(m.model_alias, ''), + COALESCE((m.model_type->>0), 'video_generate') +FROM platform_models m +JOIN integration_platforms p ON p.id = m.platform_id +WHERE m.id = $1::uuid AND p.id = $2::uuid AND p.deleted_at IS NULL`, platformModelID, platformID).Scan( + &candidate.PlatformID, &candidate.PlatformKey, &candidate.PlatformName, &candidate.Provider, + &candidate.SpecType, &candidate.BaseURL, &candidate.AuthType, &credentials, &config, + &candidate.PlatformModelID, &candidate.ProviderModelName, &candidate.ModelName, &candidate.ModelAlias, &candidate.ModelType, + ) + if IsNotFound(err) { + return RuntimeModelCandidate{}, false, nil + } + if err != nil { + return RuntimeModelCandidate{}, false, err + } + candidate.Credentials = decodeObject(credentials) + candidate.PlatformConfig = decodeObject(config) + candidate.ClientID = candidate.PlatformKey + ":" + candidate.ModelType + ":" + firstNonEmpty(candidate.ProviderModelName, candidate.ModelName) + candidate.QueueKey = candidate.ClientID + return candidate, true, nil +} diff --git a/apps/api/internal/store/tasks_runtime.go b/apps/api/internal/store/tasks_runtime.go index 3d6d804..d350497 100644 --- a/apps/api/internal/store/tasks_runtime.go +++ b/apps/api/internal/store/tasks_runtime.go @@ -530,7 +530,8 @@ WHERE id = $1::uuid UPDATE gateway_task_attempts SET remote_task_id = NULLIF($2::text, ''), response_snapshot = COALESCE(response_snapshot, '{}'::jsonb) || jsonb_build_object('remote_task_payload', $3::jsonb) -WHERE id = $1::uuid`, +WHERE id = $1::uuid + AND status = 'running'`, attemptID, remoteTaskID, string(payloadJSON), @@ -585,6 +586,72 @@ WHERE id = $1::uuid return task, true, nil } +// CancelSubmittedTask records a confirmed upstream cancellation. Callers must +// first complete the provider-side DELETE so local status never claims a remote +// task was cancelled when the upstream request was not accepted. +func (s *Store) CancelSubmittedTask(ctx context.Context, taskID string, executionToken string, message string) (GatewayTask, bool, error) { + message = strings.TrimSpace(message) + if message == "" { + message = "任务已由上游取消" + } + var task GatewayTask + changed := false + err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + var err error + task, err = scanGatewayTask(tx.QueryRow(ctx, ` +UPDATE gateway_tasks +SET status = 'cancelled', + error = NULLIF($2, ''), + error_code = 'task_cancelled', + error_message = NULLIF($2, ''), + billing_status = CASE + WHEN run_mode <> 'production' OR gateway_user_id IS NULL THEN 'not_required' + WHEN reservation_amount > 0 THEN 'pending' + ELSE 'released' + END, + billing_updated_at = now(), + locked_by = NULL, + locked_at = NULL, + heartbeat_at = NULL, + execution_token = NULL, + execution_lease_expires_at = NULL, + finished_at = now(), + updated_at = now() +WHERE id = $1::uuid + AND NULLIF(remote_task_id, '') IS NOT NULL + AND ( + (status = 'running' AND execution_token = NULLIF($3, '')::uuid) + OR status = 'queued' + ) +RETURNING `+gatewayTaskColumns, taskID, message, strings.TrimSpace(executionToken))) + if IsNotFound(err) { + return nil + } + if err != nil { + return err + } + changed = true + payloadJSON, _ := json.Marshal(map[string]any{"taskId": taskID, "reason": "upstream_cancelled"}) + _, err = tx.Exec(ctx, ` +INSERT INTO settlement_outbox ( + task_id, event_type, action, amount, currency, pricing_snapshot, payload, status, next_attempt_at +) +SELECT id, 'task.billing.release', 'release', reservation_amount, billing_currency, + pricing_snapshot, $2::jsonb, 'pending', now() +FROM gateway_tasks +WHERE id = $1::uuid + AND run_mode = 'production' + AND gateway_user_id IS NOT NULL + AND reservation_amount > 0 +ON CONFLICT (task_id, event_type) DO NOTHING`, taskID, string(payloadJSON)) + return err + }) + if err != nil { + return GatewayTask{}, false, err + } + return task, changed, nil +} + func (s *Store) ListRecoverableAsyncTasks(ctx context.Context, limit int) ([]AsyncTaskQueueItem, error) { if limit <= 0 { limit = 500 diff --git a/apps/api/internal/store/volces_compatible_tasks.go b/apps/api/internal/store/volces_compatible_tasks.go new file mode 100644 index 0000000..5293ae3 --- /dev/null +++ b/apps/api/internal/store/volces_compatible_tasks.go @@ -0,0 +1,111 @@ +package store + +import ( + "context" + "strings" + + "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" +) + +// VolcesCompatibleTaskListFilter mirrors the supported filters of Ark's +// ListContentsGenerationsTasks API. Task IDs are the gateway's public task +// IDs, which are the IDs returned by the compatibility create endpoint. +type VolcesCompatibleTaskListFilter struct { + CompatibilityMarker string + Status string + Model string + TaskIDs []string + Page int + PageSize int +} + +// ListVolcesCompatibleTasks returns only video tasks created through a named +// compatibility surface. Keeping this query separate from ListTasks avoids +// broadening the ordinary task-list API's filtering semantics. +func (s *Store) ListVolcesCompatibleTasks(ctx context.Context, user *auth.User, filter VolcesCompatibleTaskListFilter) (TaskListResult, error) { + page := filter.Page + if page < 1 { + page = 1 + } + if page > 500 { + page = 500 + } + pageSize := filter.PageSize + if pageSize < 1 { + pageSize = 20 + } + if pageSize > 500 { + pageSize = 500 + } + gatewayUserID := localGatewayUserID(user) + userID, apiKeyID := "", "" + if user != nil { + userID = strings.TrimSpace(user.ID) + apiKeyID = strings.TrimSpace(user.APIKeyID) + } + if gatewayUserID == "" && userID == "" { + return TaskListResult{}, ErrLocalUserRequired + } + taskIDs := make([]string, 0, len(filter.TaskIDs)) + seen := make(map[string]bool, len(filter.TaskIDs)) + for _, taskID := range filter.TaskIDs { + taskID = strings.TrimSpace(taskID) + if taskID != "" && !seen[taskID] { + seen[taskID] = true + taskIDs = append(taskIDs, taskID) + } + } + args := []any{ + gatewayUserID, + userID, + apiKeyID, + strings.TrimSpace(filter.CompatibilityMarker), + strings.ToLower(strings.TrimSpace(filter.Status)), + strings.TrimSpace(filter.Model), + taskIDs, + } + whereSQL := ` +WHERE ( + ( + NULLIF($1, '')::uuid IS NOT NULL + AND gateway_user_id = NULLIF($1, '')::uuid + ) + OR ( + NULLIF($1, '')::uuid IS NULL + AND NULLIF($2, '') IS NOT NULL + AND user_id = $2 + ) + ) + AND (NULLIF($3, '') IS NULL OR api_key_id = $3) + AND kind = 'videos.generations' + AND request->>'_gateway_compatibility' = $4 + AND (NULLIF($5, '') IS NULL OR LOWER(status) = $5) + AND (NULLIF($6, '') IS NULL OR model = $6 OR resolved_model = $6) + AND (COALESCE(array_length($7::text[], 1), 0) = 0 OR id::text = ANY($7::text[]))` + var total int + if err := s.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_tasks `+whereSQL, args...).Scan(&total); err != nil { + return TaskListResult{}, err + } + rows, err := s.pool.Query(ctx, ` +SELECT `+gatewayTaskColumns+` +FROM gateway_tasks +`+whereSQL+` +ORDER BY created_at DESC +LIMIT $8 OFFSET $9`, append(args, pageSize, (page-1)*pageSize)...) + if err != nil { + return TaskListResult{}, err + } + defer rows.Close() + items := make([]GatewayTask, 0) + for rows.Next() { + task, err := scanGatewayTask(rows) + if err != nil { + return TaskListResult{}, err + } + items = append(items, task) + } + if err := rows.Err(); err != nil { + return TaskListResult{}, err + } + return TaskListResult{Items: items, Total: total, Page: page, PageSize: pageSize}, nil +}