feat(volces): 接入任务兼容查询与上游取消
This commit is contained in:
@@ -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) {
|
func TestVolcesClientVideoRejectsDuplicateFirstFrameBeforeSubmit(t *testing.T) {
|
||||||
var submitted bool
|
var submitted bool
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ type Request struct {
|
|||||||
RemoteTaskID string
|
RemoteTaskID string
|
||||||
RemoteTaskPayload map[string]any
|
RemoteTaskPayload map[string]any
|
||||||
OnRemoteTaskSubmitted func(remoteTaskID string, payload map[string]any) error
|
OnRemoteTaskSubmitted func(remoteTaskID string, payload map[string]any) error
|
||||||
|
OnRemoteTaskPolled func(remoteTaskID string, payload map[string]any) error
|
||||||
Stream bool
|
Stream bool
|
||||||
StreamDelta StreamDelta
|
StreamDelta StreamDelta
|
||||||
UpstreamProtocol string
|
UpstreamProtocol string
|
||||||
|
|||||||
@@ -100,66 +100,105 @@ func (c VolcesClient) runVideo(ctx context.Context, request Request, apiKey stri
|
|||||||
timeout := volcesPollTimeout(request)
|
timeout := volcesPollTimeout(request)
|
||||||
deadline := time.NewTimer(timeout)
|
deadline := time.NewTimer(timeout)
|
||||||
defer deadline.Stop()
|
defer deadline.Stop()
|
||||||
|
nextPoll := time.NewTimer(0)
|
||||||
ticker := time.NewTicker(interval)
|
defer nextPoll.Stop()
|
||||||
defer ticker.Stop()
|
|
||||||
|
|
||||||
var lastResult map[string]any
|
var lastResult map[string]any
|
||||||
|
lastRequestID := firstNonEmpty(submitRequestID, upstreamTaskID)
|
||||||
|
transientFailures := 0
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return Response{}, &ClientError{Code: "cancelled", Message: ctx.Err().Error(), RequestID: submitRequestID, Retryable: true}
|
return Response{}, &ClientError{Code: "cancelled", Message: ctx.Err().Error(), RequestID: lastRequestID, 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}
|
|
||||||
case <-deadline.C:
|
case <-deadline.C:
|
||||||
return Response{}, &ClientError{
|
return Response{}, &ClientError{
|
||||||
Code: "timeout",
|
Code: "timeout",
|
||||||
Message: fmt.Sprintf("volces video task %s did not finish before timeout; last status: %s", upstreamTaskID, volcesTaskStatus(lastResult)),
|
Message: fmt.Sprintf("volces video task %s did not finish before timeout; last status: %s", upstreamTaskID, volcesTaskStatus(lastResult)),
|
||||||
RequestID: requestID,
|
RequestID: lastRequestID,
|
||||||
Retryable: true,
|
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 {
|
func volcesVideoTaskPath(request Request) string {
|
||||||
path := firstNonEmptyStringValue(
|
path := firstNonEmptyStringValue(
|
||||||
request.Candidate.PlatformConfig,
|
request.Candidate.PlatformConfig,
|
||||||
@@ -997,6 +1036,9 @@ func volcesTaskErrorCode(result map[string]any) string {
|
|||||||
return code
|
return code
|
||||||
}
|
}
|
||||||
status := volcesTaskStatus(result)
|
status := volcesTaskStatus(result)
|
||||||
|
if status == "cancelled" {
|
||||||
|
return "volces_task_cancelled"
|
||||||
|
}
|
||||||
if status != "" {
|
if status != "" {
|
||||||
return 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 {
|
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)
|
content, _ := raw["content"].(map[string]any)
|
||||||
videoURL := strings.TrimSpace(stringFromAny(content["video_url"]))
|
videoURL := strings.TrimSpace(stringFromAny(content["video_url"]))
|
||||||
created := intFromAny(raw["created_at"])
|
created := intFromAny(raw["created_at"])
|
||||||
@@ -1025,16 +1071,17 @@ func volcesVideoSuccessResult(request Request, upstreamTaskID string, raw map[st
|
|||||||
if videoURL != "" {
|
if videoURL != "" {
|
||||||
data = append(data, map[string]any{"url": videoURL, "type": "video"})
|
data = append(data, map[string]any{"url": videoURL, "type": "video"})
|
||||||
}
|
}
|
||||||
return map[string]any{
|
result["id"] = firstNonEmpty(stringFromAny(raw["id"]), upstreamTaskID)
|
||||||
"id": upstreamTaskID,
|
if strings.TrimSpace(stringFromAny(result["model"])) == "" {
|
||||||
"object": "video.generation",
|
result["model"] = upstreamModelName(request.Candidate)
|
||||||
"created": created,
|
|
||||||
"model": upstreamModelName(request.Candidate),
|
|
||||||
"status": "succeeded",
|
|
||||||
"upstream_task_id": upstreamTaskID,
|
|
||||||
"data": data,
|
|
||||||
"raw": raw,
|
|
||||||
}
|
}
|
||||||
|
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 {
|
func volcesVideoUsage(raw map[string]any) Usage {
|
||||||
@@ -1074,6 +1121,37 @@ func volcesPollTimeout(request Request) time.Duration {
|
|||||||
return time.Duration(seconds) * time.Second
|
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 {
|
func firstNonEmpty(values ...string) string {
|
||||||
for _, value := range values {
|
for _, value := range values {
|
||||||
if strings.TrimSpace(value) != "" {
|
if strings.TrimSpace(value) != "" {
|
||||||
|
|||||||
@@ -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/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/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/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/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/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)))
|
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("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("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("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.registerGeminiGenerateContentRoutes(mux)
|
||||||
server.registerKlingCompatibilityRoutes(mux)
|
server.registerKlingCompatibilityRoutes(mux)
|
||||||
mux.Handle("POST /upload/{version}/files", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.geminiFilesUpload)))
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -526,6 +526,20 @@ candidatesLoop:
|
|||||||
candidateBody := preprocessing.Body
|
candidateBody := preprocessing.Body
|
||||||
candidatePricing := pricingByCandidate[pricingCandidateKey(candidate)]
|
candidatePricing := pricingByCandidate[pricingCandidateKey(candidate)]
|
||||||
response, err := s.runCandidate(ctx, task, user, candidateBody, preprocessing.Log, candidate, candidatePricing, nextAttemptNo, onDelta, responseExecution, singleSourceProtected, runnerPolicy.CacheAffinityPolicy, cacheAffinityKeys.Record)
|
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 {
|
if err == nil {
|
||||||
attemptNo = nextAttemptNo
|
attemptNo = nextAttemptNo
|
||||||
var billings []any
|
var billings []any
|
||||||
@@ -592,6 +606,13 @@ candidatesLoop:
|
|||||||
ResponseDurationMS: record.ResponseDurationMS,
|
ResponseDurationMS: record.ResponseDurationMS,
|
||||||
})
|
})
|
||||||
if finishErr != nil {
|
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
|
return Result{}, finishErr
|
||||||
}
|
}
|
||||||
walletReservationFinalized = true
|
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)
|
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"),
|
Stream: boolFromMap(providerBody, "stream"),
|
||||||
StreamDelta: onDelta,
|
StreamDelta: onDelta,
|
||||||
UpstreamProtocol: candidate.ResponseProtocol,
|
UpstreamProtocol: candidate.ResponseProtocol,
|
||||||
@@ -1199,12 +1226,19 @@ func (s *Service) failTask(ctx context.Context, taskID string, executionToken st
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return store.GatewayTask{}, err
|
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 {
|
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 store.GatewayTask{}, eventErr
|
||||||
}
|
}
|
||||||
return failed, nil
|
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 {
|
type failedAttemptRecord struct {
|
||||||
Task store.GatewayTask
|
Task store.GatewayTask
|
||||||
Body map[string]any
|
Body map[string]any
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"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/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
"github.com/riverqueue/river/rivertype"
|
"github.com/riverqueue/river/rivertype"
|
||||||
)
|
)
|
||||||
@@ -104,6 +105,61 @@ func (s *Service) CancelTask(ctx context.Context, taskID string, user *auth.User
|
|||||||
}, nil
|
}, 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 {
|
func taskCancelUnavailable(task store.GatewayTask, message string) TaskCancelResult {
|
||||||
return TaskCancelResult{
|
return TaskCancelResult{
|
||||||
TaskID: task.ID,
|
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
|
||||||
|
}
|
||||||
@@ -530,7 +530,8 @@ WHERE id = $1::uuid
|
|||||||
UPDATE gateway_task_attempts
|
UPDATE gateway_task_attempts
|
||||||
SET remote_task_id = NULLIF($2::text, ''),
|
SET remote_task_id = NULLIF($2::text, ''),
|
||||||
response_snapshot = COALESCE(response_snapshot, '{}'::jsonb) || jsonb_build_object('remote_task_payload', $3::jsonb)
|
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,
|
attemptID,
|
||||||
remoteTaskID,
|
remoteTaskID,
|
||||||
string(payloadJSON),
|
string(payloadJSON),
|
||||||
@@ -585,6 +586,72 @@ WHERE id = $1::uuid
|
|||||||
return task, true, nil
|
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) {
|
func (s *Store) ListRecoverableAsyncTasks(ctx context.Context, limit int) ([]AsyncTaskQueueItem, error) {
|
||||||
if limit <= 0 {
|
if limit <= 0 {
|
||||||
limit = 500
|
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
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user