Skip to content
Closed
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
2 changes: 2 additions & 0 deletions model/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -260,6 +260,7 @@ func migrateDB() error {
&Midjourney{},
&TopUp{},
&QuotaData{},
&TaskIDMapping{},
&Task{},
&Model{},
&Vendor{},
Expand Down Expand Up @@ -294,6 +295,7 @@ func migrateDBFast() error {
{&Midjourney{}, "Midjourney"},
{&TopUp{}, "TopUp"},
{&QuotaData{}, "QuotaData"},
{&TaskIDMapping{}, "TaskIDMapping"},
{&Task{}, "Task"},
{&Model{}, "Model"},
{&Vendor{}, "Vendor"},
Expand Down
78 changes: 78 additions & 0 deletions model/task_id_mapping.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
package model

import (
"errors"
"fmt"
"strings"

"github.com/google/uuid"
"gorm.io/gorm"
)

type TaskIDMapping struct {
ID int64 `json:"id" gorm:"primary_key;AUTO_INCREMENT"`

CreatedAt int64 `json:"created_at" gorm:"index"`

LocalTaskID string `json:"local_task_id" gorm:"type:varchar(191);uniqueIndex"`
UpstreamTaskID string `json:"upstream_task_id" gorm:"type:text"`
}

func CreateTaskIDMapping(localTaskID, upstreamTaskID string) error {
localTaskID = strings.TrimSpace(localTaskID)
upstreamTaskID = strings.TrimSpace(upstreamTaskID)
if localTaskID == "" || upstreamTaskID == "" {
return fmt.Errorf("invalid task id mapping")
}
return DB.Create(&TaskIDMapping{
LocalTaskID: localTaskID,
UpstreamTaskID: upstreamTaskID,
}).Error
}

func NewLocalTaskIDWithChannel(channelID int) string {
return fmt.Sprintf("tsk_%d_%s", channelID, strings.ReplaceAll(uuid.NewString(), "-", ""))
}

func CreateTaskIDMappingWithChannel(upstreamTaskID string, channelID int) (string, error) {
upstreamTaskID = strings.TrimSpace(upstreamTaskID)
if upstreamTaskID == "" {
return "", fmt.Errorf("invalid upstream task id")
}
localTaskID := NewLocalTaskIDWithChannel(channelID)
return localTaskID, CreateTaskIDMapping(localTaskID, upstreamTaskID)
}

func GetUpstreamTaskIDByLocalTaskID(localTaskID string) (string, bool, error) {
localTaskID = strings.TrimSpace(localTaskID)
if localTaskID == "" {
return "", false, nil
}

var mapping TaskIDMapping
err := DB.Where("local_task_id = ?", localTaskID).First(&mapping).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", false, nil
}
return "", false, err
}
return mapping.UpstreamTaskID, true, nil
}

func GetLocalTaskIDByUpstreamTaskID(upstreamTaskID string) (string, bool, error) {
upstreamTaskID = strings.TrimSpace(upstreamTaskID)
if upstreamTaskID == "" {
return "", false, nil
}

var mapping TaskIDMapping
err := DB.Where("upstream_task_id = ?", upstreamTaskID).First(&mapping).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", false, nil
}
return "", false, err
}
return mapping.LocalTaskID, true, nil
}
45 changes: 36 additions & 9 deletions relay/channel/task/gemini/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -181,7 +181,13 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela
if strings.TrimSpace(s.Name) == "" {
return "", nil, service.TaskErrorWrapper(fmt.Errorf("missing operation name"), "invalid_response", http.StatusInternalServerError)
}
taskID = encodeLocalTaskID(s.Name)

localTaskID, err := model.CreateTaskIDMappingWithChannel(s.Name, info.ChannelId)
if err != nil {
return "", nil, service.TaskErrorWrapper(err, "create_task_id_mapping_failed", http.StatusInternalServerError)
}
taskID = localTaskID

ov := dto.NewOpenAIVideo()
ov.ID = taskID
ov.TaskID = taskID
Expand All @@ -206,9 +212,18 @@ func (a *TaskAdaptor) FetchTask(baseUrl, key string, body map[string]any, proxy
return nil, fmt.Errorf("invalid task_id")
}

upstreamName, err := decodeLocalTaskID(taskID)
upstreamName, found, err := model.GetUpstreamTaskIDByLocalTaskID(taskID)
if err != nil {
return nil, fmt.Errorf("decode task_id failed: %w", err)
return nil, fmt.Errorf("resolve task_id mapping failed: %w", err)
}
if !found {
if strings.HasPrefix(taskID, "tsk_") {
return nil, fmt.Errorf("task_id mapping not found")
}
upstreamName, err = decodeLocalTaskID(taskID)
if err != nil {
return nil, fmt.Errorf("decode task_id failed: %w", err)
}
}

