Commit a98e207e by Seefs Committed by GitHub

feat: add ali wan video (#2141)

* feat: add ali wan video

* refactor: use same UnmarshalBodyReusable

* feat: enhance request body metadata

* feat: opt wan convertToOpenAIVideo

* feat: add wan support other param via json metadata

* refactor: remove unused code

* fix ali

---------

Co-authored-by: feitianbubu <feitianbubu@qq.com>
parent 36b71243
...@@ -2,7 +2,6 @@ package common ...@@ -2,7 +2,6 @@ package common
import ( import (
"bytes" "bytes"
"encoding/json"
"io" "io"
"mime/multipart" "mime/multipart"
"net/http" "net/http"
...@@ -41,11 +40,11 @@ func UnmarshalBodyReusable(c *gin.Context, v any) error { ...@@ -41,11 +40,11 @@ func UnmarshalBodyReusable(c *gin.Context, v any) error {
//} //}
contentType := c.Request.Header.Get("Content-Type") contentType := c.Request.Header.Get("Content-Type")
if strings.HasPrefix(contentType, "application/json") { if strings.HasPrefix(contentType, "application/json") {
err = Unmarshal(requestBody, &v) err = Unmarshal(requestBody, v)
} else if strings.Contains(contentType, gin.MIMEPOSTForm) { } else if strings.Contains(contentType, gin.MIMEPOSTForm) {
err = parseFormData(requestBody, &v) err = parseFormData(requestBody, v)
} else if strings.Contains(contentType, gin.MIMEMultipartPOSTForm) { } else if strings.Contains(contentType, gin.MIMEMultipartPOSTForm) {
err = parseMultipartFormData(c, requestBody, &v) err = parseMultipartFormData(c, requestBody, v)
} else { } else {
// skip for now // skip for now
// TODO: someday non json request have variant model, we will need to implementation this // TODO: someday non json request have variant model, we will need to implementation this
...@@ -145,6 +144,20 @@ func ParseMultipartFormReusable(c *gin.Context) (*multipart.Form, error) { ...@@ -145,6 +144,20 @@ func ParseMultipartFormReusable(c *gin.Context) (*multipart.Form, error) {
return form, nil return form, nil
} }
func processFormMap(formMap map[string]any, v any) error {
jsonData, err := Marshal(formMap)
if err != nil {
return err
}
err = Unmarshal(jsonData, v)
if err != nil {
return err
}
return nil
}
func parseFormData(data []byte, v any) error { func parseFormData(data []byte, v any) error {
values, err := url.ParseQuery(string(data)) values, err := url.ParseQuery(string(data))
if err != nil { if err != nil {
...@@ -158,12 +171,8 @@ func parseFormData(data []byte, v any) error { ...@@ -158,12 +171,8 @@ func parseFormData(data []byte, v any) error {
formMap[key] = vals formMap[key] = vals
} }
} }
jsonData, err := json.Marshal(formMap)
if err != nil {
return err
}
return Unmarshal(jsonData, v) return processFormMap(formMap, v)
} }
func parseMultipartFormData(c *gin.Context, data []byte, v any) error { func parseMultipartFormData(c *gin.Context, data []byte, v any) error {
...@@ -191,10 +200,6 @@ func parseMultipartFormData(c *gin.Context, data []byte, v any) error { ...@@ -191,10 +200,6 @@ func parseMultipartFormData(c *gin.Context, data []byte, v any) error {
formMap[key] = vals formMap[key] = vals
} }
} }
jsonData, err := Marshal(formMap)
if err != nil {
return err
}
return Unmarshal(jsonData, v) return processFormMap(formMap, v)
} }
...@@ -91,7 +91,8 @@ func VideoProxy(c *gin.Context) { ...@@ -91,7 +91,8 @@ func VideoProxy(c *gin.Context) {
return return
} }
if channel.Type == constant.ChannelTypeGemini { switch channel.Type {
case constant.ChannelTypeGemini:
apiKey := task.PrivateData.Key apiKey := task.PrivateData.Key
if apiKey == "" { if apiKey == "" {
logger.LogError(c.Request.Context(), fmt.Sprintf("Missing stored API key for Gemini task %s", taskID)) logger.LogError(c.Request.Context(), fmt.Sprintf("Missing stored API key for Gemini task %s", taskID))
...@@ -116,7 +117,10 @@ func VideoProxy(c *gin.Context) { ...@@ -116,7 +117,10 @@ func VideoProxy(c *gin.Context) {
return return
} }
req.Header.Set("x-goog-api-key", apiKey) req.Header.Set("x-goog-api-key", apiKey)
} else { case constant.ChannelTypeAli:
// Video URL is directly in task.FailReason
videoURL = task.FailReason
default:
// Default (Sora, etc.): Use original logic // Default (Sora, etc.): Use original logic
videoURL = fmt.Sprintf("%s/v1/videos/%s/content", baseURL, task.TaskID) videoURL = fmt.Sprintf("%s/v1/videos/%s/content", baseURL, task.TaskID)
req.Header.Set("Authorization", "Bearer "+channel.Key) req.Header.Set("Authorization", "Bearer "+channel.Key)
......
...@@ -27,7 +27,7 @@ type OpenAIVideo struct { ...@@ -27,7 +27,7 @@ type OpenAIVideo struct {
Size string `json:"size,omitempty"` Size string `json:"size,omitempty"`
RemixedFromVideoID string `json:"remixed_from_video_id,omitempty"` RemixedFromVideoID string `json:"remixed_from_video_id,omitempty"`
Error *OpenAIVideoError `json:"error,omitempty"` Error *OpenAIVideoError `json:"error,omitempty"`
Metadata map[string]any `json:"meta_data,omitempty"` Metadata map[string]any `json:"metadata,omitempty"`
} }
func (m *OpenAIVideo) SetProgressStr(progress string) { func (m *OpenAIVideo) SetProgressStr(progress string) {
......
...@@ -73,20 +73,22 @@ func (t *Task) GetData(v any) error { ...@@ -73,20 +73,22 @@ func (t *Task) GetData(v any) error {
} }
type Properties struct { type Properties struct {
Input string `json:"input"` Input string `json:"input"`
UpstreamModelName string `json:"upstream_model_name,omitempty"`
OriginModelName string `json:"origin_model_name,omitempty"`
} }
func (m *Properties) Scan(val interface{}) error { func (m *Properties) Scan(val interface{}) error {
bytesValue, _ := val.([]byte) bytesValue, _ := val.([]byte)
if len(bytesValue) == 0 { if len(bytesValue) == 0 {
m.Input = "" *m = Properties{}
return nil return nil
} }
return json.Unmarshal(bytesValue, m) return json.Unmarshal(bytesValue, m)
} }
func (m Properties) Value() (driver.Value, error) { func (m Properties) Value() (driver.Value, error) {
if m.Input == "" { if m == (Properties{}) {
return nil, nil return nil, nil
} }
return json.Marshal(m) return json.Marshal(m)
...@@ -127,8 +129,16 @@ type SyncTaskQueryParams struct { ...@@ -127,8 +129,16 @@ type SyncTaskQueryParams struct {
func InitTask(platform constant.TaskPlatform, relayInfo *commonRelay.RelayInfo) *Task { func InitTask(platform constant.TaskPlatform, relayInfo *commonRelay.RelayInfo) *Task {
properties := Properties{} properties := Properties{}
privateData := TaskPrivateData{} privateData := TaskPrivateData{}
if relayInfo != nil && relayInfo.ChannelMeta != nil && relayInfo.ChannelMeta.ChannelType == constant.ChannelTypeGemini { if relayInfo != nil && relayInfo.ChannelMeta != nil {
privateData.Key = relayInfo.ChannelMeta.ApiKey if relayInfo.ChannelMeta.ChannelType == constant.ChannelTypeGemini {
privateData.Key = relayInfo.ChannelMeta.ApiKey
}
if relayInfo.UpstreamModelName != "" {
properties.UpstreamModelName = relayInfo.UpstreamModelName
}
if relayInfo.OriginModelName != "" {
properties.OriginModelName = relayInfo.OriginModelName
}
} }
t := &Task{ t := &Task{
......
package ali
var ModelList = []string{
"wan2.5-i2v-preview", // 万相2.5 preview(有声视频)推荐
"wan2.2-i2v-flash", // 万相2.2极速版(无声视频)
"wan2.2-i2v-plus", // 万相2.2专业版(无声视频)
"wanx2.1-i2v-plus", // 万相2.1专业版(无声视频)
"wanx2.1-i2v-turbo", // 万相2.1极速版(无声视频)
}
var ChannelName = "ali"
package common package common
import ( import (
"encoding/json"
"errors" "errors"
"fmt" "fmt"
"strings" "strings"
...@@ -485,14 +486,16 @@ type TaskRelayInfo struct { ...@@ -485,14 +486,16 @@ type TaskRelayInfo struct {
} }
type TaskSubmitReq struct { type TaskSubmitReq struct {
Prompt string `json:"prompt"` Prompt string `json:"prompt"`
Model string `json:"model,omitempty"` Model string `json:"model,omitempty"`
Mode string `json:"mode,omitempty"` Mode string `json:"mode,omitempty"`
Image string `json:"image,omitempty"` Image string `json:"image,omitempty"`
Images []string `json:"images,omitempty"` Images []string `json:"images,omitempty"`
Size string `json:"size,omitempty"` Size string `json:"size,omitempty"`
Duration int `json:"duration,omitempty"` Duration int `json:"duration,omitempty"`
Metadata map[string]interface{} `json:"metadata,omitempty"` Seconds string `json:"seconds,omitempty"`
InputReference string `json:"input_reference,omitempty"`
Metadata map[string]interface{} `json:"metadata,omitempty"`
} }
func (t TaskSubmitReq) GetPrompt() string { func (t TaskSubmitReq) GetPrompt() string {
...@@ -503,6 +506,38 @@ func (t TaskSubmitReq) HasImage() bool { ...@@ -503,6 +506,38 @@ func (t TaskSubmitReq) HasImage() bool {
return len(t.Images) > 0 return len(t.Images) > 0
} }
func (t *TaskSubmitReq) UnmarshalJSON(data []byte) error {
type Alias TaskSubmitReq
aux := &struct {
Metadata json.RawMessage `json:"metadata,omitempty"`
*Alias
}{
Alias: (*Alias)(t),
}
if err := common.Unmarshal(data, &aux); err != nil {
return err
}
if len(aux.Metadata) > 0 {
var metadataStr string
if err := common.Unmarshal(aux.Metadata, &metadataStr); err == nil && metadataStr != "" {
var metadataObj map[string]interface{}
if err := common.Unmarshal([]byte(metadataStr), &metadataObj); err == nil {
t.Metadata = metadataObj
return nil
}
}
var metadataObj map[string]interface{}
if err := common.Unmarshal(aux.Metadata, &metadataObj); err == nil {
t.Metadata = metadataObj
}
}
return nil
}
type TaskInfo struct { type TaskInfo struct {
Code int `json:"code"` Code int `json:"code"`
TaskID string `json:"task_id"` TaskID string `json:"task_id"`
......
...@@ -108,62 +108,33 @@ func validateMultipartTaskRequest(c *gin.Context, info *RelayInfo, action string ...@@ -108,62 +108,33 @@ func validateMultipartTaskRequest(c *gin.Context, info *RelayInfo, action string
} }
func ValidateMultipartDirect(c *gin.Context, info *RelayInfo) *dto.TaskError { func ValidateMultipartDirect(c *gin.Context, info *RelayInfo) *dto.TaskError {
contentType := c.GetHeader("Content-Type")
var prompt string var prompt string
var model string var model string
var seconds int var seconds int
var size string var size string
var hasInputReference bool var hasInputReference bool
if strings.HasPrefix(contentType, "multipart/form-data") { var req TaskSubmitReq
form, err := common.ParseMultipartFormReusable(c) if err := common.UnmarshalBodyReusable(c, &req); err != nil {
if err != nil { return createTaskError(err, "invalid_json", http.StatusBadRequest, true)
return createTaskError(err, "invalid_multipart_form", http.StatusBadRequest, true) }
}
defer form.RemoveAll()
prompts, ok := form.Value["prompt"]
if !ok || len(prompts) == 0 {
return createTaskError(fmt.Errorf("prompt field is required"), "missing_prompt", http.StatusBadRequest, true)
}
prompt = prompts[0]
if _, ok := form.Value["model"]; !ok {
return createTaskError(fmt.Errorf("model field is required"), "missing_model", http.StatusBadRequest, true)
}
model = form.Value["model"][0]
if _, ok := form.File["input_reference"]; ok {
hasInputReference = true
}
if ss, ok := form.Value["seconds"]; ok {
sInt := common.String2Int(ss[0])
if sInt > seconds {
seconds = common.String2Int(ss[0])
}
}
if sz, ok := form.Value["size"]; ok {
size = sz[0]
}
} else {
var req TaskSubmitReq
if err := common.UnmarshalBodyReusable(c, &req); err != nil {
return createTaskError(err, "invalid_json", http.StatusBadRequest, true)
}
prompt = req.Prompt prompt = req.Prompt
model = req.Model model = req.Model
seconds, _ = strconv.Atoi(req.Seconds)
if seconds == 0 {
seconds = req.Duration seconds = req.Duration
}
if req.InputReference != "" {
req.Images = []string{req.InputReference}
}
if strings.TrimSpace(req.Model) == "" { if strings.TrimSpace(req.Model) == "" {
return createTaskError(fmt.Errorf("model field is required"), "missing_model", http.StatusBadRequest, true) return createTaskError(fmt.Errorf("model field is required"), "missing_model", http.StatusBadRequest, true)
} }
if req.HasImage() { if req.HasImage() {
hasInputReference = true hasInputReference = true
}
} }
if taskErr := validatePrompt(prompt); taskErr != nil { if taskErr := validatePrompt(prompt); taskErr != nil {
......
...@@ -28,6 +28,7 @@ import ( ...@@ -28,6 +28,7 @@ import (
"github.com/QuantumNous/new-api/relay/channel/perplexity" "github.com/QuantumNous/new-api/relay/channel/perplexity"
"github.com/QuantumNous/new-api/relay/channel/siliconflow" "github.com/QuantumNous/new-api/relay/channel/siliconflow"
"github.com/QuantumNous/new-api/relay/channel/submodel" "github.com/QuantumNous/new-api/relay/channel/submodel"
taskali "github.com/QuantumNous/new-api/relay/channel/task/ali"
taskdoubao "github.com/QuantumNous/new-api/relay/channel/task/doubao" taskdoubao "github.com/QuantumNous/new-api/relay/channel/task/doubao"
taskGemini "github.com/QuantumNous/new-api/relay/channel/task/gemini" taskGemini "github.com/QuantumNous/new-api/relay/channel/task/gemini"
taskjimeng "github.com/QuantumNous/new-api/relay/channel/task/jimeng" taskjimeng "github.com/QuantumNous/new-api/relay/channel/task/jimeng"
...@@ -133,6 +134,8 @@ func GetTaskAdaptor(platform constant.TaskPlatform) channel.TaskAdaptor { ...@@ -133,6 +134,8 @@ func GetTaskAdaptor(platform constant.TaskPlatform) channel.TaskAdaptor {
} }
if channelType, err := strconv.ParseInt(string(platform), 10, 64); err == nil { if channelType, err := strconv.ParseInt(string(platform), 10, 64); err == nil {
switch channelType { switch channelType {
case constant.ChannelTypeAli:
return &taskali.TaskAdaptor{}
case constant.ChannelTypeKling: case constant.ChannelTypeKling:
return &kling.TaskAdaptor{} return &kling.TaskAdaptor{}
case constant.ChannelTypeJimeng: case constant.ChannelTypeJimeng:
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment