Não pode escolher mais do que 25 tópicos Os tópicos devem começar com uma letra ou um número, podem incluir traços ('-') e podem ter até 35 caracteres.
 
 
 

344 linhas
10 KiB

  1. package doubao
  2. import (
  3. "bytes"
  4. "fmt"
  5. "io"
  6. "net/http"
  7. "strconv"
  8. "strings"
  9. "time"
  10. "github.com/QuantumNous/new-api/common"
  11. "github.com/QuantumNous/new-api/constant"
  12. "github.com/QuantumNous/new-api/dto"
  13. "github.com/QuantumNous/new-api/model"
  14. "github.com/QuantumNous/new-api/relay/channel"
  15. taskcommon "github.com/QuantumNous/new-api/relay/channel/task/taskcommon"
  16. relaycommon "github.com/QuantumNous/new-api/relay/common"
  17. "github.com/QuantumNous/new-api/service"
  18. "github.com/gin-gonic/gin"
  19. "github.com/pkg/errors"
  20. "github.com/samber/lo"
  21. )
  22. // ============================
  23. // Request / Response structures
  24. // ============================
  25. type ContentItem struct {
  26. Type string `json:"type,omitempty"`
  27. Text string `json:"text,omitempty"`
  28. ImageURL *MediaURL `json:"image_url,omitempty"`
  29. VideoURL *MediaURL `json:"video_url,omitempty"`
  30. AudioURL *MediaURL `json:"audio_url,omitempty"`
  31. Role string `json:"role,omitempty"`
  32. }
  33. type MediaURL struct {
  34. URL string `json:"url,omitempty"`
  35. }
  36. type requestPayload struct {
  37. Model string `json:"model"`
  38. Content []ContentItem `json:"content,omitempty"`
  39. CallbackURL string `json:"callback_url,omitempty"`
  40. ReturnLastFrame *dto.BoolValue `json:"return_last_frame,omitempty"`
  41. ServiceTier string `json:"service_tier,omitempty"`
  42. ExecutionExpiresAfter *dto.IntValue `json:"execution_expires_after,omitempty"`
  43. GenerateAudio *dto.BoolValue `json:"generate_audio,omitempty"`
  44. Draft *dto.BoolValue `json:"draft,omitempty"`
  45. Tools []struct {
  46. Type string `json:"type,omitempty"`
  47. } `json:"tools,omitempty"`
  48. Resolution string `json:"resolution,omitempty"`
  49. Ratio string `json:"ratio,omitempty"`
  50. Duration *dto.IntValue `json:"duration,omitempty"`
  51. Frames *dto.IntValue `json:"frames,omitempty"`
  52. Seed *dto.IntValue `json:"seed,omitempty"`
  53. CameraFixed *dto.BoolValue `json:"camera_fixed,omitempty"`
  54. Watermark *dto.BoolValue `json:"watermark,omitempty"`
  55. }
  56. type responsePayload struct {
  57. ID string `json:"id"` // task_id
  58. }
  59. type responseTask struct {
  60. ID string `json:"id"`
  61. Model string `json:"model"`
  62. Status string `json:"status"`
  63. Content struct {
  64. VideoURL string `json:"videoUrl"`
  65. } `json:"content"`
  66. Seed int `json:"seed"`
  67. Resolution string `json:"resolution"`
  68. Duration int `json:"duration"`
  69. Ratio string `json:"ratio"`
  70. FramesPerSecond int `json:"framesPerSecond"`
  71. ServiceTier string `json:"serviceTier"`
  72. Usage struct {
  73. CompletionTokens int `json:"completionTokens"`
  74. TotalTokens int `json:"totalTokens"`
  75. } `json:"usage"`
  76. Error struct {
  77. Code string `json:"code"`
  78. Message string `json:"message"`
  79. } `json:"error"`
  80. CreatedAt int64 `json:"createdAt"`
  81. UpdatedAt int64 `json:"updatedAt"`
  82. }
  83. // ============================
  84. // Adaptor implementation
  85. // ============================
  86. type TaskAdaptor struct {
  87. taskcommon.BaseBilling
  88. ChannelType int
  89. apiKey string
  90. baseURL string
  91. }
  92. func (a *TaskAdaptor) Init(info *relaycommon.RelayInfo) {
  93. a.ChannelType = info.ChannelType
  94. a.baseURL = info.ChannelBaseUrl
  95. a.apiKey = info.ApiKey
  96. }
  97. // ValidateRequestAndSetAction parses body, validates fields and sets default action.
  98. func (a *TaskAdaptor) ValidateRequestAndSetAction(c *gin.Context, info *relaycommon.RelayInfo) (taskErr *dto.TaskError) {
  99. // The native /api/v3/contents/generations/tasks path pre-parses the
  100. // Volcengine-native body and stores it in the context; reuse it instead of
  101. // re-parsing the body as TaskSubmitReq (whose prompt validation would fail,
  102. // since the native prompt lives inside content[].text).
  103. if _, err := relaycommon.GetTaskRequest(c); err == nil {
  104. info.Action = constant.TaskActionGenerate
  105. return nil
  106. }
  107. return relaycommon.ValidateBasicTaskRequest(c, info, constant.TaskActionGenerate)
  108. }
  109. // BuildRequestURL constructs the upstream URL.
  110. func (a *TaskAdaptor) BuildRequestURL(info *relaycommon.RelayInfo) (string, error) {
  111. return fmt.Sprintf("%s/api/v3/contents/generations/tasks", a.baseURL), nil
  112. }
  113. // BuildRequestHeader sets required headers.
  114. func (a *TaskAdaptor) BuildRequestHeader(c *gin.Context, req *http.Request, info *relaycommon.RelayInfo) error {
  115. req.Header.Set("Content-Type", "application/json")
  116. req.Header.Set("Accept", "application/json")
  117. req.Header.Set("Authorization", "Bearer "+a.apiKey)
  118. return nil
  119. }
  120. // BuildRequestBody converts request into Doubao specific format.
  121. func (a *TaskAdaptor) BuildRequestBody(c *gin.Context, info *relaycommon.RelayInfo) (io.Reader, error) {
  122. req, err := relaycommon.GetTaskRequest(c)
  123. if err != nil {
  124. return nil, err
  125. }
  126. body, err := a.convertToRequestPayload(&req)
  127. if err != nil {
  128. return nil, errors.Wrap(err, "convert request payload failed")
  129. }
  130. if info.IsModelMapped {
  131. body.Model = info.UpstreamModelName
  132. } else {
  133. info.UpstreamModelName = body.Model
  134. }
  135. data, err := common.Marshal(body)
  136. if err != nil {
  137. return nil, err
  138. }
  139. return bytes.NewReader(data), nil
  140. }
  141. // DoRequest delegates to common helper.
  142. func (a *TaskAdaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error) {
  143. return channel.DoTaskApiRequest(a, c, info, requestBody)
  144. }
  145. // DoResponse handles upstream response, returns taskID etc.
  146. func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (taskID string, taskData []byte, taskErr *dto.TaskError) {
  147. responseBody, err := io.ReadAll(resp.Body)
  148. if err != nil {
  149. taskErr = service.TaskErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError)
  150. return
  151. }
  152. _ = resp.Body.Close()
  153. // Parse Doubao response
  154. var dResp responsePayload
  155. if err := common.Unmarshal(responseBody, &dResp); err != nil {
  156. taskErr = service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError)
  157. return
  158. }
  159. if dResp.ID == "" {
  160. taskErr = service.TaskErrorWrapper(fmt.Errorf("task_id is empty"), "invalid_response", http.StatusInternalServerError)
  161. return
  162. }
  163. ov := dto.NewOpenAIVideo()
  164. ov.ID = info.PublicTaskID
  165. ov.TaskID = info.PublicTaskID
  166. ov.CreatedAt = time.Now().Unix()
  167. ov.Model = info.OriginModelName
  168. c.JSON(http.StatusOK, ov)
  169. return dResp.ID, responseBody, nil
  170. }
  171. // FetchTask fetch task status
  172. func (a *TaskAdaptor) FetchTask(baseUrl, key string, body map[string]any, proxy string) (*http.Response, error) {
  173. taskID, ok := body["task_id"].(string)
  174. if !ok {
  175. return nil, fmt.Errorf("invalid task_id")
  176. }
  177. uri := fmt.Sprintf("%s/api/v3/contents/generations/tasks/%s", baseUrl, taskID)
  178. req, err := http.NewRequest(http.MethodGet, uri, nil)
  179. if err != nil {
  180. return nil, err
  181. }
  182. req.Header.Set("Accept", "application/json")
  183. req.Header.Set("Content-Type", "application/json")
  184. req.Header.Set("Authorization", "Bearer "+key)
  185. client, err := service.GetHttpClientWithProxy(proxy)
  186. if err != nil {
  187. return nil, fmt.Errorf("new proxy http client failed: %w", err)
  188. }
  189. return client.Do(req)
  190. }
  191. func (a *TaskAdaptor) GetModelList() []string {
  192. return ModelList
  193. }
  194. func (a *TaskAdaptor) GetChannelName() string {
  195. return ChannelName
  196. }
  197. func (a *TaskAdaptor) convertToRequestPayload(req *relaycommon.TaskSubmitReq) (*requestPayload, error) {
  198. r := requestPayload{
  199. Model: req.Model,
  200. Content: []ContentItem{},
  201. }
  202. // Add images if present
  203. if req.HasImage() {
  204. for _, imgURL := range req.Images {
  205. r.Content = append(r.Content, ContentItem{
  206. Type: "image_url",
  207. ImageURL: &MediaURL{
  208. URL: imgURL,
  209. },
  210. })
  211. }
  212. }
  213. metadata := req.Metadata
  214. if err := taskcommon.UnmarshalMetadata(metadata, &r); err != nil {
  215. return nil, errors.Wrap(err, "unmarshal metadata failed")
  216. }
  217. if sec, _ := strconv.Atoi(req.Seconds); sec > 0 {
  218. r.Duration = lo.ToPtr(dto.IntValue(sec))
  219. }
  220. // An explicit prompt replaces any text item from metadata. An empty prompt
  221. // only happens on the native /api/v3 path, where the prompt already lives
  222. // in a content text item that must be preserved as-is.
  223. if strings.TrimSpace(req.Prompt) != "" {
  224. r.Content = lo.Reject(r.Content, func(c ContentItem, _ int) bool { return c.Type == "text" })
  225. r.Content = append(r.Content, ContentItem{
  226. Type: "text",
  227. Text: req.Prompt,
  228. })
  229. }
  230. return &r, nil
  231. }
  232. func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, error) {
  233. resTask := responseTask{}
  234. if err := common.Unmarshal(respBody, &resTask); err != nil {
  235. return nil, errors.Wrap(err, "unmarshal task result failed")
  236. }
  237. taskResult := relaycommon.TaskInfo{
  238. Code: 0,
  239. }
  240. // Map Doubao status to internal status
  241. switch resTask.Status {
  242. case "pending", "queued":
  243. taskResult.Status = model.TaskStatusQueued
  244. taskResult.Progress = "10%"
  245. case "processing", "running":
  246. taskResult.Status = model.TaskStatusInProgress
  247. taskResult.Progress = "50%"
  248. case "succeeded":
  249. taskResult.Status = model.TaskStatusSuccess
  250. taskResult.Progress = "100%"
  251. taskResult.Url = resTask.Content.VideoURL
  252. // 解析 usage 信息用于按倍率计费
  253. taskResult.CompletionTokens = resTask.Usage.CompletionTokens
  254. taskResult.TotalTokens = resTask.Usage.TotalTokens
  255. case "failed":
  256. taskResult.Status = model.TaskStatusFailure
  257. taskResult.Progress = "100%"
  258. taskResult.Reason = resTask.Error.Message
  259. if taskResult.Reason == "" {
  260. taskResult.Reason = "task failed"
  261. }
  262. default:
  263. // Unknown status, treat as processing
  264. taskResult.Status = model.TaskStatusInProgress
  265. taskResult.Progress = "30%"
  266. }
  267. return &taskResult, nil
  268. }
  269. func (a *TaskAdaptor) ConvertToOpenAIVideo(originTask *model.Task) ([]byte, error) {
  270. var dResp responseTask
  271. if err := common.Unmarshal(originTask.Data, &dResp); err != nil {
  272. return nil, errors.Wrap(err, "unmarshal doubao task data failed")
  273. }
  274. openAIVideo := dto.NewOpenAIVideo()
  275. openAIVideo.ID = originTask.TaskID
  276. openAIVideo.TaskID = originTask.TaskID
  277. openAIVideo.Status = originTask.Status.ToVideoStatus()
  278. openAIVideo.SetProgressStr(originTask.Progress)
  279. openAIVideo.SetMetadata("url", dResp.Content.VideoURL)
  280. openAIVideo.CreatedAt = originTask.CreatedAt
  281. openAIVideo.CompletedAt = originTask.UpdatedAt
  282. openAIVideo.Model = originTask.Properties.OriginModelName
  283. if dResp.Status == "failed" {
  284. message := dResp.Error.Message
  285. if message == "" {
  286. message = "task failed"
  287. }
  288. code := dResp.Error.Code
  289. if code == "" {
  290. code = "failed"
  291. }
  292. openAIVideo.Error = &dto.OpenAIVideoError{
  293. Message: message,
  294. Code: code,
  295. }
  296. }
  297. return common.Marshal(openAIVideo)
  298. }