diff --git a/common/gin.go b/common/gin.go index 315e86613ec7..8fb67104933d 100644 --- a/common/gin.go +++ b/common/gin.go @@ -19,6 +19,7 @@ import ( const KeyRequestBody = "key_request_body" const KeyBodyStorage = "key_body_storage" +const KeySeedanceOfficialAPI = "seedance_official_api" var ErrRequestBodyTooLarge = errors.New("request body too large") diff --git a/controller/seedance.go b/controller/seedance.go new file mode 100644 index 000000000000..f240ea4bf8fc --- /dev/null +++ b/controller/seedance.go @@ -0,0 +1,24 @@ +package controller + +import ( + "net/http" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/relay" + "github.com/gin-gonic/gin" +) + +func RelaySeedanceTask(c *gin.Context) { + c.Set(common.KeySeedanceOfficialAPI, true) + RelayTask(c) +} + +func RelaySeedanceTaskFetch(c *gin.Context) { + c.Set(common.KeySeedanceOfficialAPI, true) + respBody, taskErr := relay.SeedanceTaskFetch(c) + if taskErr != nil { + respondTaskError(c, taskErr) + return + } + c.Data(http.StatusOK, "application/json", respBody) +} diff --git a/middleware/seedance_adapter.go b/middleware/seedance_adapter.go new file mode 100644 index 000000000000..5ddbac7189f0 --- /dev/null +++ b/middleware/seedance_adapter.go @@ -0,0 +1,23 @@ +package middleware + +import ( + "net/http" + + "github.com/QuantumNous/new-api/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + + "github.com/gin-gonic/gin" +) + +func SeedanceRequestConvert() func(c *gin.Context) { + return func(c *gin.Context) { + c.Set(common.KeySeedanceOfficialAPI, true) + + if c.Request.Method == http.MethodPost { + c.Request.URL.Path = "/v1/video/generations" + c.Set("relay_mode", relayconstant.RelayModeVideoSubmit) + } + + c.Next() + } +} diff --git a/middleware/seedance_adapter_test.go b/middleware/seedance_adapter_test.go new file mode 100644 index 000000000000..6b8d1a19e977 --- /dev/null +++ b/middleware/seedance_adapter_test.go @@ -0,0 +1,65 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestSeedanceRequestConvertSubmit(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(SeedanceRequestConvert()) + router.POST("/seedance/api/v3/contents/generations/tasks", func(c *gin.Context) { + require.True(t, c.GetBool(common.KeySeedanceOfficialAPI)) + require.Equal(t, "/v1/video/generations", c.Request.URL.Path) + require.Equal(t, relayconstant.RelayModeVideoSubmit, c.GetInt("relay_mode")) + c.Status(http.StatusNoContent) + }) + + req := httptest.NewRequest(http.MethodPost, "/seedance/api/v3/contents/generations/tasks", nil) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusNoContent, recorder.Code) +} + +func TestSeedanceRequestConvertFetchByID(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(SeedanceRequestConvert()) + router.GET("/seedance/api/v3/contents/generations/tasks/:task_id", func(c *gin.Context) { + require.True(t, c.GetBool(common.KeySeedanceOfficialAPI)) + require.Equal(t, "/seedance/api/v3/contents/generations/tasks/task_public", c.Request.URL.Path) + require.Equal(t, "task_public", c.Param("task_id")) + c.Status(http.StatusNoContent) + }) + + req := httptest.NewRequest(http.MethodGet, "/seedance/api/v3/contents/generations/tasks/task_public", nil) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusNoContent, recorder.Code) +} + +func TestSeedanceRequestConvertFetchList(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(SeedanceRequestConvert()) + router.GET("/seedance/api/v3/contents/generations/tasks", func(c *gin.Context) { + require.True(t, c.GetBool(common.KeySeedanceOfficialAPI)) + require.Equal(t, "/seedance/api/v3/contents/generations/tasks", c.Request.URL.Path) + c.Status(http.StatusNoContent) + }) + + req := httptest.NewRequest(http.MethodGet, "/seedance/api/v3/contents/generations/tasks", nil) + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusNoContent, recorder.Code) +} diff --git a/relay/channel/task/doubao/adaptor.go b/relay/channel/task/doubao/adaptor.go index a6dabb5f1086..119e42061172 100644 --- a/relay/channel/task/doubao/adaptor.go +++ b/relay/channel/task/doubao/adaptor.go @@ -6,6 +6,7 @@ import ( "io" "net/http" "strconv" + "strings" "time" "github.com/QuantumNous/new-api/common" @@ -115,6 +116,29 @@ func (a *TaskAdaptor) Init(info *relaycommon.RelayInfo) { // ValidateRequestAndSetAction parses body, validates fields and sets default action. func (a *TaskAdaptor) ValidateRequestAndSetAction(c *gin.Context, info *relaycommon.RelayInfo) (taskErr *dto.TaskError) { + if c.GetBool(common.KeySeedanceOfficialAPI) { + var body map[string]interface{} + if err := common.UnmarshalBodyReusable(c, &body); err != nil { + return service.TaskErrorWrapperLocal(err, "invalid_request", http.StatusBadRequest) + } + + modelName, _ := body["model"].(string) + if strings.TrimSpace(modelName) == "" { + return service.TaskErrorWrapperLocal(fmt.Errorf("field model is required"), "missing_model", http.StatusBadRequest) + } + if _, ok := body["content"]; !ok { + return service.TaskErrorWrapperLocal(fmt.Errorf("field content is required"), "missing_content", http.StatusBadRequest) + } + + info.Action = constant.TaskActionGenerate + c.Set("task_request", relaycommon.TaskSubmitReq{ + Model: modelName, + Prompt: seedanceTextPrompt(body), + Metadata: body, + }) + return nil + } + // Accept only POST /v1/video/generations as "generate" action. return relaycommon.ValidateBasicTaskRequest(c, info, constant.TaskActionGenerate) } @@ -177,6 +201,30 @@ func hasVideoInMetadata(metadata map[string]interface{}) bool { // BuildRequestBody converts request into Doubao specific format. func (a *TaskAdaptor) BuildRequestBody(c *gin.Context, info *relaycommon.RelayInfo) (io.Reader, error) { + if c.GetBool(common.KeySeedanceOfficialAPI) { + storage, err := common.GetBodyStorage(c) + if err != nil { + return nil, err + } + cachedBody, err := storage.Bytes() + if err != nil { + return nil, err + } + + var bodyMap map[string]interface{} + if err := common.Unmarshal(cachedBody, &bodyMap); err != nil { + return bytes.NewReader(cachedBody), nil + } + if info.UpstreamModelName != "" { + bodyMap["model"] = info.UpstreamModelName + } + data, err := common.Marshal(bodyMap) + if err != nil { + return nil, err + } + return bytes.NewReader(data), nil + } + req, err := relaycommon.GetTaskRequest(c) if err != nil { return nil, err @@ -224,6 +272,13 @@ func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *rela return } + if c.GetBool(common.KeySeedanceOfficialAPI) { + c.JSON(http.StatusOK, gin.H{ + "id": dResp.ID, + }) + return dResp.ID, responseBody, nil + } + ov := dto.NewOpenAIVideo() ov.ID = info.PublicTaskID ov.TaskID = info.PublicTaskID @@ -303,6 +358,27 @@ func (a *TaskAdaptor) convertToRequestPayload(req *relaycommon.TaskSubmitReq) (* return &r, nil } +func seedanceTextPrompt(body map[string]interface{}) string { + content, ok := body["content"].([]interface{}) + if !ok { + return "" + } + for _, item := range content { + itemMap, ok := item.(map[string]interface{}) + if !ok { + continue + } + if itemMap["type"] != "text" { + continue + } + text, _ := itemMap["text"].(string) + if strings.TrimSpace(text) != "" { + return text + } + } + return "" +} + func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, error) { resTask := responseTask{} if err := common.Unmarshal(respBody, &resTask); err != nil { @@ -332,6 +408,10 @@ func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, e taskResult.Status = model.TaskStatusFailure taskResult.Progress = "100%" taskResult.Reason = resTask.Error.Message + case "expired", "cancelled": + taskResult.Status = model.TaskStatusFailure + taskResult.Progress = "100%" + taskResult.Reason = resTask.Status default: // Unknown status, treat as processing taskResult.Status = model.TaskStatusInProgress diff --git a/relay/channel/task/doubao/adaptor_test.go b/relay/channel/task/doubao/adaptor_test.go new file mode 100644 index 000000000000..419cb146c3ae --- /dev/null +++ b/relay/channel/task/doubao/adaptor_test.go @@ -0,0 +1,78 @@ +package doubao + +import ( + "bytes" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSeedanceOfficialBuildRequestBodyPreservesNativeContent(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Set(common.KeySeedanceOfficialAPI, true) + ctx.Request = httptest.NewRequest(http.MethodPost, "/seedance/api/v3/contents/generations/tasks", bytes.NewBufferString(`{ + "model":"doubao-seedance-1-5-pro", + "content":[ + {"type":"image_url","image_url":{"url":"https://example.com/a.png"},"role":"first_frame"}, + {"type":"text","text":"make a video"} + ], + "duration":5, + "watermark":false + }`)) + ctx.Request.Header.Set("Content-Type", "application/json") + + adaptor := &TaskAdaptor{} + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "upstream-seedance-model"}, + } + + body, err := adaptor.BuildRequestBody(ctx, info) + require.NoError(t, err) + + raw, err := io.ReadAll(body) + require.NoError(t, err) + + var payload map[string]any + require.NoError(t, common.Unmarshal(raw, &payload)) + assert.Equal(t, "upstream-seedance-model", payload["model"]) + assert.Equal(t, float64(5), payload["duration"]) + assert.Equal(t, false, payload["watermark"]) + + content, ok := payload["content"].([]any) + require.True(t, ok) + require.Len(t, content, 2) + first, ok := content[0].(map[string]any) + require.True(t, ok) + assert.Equal(t, "image_url", first["type"]) + assert.Equal(t, "first_frame", first["role"]) +} + +func TestSeedanceOfficialDoResponseReturnsUpstreamTaskID(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Set(common.KeySeedanceOfficialAPI, true) + + adaptor := &TaskAdaptor{} + resp := &http.Response{ + Body: io.NopCloser(bytes.NewBufferString(`{"id":"cgt-upstream"}`)), + } + info := &relaycommon.RelayInfo{ + TaskRelayInfo: &relaycommon.TaskRelayInfo{PublicTaskID: "task_public"}, + } + + taskID, taskData, taskErr := adaptor.DoResponse(ctx, resp, info) + require.Nil(t, taskErr) + assert.Equal(t, "cgt-upstream", taskID) + assert.JSONEq(t, `{"id":"cgt-upstream"}`, string(taskData)) + assert.JSONEq(t, `{"id":"cgt-upstream"}`, recorder.Body.String()) +} diff --git a/relay/relay_task_seedance_test.go b/relay/relay_task_seedance_test.go new file mode 100644 index 000000000000..01c395abb81c --- /dev/null +++ b/relay/relay_task_seedance_test.go @@ -0,0 +1,54 @@ +package relay + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/model" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSeedanceTaskIDFilters(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest( + http.MethodGet, + "/seedance/api/v3/contents/generations/tasks?filter.task_ids=cgt-a,cgt-b&filter.task_ids=task_c&filter.task_ids[]=cgt-d", + nil, + ) + + require.Equal(t, []string{"cgt-a", "cgt-b", "task_c", "cgt-d"}, seedanceTaskIDFilters(ctx)) +} + +func TestSeedanceTaskResponseUsesUpstreamShape(t *testing.T) { + task := &model.Task{ + TaskID: "task_public", + Status: model.TaskStatusSuccess, + SubmitTime: 1710000000, + UpdatedAt: 1710000100, + Properties: model.Properties{ + OriginModelName: "doubao-seedance-1-5-pro", + }, + PrivateData: model.TaskPrivateData{ + UpstreamTaskID: "cgt-upstream", + ResultURL: "https://example.com/video.mp4", + }, + Data: json.RawMessage(`{"id":"cgt-upstream","status":"running","content":{},"service_tier":"default"}`), + } + + resp := seedanceTaskResponse(task) + assert.Equal(t, "cgt-upstream", resp["id"]) + assert.Equal(t, "doubao-seedance-1-5-pro", resp["model"]) + assert.Equal(t, "succeeded", resp["status"]) + assert.Equal(t, int64(1710000000), resp["created_at"]) + assert.Equal(t, int64(1710000100), resp["updated_at"]) + + content, ok := resp["content"].(map[string]any) + require.True(t, ok) + assert.Equal(t, "https://example.com/video.mp4", content["video_url"]) +} diff --git a/relay/seedance_task.go b/relay/seedance_task.go new file mode 100644 index 000000000000..6a6e0616747e --- /dev/null +++ b/relay/seedance_task.go @@ -0,0 +1,246 @@ +package relay + +import ( + "errors" + "net/http" + "strconv" + "strings" + "time" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/model" + "github.com/QuantumNous/new-api/service" + "github.com/gin-gonic/gin" +) + +func SeedanceTaskFetch(c *gin.Context) (respBody []byte, taskResp *dto.TaskError) { + taskID := strings.TrimSpace(c.Param("task_id")) + if taskID != "" { + return seedanceFetchTaskByID(c, taskID) + } + return seedanceFetchTaskList(c) +} + +func seedanceFetchTaskByID(c *gin.Context, taskID string) (respBody []byte, taskResp *dto.TaskError) { + originTask, exist, err := seedanceGetTaskByID(c.GetInt("id"), taskID) + if err != nil { + taskResp = service.TaskErrorWrapper(err, "get_task_failed", http.StatusInternalServerError) + return + } + if !exist { + taskResp = service.TaskErrorWrapperLocal(errors.New("task_not_exist"), "task_not_exist", http.StatusBadRequest) + return + } + + respBody, err = common.Marshal(seedanceTaskResponse(originTask)) + if err != nil { + taskResp = service.TaskErrorWrapper(err, "marshal_response_failed", http.StatusInternalServerError) + } + return +} + +func seedanceFetchTaskList(c *gin.Context) (respBody []byte, taskResp *dto.TaskError) { + pageNum := parseSeedancePositiveInt(c.Query("page_num"), 1, 500) + pageSize := parseSeedancePositiveInt(c.Query("page_size"), 20, 500) + offset := (pageNum - 1) * pageSize + + query := model.DB. + Where("user_id = ?", c.GetInt("id")). + Where("platform in ?", seedanceTaskPlatforms()). + Where("submit_time >= ?", time.Now().Add(-7*24*time.Hour).Unix()). + Order("id desc") + + var tasks []*model.Task + if err := query.Find(&tasks).Error; err != nil { + taskResp = service.TaskErrorWrapper(err, "get_tasks_failed", http.StatusInternalServerError) + return + } + + statusFilter := strings.TrimSpace(c.Query("filter.status")) + modelFilter := strings.TrimSpace(c.Query("filter.model")) + serviceTierFilter := strings.TrimSpace(c.Query("filter.service_tier")) + taskIDFilter := seedanceTaskIDFilters(c) + filtered := make([]*model.Task, 0, len(tasks)) + for _, task := range tasks { + if len(taskIDFilter) > 0 && !seedanceTaskMatchesID(task, taskIDFilter) { + continue + } + if statusFilter != "" && seedanceTaskStatus(task.Status) != statusFilter { + continue + } + if modelFilter != "" && !seedanceTaskMatchesModel(task, modelFilter) { + continue + } + if serviceTierFilter != "" && !seedanceTaskFieldEquals(task, "service_tier", serviceTierFilter) { + continue + } + filtered = append(filtered, task) + } + + total := len(filtered) + if offset > total { + filtered = []*model.Task{} + } else { + end := offset + pageSize + if end > total { + end = total + } + filtered = filtered[offset:end] + } + + items := make([]map[string]any, 0, len(filtered)) + for _, task := range filtered { + items = append(items, seedanceTaskResponse(task)) + } + + respBody, err := common.Marshal(map[string]any{ + "items": items, + "total": total, + }) + if err != nil { + taskResp = service.TaskErrorWrapper(err, "marshal_response_failed", http.StatusInternalServerError) + } + return +} + +func seedanceGetTaskByID(userID int, taskID string) (*model.Task, bool, error) { + task, exist, err := model.GetByTaskId(userID, taskID) + if err != nil || exist { + return task, exist, err + } + + var tasks []*model.Task + err = model.DB. + Where("user_id = ?", userID). + Where("platform in ?", seedanceTaskPlatforms()). + Where("submit_time >= ?", time.Now().Add(-7*24*time.Hour).Unix()). + Find(&tasks).Error + if err != nil { + return nil, false, err + } + for _, candidate := range tasks { + if candidate.GetUpstreamTaskID() == taskID { + return candidate, true, nil + } + } + return nil, false, nil +} + +func seedanceTaskPlatforms() []string { + return []string{ + strconv.Itoa(constant.ChannelTypeVolcEngine), + strconv.Itoa(constant.ChannelTypeDoubaoVideo), + } +} + +func seedanceTaskIDFilters(c *gin.Context) []string { + rawIDs := append(c.QueryArray("filter.task_ids"), c.QueryArray("filter.task_ids[]")...) + taskIDs := make([]string, 0, len(rawIDs)) + for _, rawID := range rawIDs { + for _, taskID := range strings.Split(rawID, ",") { + taskID = strings.TrimSpace(taskID) + if taskID != "" { + taskIDs = append(taskIDs, taskID) + } + } + } + return taskIDs +} + +func seedanceTaskMatchesID(task *model.Task, taskIDs []string) bool { + for _, taskID := range taskIDs { + if task.TaskID == taskID || task.GetUpstreamTaskID() == taskID { + return true + } + } + return false +} + +func parseSeedancePositiveInt(raw string, fallback, maxValue int) int { + value, err := strconv.Atoi(raw) + if err != nil || value <= 0 { + value = fallback + } + if value > maxValue { + return maxValue + } + return value +} + +func seedanceTaskMatchesModel(task *model.Task, modelName string) bool { + if task.Properties.OriginModelName == modelName || task.Properties.UpstreamModelName == modelName { + return true + } + return seedanceTaskFieldEquals(task, "model", modelName) +} + +func seedanceTaskFieldEquals(task *model.Task, field string, value string) bool { + var data map[string]any + if err := common.Unmarshal(task.Data, &data); err != nil { + return false + } + fieldValue, _ := data[field].(string) + return fieldValue == value +} + +func seedanceTaskResponse(task *model.Task) map[string]any { + resp := map[string]any{} + _ = common.Unmarshal(task.Data, &resp) + + resp["id"] = task.GetUpstreamTaskID() + if modelName := task.Properties.OriginModelName; modelName != "" { + resp["model"] = modelName + } else if modelName = task.Properties.UpstreamModelName; modelName != "" { + resp["model"] = modelName + } + resp["status"] = seedanceTaskStatus(task.Status) + + if createdAt := nonzeroSeedanceInt64(task.SubmitTime, task.CreatedAt); createdAt > 0 { + resp["created_at"] = createdAt + } + if task.UpdatedAt > 0 { + resp["updated_at"] = task.UpdatedAt + } + + if resultURL := task.GetResultURL(); resultURL != "" && task.Status == model.TaskStatusSuccess { + content, _ := resp["content"].(map[string]any) + if content == nil { + content = map[string]any{} + } + if _, ok := content["video_url"]; !ok { + content["video_url"] = resultURL + } + resp["content"] = content + } + + if task.Status == model.TaskStatusFailure && resp["error"] == nil && task.FailReason != "" { + resp["error"] = map[string]any{ + "message": task.FailReason, + } + } + return resp +} + +func nonzeroSeedanceInt64(values ...int64) int64 { + for _, value := range values { + if value != 0 { + return value + } + } + return 0 +} + +func seedanceTaskStatus(status model.TaskStatus) string { + switch status { + case model.TaskStatusSuccess: + return "succeeded" + case model.TaskStatusFailure: + return "failed" + case model.TaskStatusInProgress: + return "running" + default: + return "queued" + } +} diff --git a/router/video-router.go b/router/video-router.go index 461451104520..bbf29e7dcd8f 100644 --- a/router/video-router.go +++ b/router/video-router.go @@ -41,6 +41,18 @@ func SetVideoRouter(router *gin.Engine) { klingV1Router.GET("/videos/image2video/:task_id", controller.RelayTaskFetch) } + // Seedance official API routes - use a dedicated prefix like /kling to avoid + // changing response formats on the generic video endpoints. + seedanceOfficialGroup := router.Group("/seedance/api/v3/contents/generations") + seedanceOfficialGroup.Use(middleware.RouteTag("relay")) + seedanceOfficialGroup.Use(middleware.SeedanceRequestConvert(), middleware.TokenAuth()) + { + // Maps to: POST/GET /api/v3/contents/generations/tasks + seedanceOfficialGroup.POST("/tasks", middleware.Distribute(), controller.RelaySeedanceTask) + seedanceOfficialGroup.GET("/tasks", controller.RelaySeedanceTaskFetch) + seedanceOfficialGroup.GET("/tasks/:task_id", controller.RelaySeedanceTaskFetch) + } + // Jimeng official API routes - direct mapping to official API format jimengOfficialGroup := router.Group("jimeng") jimengOfficialGroup.Use(middleware.RouteTag("relay"))