From 05b5ff1ee617716bffc4d3ed2a503e1c4f6b02f5 Mon Sep 17 00:00:00 2001 From: "shujian.o" Date: Tue, 7 Jul 2026 15:46:50 +0800 Subject: [PATCH] fix task polling result urls and reset locking --- model/subscription.go | 4 +-- service/task_polling.go | 9 ++++--- service/task_polling_test.go | 51 ++++++++++++++++++++++++++++++++++++ 3 files changed, 59 insertions(+), 5 deletions(-) diff --git a/model/subscription.go b/model/subscription.go index b8ae97746ee8..642c7b0355c6 100644 --- a/model/subscription.go +++ b/model/subscription.go @@ -1027,7 +1027,7 @@ func adminResetUserSubscriptionsByPlanTx(tx *gorm.DB, userId int, plan *Subscrip return nil, errors.New("invalid reset args") } var subs []UserSubscription - if err := tx.Set("gorm:query_option", "FOR UPDATE"). + if err := lockForUpdate(tx). Where("user_id = ? AND plan_id = ? AND status = ? AND end_time > ?", userId, plan.Id, "active", now). Order("end_time asc, id asc"). Find(&subs).Error; err != nil { @@ -1049,7 +1049,7 @@ func adminResetPlanSubscriptionsTx(tx *gorm.DB, plan *SubscriptionPlan, now int6 return nil, errors.New("invalid reset args") } var subs []UserSubscription - if err := tx.Set("gorm:query_option", "FOR UPDATE"). + if err := lockForUpdate(tx). Where("plan_id = ? AND status = ? AND end_time > ?", plan.Id, "active", now). Order("user_id asc, end_time asc, id asc"). Find(&subs).Error; err != nil { diff --git a/service/task_polling.go b/service/task_polling.go index 179fa7342082..d8186e328457 100644 --- a/service/task_polling.go +++ b/service/task_polling.go @@ -468,13 +468,16 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch * taskResult := &relaycommon.TaskInfo{} // try parse as New API response format - var responseItems dto.TaskResponse[model.Task] + var responseItems dto.TaskResponse[dto.TaskDto] if err = common.Unmarshal(responseBody, &responseItems); err == nil && responseItems.IsSuccess() { logger.LogDebug(ctx, "updateVideoSingleTask parsed as new api response format: %+v", responseItems) t := responseItems.Data taskResult.TaskID = t.TaskID - taskResult.Status = string(t.Status) - taskResult.Url = t.GetResultURL() + taskResult.Status = t.Status + taskResult.Url = strings.TrimSpace(t.ResultURL) + if taskResult.Url == "" && model.TaskStatus(t.Status) == model.TaskStatusSuccess { + taskResult.Url = strings.TrimSpace(t.FailReason) + } taskResult.Progress = t.Progress taskResult.Reason = t.FailReason task.Data = t.Data diff --git a/service/task_polling_test.go b/service/task_polling_test.go index 3164d0a9e29c..41020c8c391c 100644 --- a/service/task_polling_test.go +++ b/service/task_polling_test.go @@ -22,6 +22,7 @@ import ( type taskPollingFetchAdaptor struct { mu sync.Mutex taskIDs []string + responseBody map[string][]byte fetched chan string blockTaskID string blockStarted chan struct{} @@ -51,6 +52,12 @@ func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, body map[string] default: } } + if body, ok := a.responseBody[taskID]; ok { + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + }, nil + } response := dto.TaskResponse[model.Task]{ Code: dto.TaskSuccessCode, @@ -125,6 +132,50 @@ func seedPollingTask(t *testing.T, channelID int, publicID string, upstreamID st return task } +func TestUpdateVideoTasksPreservesNewAPITaskDtoResultURL(t *testing.T) { + truncate(t) + + const channelID = 1001 + const resultURL = "https://cdn.example.com/videos/result.mp4" + seedTaskPollingChannel(t, channelID, true) + task := seedPollingTask(t, channelID, "task_public_result_url", "upstream_result_url") + response := dto.TaskResponse[dto.TaskDto]{ + Code: dto.TaskSuccessCode, + Data: dto.TaskDto{ + TaskID: task.TaskID, + Status: string(model.TaskStatusSuccess), + ResultURL: resultURL, + Progress: "100%", + }, + } + body, err := common.Marshal(response) + require.NoError(t, err) + + adaptor := &taskPollingFetchAdaptor{ + responseBody: map[string][]byte{ + task.GetUpstreamTaskID(): body, + }, + } + previousFactory := GetTaskAdaptorFunc + GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor } + t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory }) + + err = UpdateVideoTasks(context.Background(), constant.TaskPlatform("kling"), map[int][]string{ + channelID: { + task.GetUpstreamTaskID(), + }, + }, map[string]*model.Task{ + task.GetUpstreamTaskID(): task, + }) + + require.NoError(t, err) + var saved model.Task + require.NoError(t, model.DB.Where("task_id = ?", task.TaskID).First(&saved).Error) + assert.Equal(t, model.TaskStatus(model.TaskStatusSuccess), saved.Status) + assert.Equal(t, resultURL, saved.PrivateData.ResultURL) + assert.Equal(t, "100%", saved.Progress) +} + func TestUpdateVideoTasksDefaultSleepWaitsBetweenTasks(t *testing.T) { truncate(t)