Nelze vybrat více než 25 témat Téma musí začínat písmenem nebo číslem, může obsahovat pomlčky („-“) a může být dlouhé až 35 znaků.
 
 
 

183 řádky
5.8 KiB

  1. package doubao_tianyiyun
  2. import (
  3. "io"
  4. "net/http"
  5. "net/http/httptest"
  6. "strings"
  7. "testing"
  8. "github.com/QuantumNous/new-api/common"
  9. "github.com/QuantumNous/new-api/constant"
  10. "github.com/QuantumNous/new-api/model"
  11. relaycommon "github.com/QuantumNous/new-api/relay/common"
  12. "github.com/gin-gonic/gin"
  13. "github.com/stretchr/testify/require"
  14. )
  15. func TestBuildRequestURLUsesTianyiYunEndpoint(t *testing.T) {
  16. adaptor := &TaskAdaptor{}
  17. adaptor.Init(&relaycommon.RelayInfo{
  18. ChannelMeta: &relaycommon.ChannelMeta{
  19. ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun,
  20. ChannelBaseUrl: "https://ai.ctaigw.cn",
  21. ApiKey: "sk-test",
  22. },
  23. })
  24. got, err := adaptor.BuildRequestURL(&relaycommon.RelayInfo{
  25. ChannelMeta: &relaycommon.ChannelMeta{ChannelBaseUrl: "https://ai.ctaigw.cn"},
  26. })
  27. require.NoError(t, err)
  28. require.Equal(t, "https://ai.ctaigw.cn/v1/contents/generations/tasks", got)
  29. }
  30. func TestFetchTaskUsesTianyiYunEndpoint(t *testing.T) {
  31. adaptor := &TaskAdaptor{}
  32. req, err := adaptor.buildFetchRequest("https://ai.ctaigw.cn", "sk-test", map[string]any{"task_id": "task-123"})
  33. require.NoError(t, err)
  34. require.Equal(t, http.MethodGet, req.Method)
  35. require.Equal(t, "https://ai.ctaigw.cn/v1/contents/generations/tasks/task-123", req.URL.String())
  36. require.Equal(t, "Bearer sk-test", req.Header.Get("Authorization"))
  37. require.Equal(t, "application/json", req.Header.Get("Accept"))
  38. }
  39. func TestGetModelListIncludesTianyiYunSeedanceModels(t *testing.T) {
  40. models := (&TaskAdaptor{}).GetModelList()
  41. require.Contains(t, models, "cdance2.0-0611")
  42. require.Contains(t, models, "cdance2.0-fast-0611")
  43. require.Equal(t, "DoubaoVideoCompatibleTianyiYun", (&TaskAdaptor{}).GetChannelName())
  44. }
  45. func TestBuildRequestBodyPreservesTianyiYunNativeFields(t *testing.T) {
  46. adaptor := &TaskAdaptor{}
  47. c := newTianyiYunTaskRequestContext(t, `{
  48. "model":"cdance2.0-0611",
  49. "prompt":"write a clean product video",
  50. "metadata":{
  51. "content":[
  52. {"type":"image_url","image_url":{"url":"https://example.test/input.png"},"role":"reference_image"}
  53. ],
  54. "ratio":"16:9",
  55. "duration":5,
  56. "watermark":false
  57. }
  58. }`)
  59. info := &relaycommon.RelayInfo{
  60. ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"},
  61. }
  62. body, err := adaptor.BuildRequestBody(c, info)
  63. require.NoError(t, err)
  64. data, err := io.ReadAll(body)
  65. require.NoError(t, err)
  66. require.Contains(t, string(data), `"model":"cdance2.0-0611"`)
  67. require.Contains(t, string(data), `"image_url":{"url":"https://example.test/input.png"}`)
  68. require.Contains(t, string(data), `"text":"write a clean product video"`)
  69. require.Contains(t, string(data), `"ratio":"16:9"`)
  70. require.Contains(t, string(data), `"duration":5`)
  71. require.Contains(t, string(data), `"watermark":false`)
  72. require.NotContains(t, string(data), `"seconds"`)
  73. }
  74. func TestDoResponseRewritesIDToPublicTaskID(t *testing.T) {
  75. adaptor := &TaskAdaptor{}
  76. w := httptest.NewRecorder()
  77. c, _ := gin.CreateTestContext(w)
  78. info := &relaycommon.RelayInfo{
  79. OriginModelName: "cdance2.0-0611",
  80. TaskRelayInfo: &relaycommon.TaskRelayInfo{PublicTaskID: "task_public"},
  81. }
  82. resp := &http.Response{
  83. StatusCode: http.StatusOK,
  84. Body: io.NopCloser(strings.NewReader(`{
  85. "id":"task-upstream",
  86. "status":"queued",
  87. "created_at":1781496040,
  88. "model":"cdance2.0-0611"
  89. }`)),
  90. }
  91. upstreamID, data, taskErr := adaptor.DoResponse(c, resp, info)
  92. require.Nil(t, taskErr)
  93. require.Equal(t, "task-upstream", upstreamID)
  94. require.Contains(t, string(data), `"id":"task-upstream"`)
  95. require.Equal(t, http.StatusOK, w.Code)
  96. require.Contains(t, w.Body.String(), `"id":"task_public"`)
  97. require.Contains(t, w.Body.String(), `"status":"queued"`)
  98. require.Contains(t, w.Body.String(), `"model":"cdance2.0-0611"`)
  99. }
  100. func TestParseTaskResultMapsTianyiYunStatuses(t *testing.T) {
  101. adaptor := &TaskAdaptor{}
  102. cases := []struct {
  103. name string
  104. body string
  105. wantStatus model.TaskStatus
  106. wantURL string
  107. wantReason string
  108. wantUsage bool
  109. }{
  110. {
  111. name: "success",
  112. body: `{"id":"task-ok","status":"success","content":{"video_url":"https://example.test/out.mp4"},"usage":{"completion_tokens":7,"total_tokens":11}}`,
  113. wantStatus: model.TaskStatusSuccess,
  114. wantURL: "https://example.test/out.mp4",
  115. wantUsage: true,
  116. },
  117. {
  118. name: "succeeded",
  119. body: `{"id":"task-ok","status":"succeeded","content":{"video_url":"https://example.test/out.mp4"}}`,
  120. wantStatus: model.TaskStatusSuccess,
  121. wantURL: "https://example.test/out.mp4",
  122. },
  123. {
  124. name: "running",
  125. body: `{"id":"task-run","status":"running"}`,
  126. wantStatus: model.TaskStatusInProgress,
  127. },
  128. {
  129. name: "failed",
  130. body: `{"id":"task-fail","status":"failed","error":{"code":"BadRequest","message":"bad prompt"}}`,
  131. wantStatus: model.TaskStatusFailure,
  132. wantReason: "bad prompt",
  133. },
  134. }
  135. for _, tc := range cases {
  136. t.Run(tc.name, func(t *testing.T) {
  137. got, err := adaptor.ParseTaskResult([]byte(tc.body))
  138. require.NoError(t, err)
  139. require.Equal(t, string(tc.wantStatus), got.Status)
  140. if tc.wantURL != "" {
  141. require.Equal(t, tc.wantURL, got.Url)
  142. }
  143. if tc.wantUsage {
  144. require.Equal(t, 7, got.CompletionTokens)
  145. require.Equal(t, 11, got.TotalTokens)
  146. }
  147. if tc.wantReason != "" {
  148. require.Equal(t, tc.wantReason, got.Reason)
  149. }
  150. })
  151. }
  152. }
  153. func newTianyiYunTaskRequestContext(t *testing.T, body string) *gin.Context {
  154. t.Helper()
  155. w := httptest.NewRecorder()
  156. c, _ := gin.CreateTestContext(w)
  157. c.Request = httptest.NewRequest(http.MethodPost, "/api/v3/contents/generations/tasks", strings.NewReader(body))
  158. c.Request.Header.Set("Content-Type", "application/json")
  159. var req relaycommon.TaskSubmitReq
  160. require.NoError(t, common.Unmarshal([]byte(body), &req))
  161. relaycommon.StoreTaskRequest(c, &relaycommon.RelayInfo{}, constant.TaskActionGenerate, req)
  162. return c
  163. }