您最多选择25个主题 主题必须以字母或数字开头,可以包含连字符 (-),并且长度不得超过35个字符
 
 
 

552 行
16 KiB

  1. package relay
  2. import (
  3. "bytes"
  4. "errors"
  5. "fmt"
  6. "io"
  7. "net/http"
  8. "strconv"
  9. "strings"
  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. "github.com/QuantumNous/new-api/relay/channel/task/taskcommon"
  16. relaycommon "github.com/QuantumNous/new-api/relay/common"
  17. relayconstant "github.com/QuantumNous/new-api/relay/constant"
  18. "github.com/QuantumNous/new-api/relay/helper"
  19. "github.com/QuantumNous/new-api/service"
  20. "github.com/QuantumNous/new-api/types"
  21. "github.com/gin-gonic/gin"
  22. )
  23. type TaskSubmitResult struct {
  24. UpstreamTaskID string
  25. TaskData []byte
  26. Platform constant.TaskPlatform
  27. Quota int
  28. UpstreamReqJSON []byte
  29. //PerCallPrice types.PriceData
  30. }
  31. func ResolveOriginTask(c *gin.Context, info *relaycommon.RelayInfo) *dto.TaskError {
  32. path := c.Request.URL.Path
  33. if strings.Contains(path, "/v1/videos/") && strings.HasSuffix(path, "/remix") {
  34. info.Action = constant.TaskActionRemix
  35. }
  36. if info.Action == constant.TaskActionRemix {
  37. videoID := c.Param("video_id")
  38. if strings.TrimSpace(videoID) == "" {
  39. return service.TaskErrorWrapperLocal(fmt.Errorf("video_id is required"), "invalid_request", http.StatusBadRequest)
  40. }
  41. info.OriginTaskID = videoID
  42. }
  43. if info.OriginTaskID == "" {
  44. return nil
  45. }
  46. originTask, exist, err := model.GetByTaskId(info.UserId, info.OriginTaskID)
  47. if err != nil {
  48. return service.TaskErrorWrapper(err, "get_origin_task_failed", http.StatusInternalServerError)
  49. }
  50. if !exist {
  51. return service.TaskErrorWrapperLocal(errors.New("task_origin_not_exist"), "task_not_exist", http.StatusBadRequest)
  52. }
  53. // 娴犲骸甯慨瀣╂崲閸斺剝甯圭€靛吋膩閸ㄥ鎮曢敓?
  54. if info.OriginModelName == "" {
  55. if originTask.Properties.OriginModelName != "" {
  56. info.OriginModelName = originTask.Properties.OriginModelName
  57. } else if originTask.Properties.UpstreamModelName != "" {
  58. info.OriginModelName = originTask.Properties.UpstreamModelName
  59. } else {
  60. var taskData map[string]interface{}
  61. _ = common.Unmarshal(originTask.Data, &taskData)
  62. if m, ok := taskData["model"].(string); ok && m != "" {
  63. info.OriginModelName = m
  64. }
  65. }
  66. }
  67. ch, err := model.GetChannelById(originTask.ChannelId, true)
  68. if err != nil {
  69. return service.TaskErrorWrapperLocal(err, "channel_not_found", http.StatusBadRequest)
  70. }
  71. if ch.Status != common.ChannelStatusEnabled {
  72. return service.TaskErrorWrapperLocal(errors.New("the channel of the origin task is disabled"), "task_channel_disable", http.StatusBadRequest)
  73. }
  74. info.LockedChannel = ch
  75. if info.Action == constant.TaskActionRemix {
  76. if originTask.PrivateData.BillingContext != nil {
  77. bc := originTask.PrivateData.BillingContext
  78. info.OriginPricing = &types.OriginPricingSnapshot{
  79. BillingMode: bc.BillingMode,
  80. BillingUnit: bc.BillingUnit,
  81. PriceUSD: bc.ModelPrice,
  82. TokenUnitPriceUSD: bc.TokenUnitPriceUSD,
  83. Snapshot: types.CloneMapAny(bc.PricingSnapshot),
  84. OtherRatios: types.CloneRatios(bc.OtherRatios),
  85. PerCallBilling: bc.PerCallBilling,
  86. }
  87. } else {
  88. var taskData map[string]interface{}
  89. _ = common.Unmarshal(originTask.Data, &taskData)
  90. secondsStr, _ := taskData["seconds"].(string)
  91. seconds, _ := strconv.Atoi(secondsStr)
  92. if seconds <= 0 {
  93. seconds = 4
  94. }
  95. sizeStr, _ := taskData["size"].(string)
  96. otherRatios := map[string]float64{
  97. "seconds": float64(seconds),
  98. "size": 1,
  99. }
  100. if sizeStr == "1792x1024" || sizeStr == "1024x1792" {
  101. otherRatios["size"] = 1.666667
  102. }
  103. info.OriginPricing = &types.OriginPricingSnapshot{
  104. OtherRatios: otherRatios,
  105. PerCallBilling: false,
  106. }
  107. }
  108. }
  109. return nil
  110. }
  111. func RelayTaskSubmit(c *gin.Context, info *relaycommon.RelayInfo) (*TaskSubmitResult, *dto.TaskError) {
  112. info.InitChannelMeta(c)
  113. platform := constant.TaskPlatform(c.GetString("platform"))
  114. if platform == "" {
  115. platform = GetTaskPlatform(c)
  116. }
  117. adaptor := GetTaskAdaptor(platform)
  118. if adaptor == nil {
  119. return nil, service.TaskErrorWrapperLocal(fmt.Errorf("invalid api platform: %s", platform), "invalid_api_platform", http.StatusBadRequest)
  120. }
  121. adaptor.Init(info)
  122. if taskErr := adaptor.ValidateRequestAndSetAction(c, info); taskErr != nil {
  123. return nil, taskErr
  124. }
  125. modelName := info.OriginModelName
  126. if modelName == "" {
  127. modelName = service.CoverTaskActionToModelName(platform, info.Action)
  128. }
  129. info.UpstreamModelName = modelName
  130. if err := helper.ModelMappedHelper(c, info, nil); err != nil {
  131. return nil, service.TaskErrorWrapperLocal(err, "model_mapping_failed", http.StatusBadRequest)
  132. }
  133. if info.PublicTaskID == "" {
  134. info.PublicTaskID = model.GenerateTaskID()
  135. }
  136. // 4. 娴犻攱鐗哥拋锛勭暬閿涙艾鐔€绾偓濡€崇€锋禒閿嬬壐
  137. info.OriginModelName = modelName
  138. priceData, pricingErr := helper.ModelPriceHelperPerCall(c, info)
  139. if pricingErr != nil {
  140. return nil, pricingErr
  141. }
  142. info.PriceData = priceData
  143. matrixOrOriginPricing := info.OriginPricing != nil ||
  144. (info.PricingDecisionFrozen != nil && info.PricingDecisionFrozen.BillingMode == types.BillingModeMatrix)
  145. if !matrixOrOriginPricing {
  146. if estimatedRatios := adaptor.EstimateBilling(c, info); len(estimatedRatios) > 0 {
  147. for k, v := range estimatedRatios {
  148. info.PriceData.AddOtherRatio(k, v)
  149. }
  150. }
  151. }
  152. if !matrixOrOriginPricing && !common.StringsContains(constant.TaskPricePatches, modelName) {
  153. for _, ra := range info.PriceData.OtherRatios {
  154. if ra != 1.0 {
  155. info.PriceData.Quota = int(float64(info.PriceData.Quota) * ra)
  156. }
  157. }
  158. }
  159. if info.Billing == nil && !info.PriceData.FreeModel {
  160. info.ForcePreConsume = true
  161. if apiErr := service.PreConsumeBilling(c, info.PriceData.Quota, info); apiErr != nil {
  162. return nil, service.TaskErrorFromAPIError(apiErr)
  163. }
  164. }
  165. requestBody, err := adaptor.BuildRequestBody(c, info)
  166. if err != nil {
  167. return nil, service.TaskErrorWrapper(err, "build_request_failed", http.StatusInternalServerError)
  168. }
  169. var upstreamReqBytes []byte
  170. if requestBody != nil {
  171. upstreamReqBytes, err = io.ReadAll(requestBody)
  172. if err != nil {
  173. return nil, service.TaskErrorWrapper(err, "read_request_body_failed", http.StatusInternalServerError)
  174. }
  175. requestBody = bytes.NewReader(upstreamReqBytes)
  176. }
  177. resp, err := adaptor.DoRequest(c, info, requestBody)
  178. if err != nil {
  179. return nil, service.TaskErrorWrapper(err, "do_request_failed", http.StatusInternalServerError)
  180. }
  181. if resp != nil && resp.StatusCode != http.StatusOK {
  182. responseBody, _ := io.ReadAll(resp.Body)
  183. return nil, service.TaskErrorWrapper(fmt.Errorf("%s", string(responseBody)), "fail_to_fetch_task", resp.StatusCode)
  184. }
  185. otherRatios := info.PriceData.OtherRatios
  186. if otherRatios == nil {
  187. otherRatios = map[string]float64{}
  188. }
  189. ratiosJSON, _ := common.Marshal(otherRatios)
  190. c.Header("X-New-Api-Other-Ratios", string(ratiosJSON))
  191. // 11. 鐟欙絾鐎介崫宥呯安
  192. upstreamTaskID, taskData, taskErr := adaptor.DoResponse(c, resp, info)
  193. if taskErr != nil {
  194. return nil, taskErr
  195. }
  196. finalQuota := info.PriceData.Quota
  197. if !matrixOrOriginPricing {
  198. if adjustedRatios := adaptor.AdjustBillingOnSubmit(info, taskData); len(adjustedRatios) > 0 {
  199. finalQuota = recalcQuotaFromRatios(info, adjustedRatios)
  200. info.PriceData.OtherRatios = adjustedRatios
  201. info.PriceData.Quota = finalQuota
  202. }
  203. }
  204. return &TaskSubmitResult{
  205. UpstreamTaskID: upstreamTaskID,
  206. TaskData: taskData,
  207. Platform: platform,
  208. Quota: finalQuota,
  209. UpstreamReqJSON: upstreamReqBytes,
  210. }, nil
  211. }
  212. func recalcQuotaFromRatios(info *relaycommon.RelayInfo, ratios map[string]float64) int {
  213. baseQuota := info.PriceData.Quota
  214. for _, ra := range info.PriceData.OtherRatios {
  215. if ra != 1.0 && ra > 0 {
  216. baseQuota = int(float64(baseQuota) / ra)
  217. }
  218. }
  219. // 鎼存梻鏁ら弬鎵畱 ratios
  220. result := float64(baseQuota)
  221. for _, ra := range ratios {
  222. if ra != 1.0 {
  223. result *= ra
  224. }
  225. }
  226. return int(result)
  227. }
  228. var fetchRespBuilders = map[int]func(c *gin.Context) (respBody []byte, taskResp *dto.TaskError){
  229. relayconstant.RelayModeSunoFetchByID: sunoFetchByIDRespBodyBuilder,
  230. relayconstant.RelayModeSunoFetch: sunoFetchRespBodyBuilder,
  231. relayconstant.RelayModeVideoFetchByID: videoFetchByIDRespBodyBuilder,
  232. }
  233. func RelayTaskFetch(c *gin.Context, relayMode int) (taskResp *dto.TaskError) {
  234. respBuilder, ok := fetchRespBuilders[relayMode]
  235. if !ok {
  236. taskResp = service.TaskErrorWrapperLocal(errors.New("invalid_relay_mode"), "invalid_relay_mode", http.StatusBadRequest)
  237. }
  238. respBody, taskErr := respBuilder(c)
  239. if taskErr != nil {
  240. return taskErr
  241. }
  242. if len(respBody) == 0 {
  243. respBody = []byte("{\"code\":\"success\",\"data\":null}")
  244. }
  245. c.Writer.Header().Set("Content-Type", "application/json")
  246. _, err := io.Copy(c.Writer, bytes.NewBuffer(respBody))
  247. if err != nil {
  248. taskResp = service.TaskErrorWrapper(err, "copy_response_body_failed", http.StatusInternalServerError)
  249. return
  250. }
  251. return
  252. }
  253. func sunoFetchRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *dto.TaskError) {
  254. userId := c.GetInt("id")
  255. var condition = struct {
  256. IDs []any `json:"ids"`
  257. Action string `json:"action"`
  258. }{}
  259. err := c.BindJSON(&condition)
  260. if err != nil {
  261. taskResp = service.TaskErrorWrapper(err, "invalid_request", http.StatusBadRequest)
  262. return
  263. }
  264. var tasks []any
  265. if len(condition.IDs) > 0 {
  266. taskModels, err := model.GetByTaskIds(userId, condition.IDs)
  267. if err != nil {
  268. taskResp = service.TaskErrorWrapper(err, "get_tasks_failed", http.StatusInternalServerError)
  269. return
  270. }
  271. for _, task := range taskModels {
  272. tasks = append(tasks, TaskModel2Dto(task))
  273. }
  274. } else {
  275. tasks = make([]any, 0)
  276. }
  277. respBody, err = common.Marshal(dto.TaskResponse[[]any]{
  278. Code: "success",
  279. Data: tasks,
  280. })
  281. return
  282. }
  283. func sunoFetchByIDRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *dto.TaskError) {
  284. taskId := c.Param("id")
  285. userId := c.GetInt("id")
  286. originTask, exist, err := model.GetByTaskId(userId, taskId)
  287. if err != nil {
  288. taskResp = service.TaskErrorWrapper(err, "get_task_failed", http.StatusInternalServerError)
  289. return
  290. }
  291. if !exist {
  292. taskResp = service.TaskErrorWrapperLocal(errors.New("task_not_exist"), "task_not_exist", http.StatusBadRequest)
  293. return
  294. }
  295. respBody, err = common.Marshal(dto.TaskResponse[any]{
  296. Code: "success",
  297. Data: TaskModel2Dto(originTask),
  298. })
  299. return
  300. }
  301. func videoFetchByIDRespBodyBuilder(c *gin.Context) (respBody []byte, taskResp *dto.TaskError) {
  302. taskId := c.Param("task_id")
  303. if taskId == "" {
  304. taskId = c.GetString("task_id")
  305. }
  306. userId := c.GetInt("id")
  307. originTask, exist, err := model.GetByTaskId(userId, taskId)
  308. if err != nil {
  309. taskResp = service.TaskErrorWrapper(err, "get_task_failed", http.StatusInternalServerError)
  310. return
  311. }
  312. if !exist {
  313. taskResp = service.TaskErrorWrapperLocal(errors.New("task_not_exist"), "task_not_exist", http.StatusBadRequest)
  314. return
  315. }
  316. isOpenAIVideoAPI := strings.HasPrefix(c.Request.RequestURI, "/v1/videos/")
  317. // Gemini/Vertex 閺€顖涘瘮鐎圭偞妞傞弻銉嚄閿涙氨鏁ら敓?fetch 閺冨墎娲块幒銉ょ矤娑撳﹥鐖堕幏澶婂絿閺堚偓閺傛壆濮搁敓?
  318. if realtimeResp := tryRealtimeFetch(originTask, isOpenAIVideoAPI); len(realtimeResp) > 0 {
  319. respBody = realtimeResp
  320. return
  321. }
  322. if isOpenAIVideoAPI {
  323. adaptor := GetTaskAdaptor(originTask.Platform)
  324. if adaptor == nil {
  325. taskResp = service.TaskErrorWrapperLocal(fmt.Errorf("invalid channel id: %d", originTask.ChannelId), "invalid_channel_id", http.StatusBadRequest)
  326. return
  327. }
  328. if converter, ok := adaptor.(channel.OpenAIVideoConverter); ok {
  329. openAIVideoData, err := converter.ConvertToOpenAIVideo(originTask)
  330. if err != nil {
  331. taskResp = service.TaskErrorWrapper(err, "convert_to_openai_video_failed", http.StatusInternalServerError)
  332. return
  333. }
  334. respBody = openAIVideoData
  335. return
  336. }
  337. taskResp = service.TaskErrorWrapperLocal(fmt.Errorf("not_implemented:%s", originTask.Platform), "not_implemented", http.StatusNotImplemented)
  338. return
  339. }
  340. respBody, err = common.Marshal(dto.TaskResponse[any]{
  341. Code: "success",
  342. Data: TaskModel2Dto(originTask),
  343. })
  344. if err != nil {
  345. taskResp = service.TaskErrorWrapper(err, "marshal_response_failed", http.StatusInternalServerError)
  346. }
  347. return
  348. }
  349. func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte {
  350. channelModel, err := model.GetChannelById(task.ChannelId, true)
  351. if err != nil {
  352. return nil
  353. }
  354. if channelModel.Type != constant.ChannelTypeVertexAi && channelModel.Type != constant.ChannelTypeGemini {
  355. return nil
  356. }
  357. baseURL := constant.ChannelBaseURLs[channelModel.Type]
  358. if channelModel.GetBaseURL() != "" {
  359. baseURL = channelModel.GetBaseURL()
  360. }
  361. proxy := channelModel.GetSetting().Proxy
  362. adaptor := GetTaskAdaptor(constant.TaskPlatform(strconv.Itoa(channelModel.Type)))
  363. if adaptor == nil {
  364. return nil
  365. }
  366. fetchModel := task.Properties.UpstreamModelName
  367. if strings.TrimSpace(fetchModel) == "" {
  368. fetchModel = task.Properties.OriginModelName
  369. }
  370. resp, err := adaptor.FetchTask(baseURL, channelModel.Key, map[string]any{
  371. "task_id": task.GetUpstreamTaskID(),
  372. "action": task.Action,
  373. "model": fetchModel,
  374. }, proxy)
  375. if err != nil || resp == nil {
  376. return nil
  377. }
  378. defer resp.Body.Close()
  379. body, err := io.ReadAll(resp.Body)
  380. if err != nil {
  381. return nil
  382. }
  383. ti, err := adaptor.ParseTaskResult(body)
  384. if err != nil || ti == nil {
  385. return nil
  386. }
  387. snap := task.Snapshot()
  388. if ti.Status != "" {
  389. task.Status = model.TaskStatus(ti.Status)
  390. }
  391. if ti.Progress != "" {
  392. task.Progress = ti.Progress
  393. }
  394. if strings.HasPrefix(ti.Url, "data:") {
  395. } else if ti.Url != "" {
  396. task.PrivateData.ResultURL = ti.Url
  397. } else if task.Status == model.TaskStatusSuccess {
  398. task.PrivateData.ResultURL = taskcommon.BuildProxyURL(task.TaskID)
  399. }
  400. if !snap.Equal(task.Snapshot()) {
  401. _, _ = task.UpdateWithStatus(snap.Status)
  402. }
  403. if isOpenAIVideoAPI {
  404. return nil
  405. }
  406. format := detectVideoFormat(body)
  407. out := map[string]any{
  408. "error": nil,
  409. "format": format,
  410. "metadata": nil,
  411. "status": mapTaskStatusToSimple(task.Status),
  412. "task_id": task.TaskID,
  413. "url": task.GetResultURL(),
  414. }
  415. respBody, _ := common.Marshal(dto.TaskResponse[any]{
  416. Code: "success",
  417. Data: out,
  418. })
  419. return respBody
  420. }
  421. func detectVideoFormat(rawBody []byte) string {
  422. var raw map[string]any
  423. if err := common.Unmarshal(rawBody, &raw); err != nil {
  424. return "mp4"
  425. }
  426. respObj, ok := raw["response"].(map[string]any)
  427. if !ok {
  428. return "mp4"
  429. }
  430. vids, ok := respObj["videos"].([]any)
  431. if !ok || len(vids) == 0 {
  432. return "mp4"
  433. }
  434. v0, ok := vids[0].(map[string]any)
  435. if !ok {
  436. return "mp4"
  437. }
  438. mt, ok := v0["mimeType"].(string)
  439. if !ok || mt == "" || strings.Contains(mt, "mp4") {
  440. return "mp4"
  441. }
  442. return mt
  443. }
  444. func mapTaskStatusToSimple(status model.TaskStatus) string {
  445. switch status {
  446. case model.TaskStatusSuccess:
  447. return "succeeded"
  448. case model.TaskStatusFailure:
  449. return "failed"
  450. case model.TaskStatusQueued, model.TaskStatusSubmitted:
  451. return "queued"
  452. default:
  453. return "processing"
  454. }
  455. }
  456. func TaskModel2Dto(task *model.Task) *dto.TaskDto {
  457. return &dto.TaskDto{
  458. ID: task.ID,
  459. CreatedAt: task.CreatedAt,
  460. UpdatedAt: task.UpdatedAt,
  461. TaskID: task.TaskID,
  462. Platform: string(task.Platform),
  463. UserId: task.UserId,
  464. Group: task.Group,
  465. ChannelId: task.ChannelId,
  466. Quota: task.Quota,
  467. Action: task.Action,
  468. Status: string(task.Status),
  469. FailReason: task.FailReason,
  470. ResultURL: task.GetResultURL(),
  471. SubmitTime: task.SubmitTime,
  472. StartTime: task.StartTime,
  473. FinishTime: task.FinishTime,
  474. Progress: task.Progress,
  475. Properties: task.Properties,
  476. Username: task.Username,
  477. Data: sanitizeTaskDtoData(task),
  478. }
  479. }
  480. func sanitizeTaskDtoData(task *model.Task) []byte {
  481. if task == nil || len(task.Data) == 0 {
  482. return nil
  483. }
  484. payload := map[string]any{}
  485. if err := common.Unmarshal(task.Data, &payload); err != nil {
  486. return task.Data
  487. }
  488. if _, ok := payload["aiping_id"]; !ok {
  489. return task.Data
  490. }
  491. delete(payload, "aiping_id")
  492. payload["id"] = task.TaskID
  493. data, err := common.Marshal(payload)
  494. if err != nil {
  495. return task.Data
  496. }
  497. return data
  498. }