Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions model/subscription.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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 {
Expand Down
9 changes: 6 additions & 3 deletions service/task_polling.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
51 changes: 51 additions & 0 deletions service/task_polling_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{}
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)

Expand Down