feat(volces): 接入任务兼容查询与上游取消

This commit is contained in:
2026-07-22 00:25:22 +08:00
parent ddd68cfebd
commit d7951cfdd2
11 changed files with 885 additions and 55 deletions
+71
View File
@@ -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) {
+1
View File
@@ -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
+132 -54
View File
@@ -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) != "" {
+10
View File
@@ -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)))
@@ -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 ""
}
}
@@ -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)
}
}
+34
View File
@@ -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
+56
View File
@@ -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,
@@ -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
}
+68 -1
View File
@@ -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
@@ -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
}