// For Gemini API, we use GET request to the operations endpoint
Expand Down Expand Up @@ -254,9 +269,12 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e
ti.Status = model.TaskStatusSuccess
ti.Progress = "100%"

taskID := encodeLocalTaskID(op.Name)
ti.TaskID = taskID
ti.Url = fmt.Sprintf("%s/v1/videos/%s/content", system_setting.ServerAddress, taskID)
if system_setting.ServerAddress != "" {
if localID, found, err := model.GetLocalTaskIDByUpstreamTaskID(op.Name); err == nil && found {
ti.TaskID = localID
ti.Url = fmt.Sprintf("%s/v1/videos/%s/content", system_setting.ServerAddress, localID)
}
}

// Extract URL from generateVideoResponse if available
if len(op.Response.GenerateVideoResponse.GeneratedSamples) > 0 {
Expand All @@ -269,9 +287,18 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e
}

func (a *TaskAdaptor) ConvertToOpenAIVideo(task *model.Task) ([]byte, error) {
upstreamName, err := decodeLocalTaskID(task.TaskID)
if err != nil {
upstreamName = ""
upstreamName := ""
if task != nil {
if mapped, found, err := model.GetUpstreamTaskIDByLocalTaskID(task.TaskID); err == nil && found {
upstreamName = mapped
} else {
if !strings.HasPrefix(task.TaskID, "tsk_") {
decoded, err := decodeLocalTaskID(task.TaskID)
if err == nil {
upstreamName = decoded
}
}
}
}
modelName := extractModelFromOperationName(upstreamName)
if strings.TrimSpace(modelName) == "" {
Expand Down
37 changes: 29 additions & 8 deletions relay/channel/task/vertex/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -211,9 +211,12 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela
if strings.TrimSpace(s.Name) == "" {
return "", nil, service.TaskErrorWrapper(fmt.Errorf("missing operation name"), "invalid_response", http.StatusInternalServerError)
}
localID := encodeLocalTaskID(s.Name)
c.JSON(http.StatusOK, gin.H{"task_id": localID})
return localID, responseBody, nil
localTaskID, err := model.CreateTaskIDMappingWithChannel(s.Name, info.ChannelId)
if err != nil {
return "", nil, service.TaskErrorWrapper(err, "create_task_id_mapping_failed", http.StatusInternalServerError)
}
c.JSON(http.StatusOK, gin.H{"task_id": localTaskID})
return localTaskID, responseBody, nil
}

func (a *TaskAdaptor) GetModelList() []string { return []string{"veo-3.0-generate-001"} }
Expand All @@ -225,9 +228,18 @@ func (a *TaskAdaptor) FetchTask(baseUrl, key string, body map[string]any, proxy
if !ok {
return nil, fmt.Errorf("invalid task_id")
}
upstreamName, err := decodeLocalTaskID(taskID)
upstreamName, found, err := model.GetUpstreamTaskIDByLocalTaskID(taskID)
if err != nil {
return nil, fmt.Errorf("decode task_id failed: %w", err)
return nil, fmt.Errorf("resolve task_id mapping failed: %w", err)
}
if !found {
if strings.HasPrefix(taskID, "tsk_") {
return nil, fmt.Errorf("task_id mapping not found")
}
upstreamName, err = decodeLocalTaskID(taskID)
if err != nil {
return nil, fmt.Errorf("decode task_id failed: %w", err)
}
}
region := extractRegionFromOperationName(upstreamName)
if region == "" {
Expand Down Expand Up @@ -338,9 +350,18 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e
}

func (a *TaskAdaptor) ConvertToOpenAIVideo(task *model.Task) ([]byte, error) {
upstreamName, err := decodeLocalTaskID(task.TaskID)
if err != nil {
upstreamName = ""
upstreamName := ""
if task != nil {
if mapped, found, err := model.GetUpstreamTaskIDByLocalTaskID(task.TaskID); err == nil && found {
upstreamName = mapped
} else {
if !strings.HasPrefix(task.TaskID, "tsk_") {
decoded, err := decodeLocalTaskID(task.TaskID)
if err == nil {
upstreamName = decoded
}
}
}
}
modelName := extractModelFromOperationName(upstreamName)
if strings.TrimSpace(modelName) == "" {
Expand Down