You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 

169 lines
5.5 KiB

  1. package aiping
  2. import (
  3. "io"
  4. "net/http"
  5. "net/http/httptest"
  6. "testing"
  7. "github.com/QuantumNous/new-api/common"
  8. relaycommon "github.com/QuantumNous/new-api/relay/common"
  9. "github.com/QuantumNous/new-api/service"
  10. "github.com/gin-gonic/gin"
  11. "github.com/stretchr/testify/require"
  12. )
  13. func TestRoutesUseOfficialExternalAndAipingUpstreamPaths(t *testing.T) {
  14. route, ok := FindRoute(http.MethodPost, "/v1/videos/text2video", RouteKindSubmit)
  15. require.True(t, ok)
  16. require.Equal(t, ActionText2Video, route.Action)
  17. require.Equal(t, "/v1/videos/kling/text2video", route.UpstreamPath)
  18. voice, ok := FindRoute(http.MethodPost, "/v1/kling/general/custom-voices", RouteKindSubmit)
  19. require.True(t, ok)
  20. require.Equal(t, "/v1/kling/general/custom-voices", voice.UpstreamPath)
  21. presets, ok := FindRoute(http.MethodGet, "/v1/kling/general/presets-voices", RouteKindProxy)
  22. require.True(t, ok)
  23. require.Equal(t, "/v1/kling/general/presets-voices", presets.UpstreamPath)
  24. deleteVoices, ok := FindRoute(http.MethodPost, "/v1/kling/general/delete-voices", RouteKindProxy)
  25. require.True(t, ok)
  26. require.Equal(t, ActionDeleteVoices, deleteVoices.Action)
  27. deleteElements, ok := FindRoute(http.MethodPost, "/v1/kling/general/delete-elements", RouteKindProxy)
  28. require.True(t, ok)
  29. require.Equal(t, ActionDeleteElements, deleteElements.Action)
  30. }
  31. func TestBuildRequestURLUsesAipingKlingPath(t *testing.T) {
  32. adaptor := &TaskAdaptor{}
  33. info := &relaycommon.RelayInfo{
  34. ChannelMeta: &relaycommon.ChannelMeta{ChannelBaseUrl: "https://aiping.cn/api"},
  35. TaskRelayInfo: &relaycommon.TaskRelayInfo{
  36. Action: ActionText2Video,
  37. },
  38. }
  39. adaptor.Init(info)
  40. url, err := adaptor.BuildRequestURL(info)
  41. require.NoError(t, err)
  42. require.Equal(t, "https://aiping.cn/api/v1/videos/kling/text2video", url)
  43. }
  44. func TestFetchTaskUsesSingleTaskPath(t *testing.T) {
  45. service.InitHttpClient()
  46. var gotPath string
  47. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  48. gotPath = r.URL.Path
  49. w.Header().Set("Content-Type", "application/json")
  50. _, _ = w.Write([]byte(`{"code":0,"data":{"task_id":"upstream-task","task_status":"submitted"}}`))
  51. }))
  52. defer server.Close()
  53. resp, err := (&TaskAdaptor{}).FetchTask(server.URL, "key", map[string]any{
  54. "task_id": "upstream-task",
  55. "action": ActionText2Video,
  56. }, "")
  57. require.NoError(t, err)
  58. defer resp.Body.Close()
  59. require.Equal(t, "/v1/videos/kling/text2video/upstream-task", gotPath)
  60. }
  61. func TestBuildRequestBodyNormalizesOfficialCompatibleFields(t *testing.T) {
  62. gin.SetMode(gin.TestMode)
  63. c, _ := gin.CreateTestContext(nil)
  64. info := &relaycommon.RelayInfo{
  65. ChannelMeta: &relaycommon.ChannelMeta{},
  66. TaskRelayInfo: &relaycommon.TaskRelayInfo{
  67. Action: ActionMultiImage2Video,
  68. },
  69. OriginModelName: "Kling-V2.6",
  70. }
  71. relaycommon.StoreTaskRequest(c, info, ActionMultiImage2Video, relaycommon.TaskSubmitReq{
  72. Model: "Kling-V2.6",
  73. Metadata: map[string]any{
  74. "model": "Kling-V1.6",
  75. "model_name": "Kling-V2.6",
  76. "seconds": 3,
  77. "duration": 5,
  78. "action_control": map[string]any{"type": "continuous"},
  79. "uid": "user-1",
  80. "create_at": float64(1750000000000),
  81. "_standard_model": "Kling-V-2-6",
  82. "image_list": []any{
  83. "https://example.com/1.png",
  84. map[string]any{"image_url": "https://example.com/2.png"},
  85. map[string]any{"url": "https://example.com/3.png"},
  86. map[string]any{"base64": "abc"},
  87. },
  88. },
  89. })
  90. body, err := (&TaskAdaptor{}).BuildRequestBody(c, info)
  91. require.NoError(t, err)
  92. data, err := io.ReadAll(body)
  93. require.NoError(t, err)
  94. var got map[string]any
  95. require.NoError(t, common.Unmarshal(data, &got))
  96. require.Equal(t, "Kling-V2.6", got["model_name"])
  97. require.NotContains(t, got, "model")
  98. require.Equal(t, "5", got["duration"])
  99. require.NotContains(t, got, "seconds")
  100. require.NotContains(t, got, "action_control")
  101. require.NotContains(t, got, "uid")
  102. require.NotContains(t, got, "create_at")
  103. require.NotContains(t, got, "_standard_model")
  104. require.Equal(t, []any{
  105. map[string]any{"image": "https://example.com/1.png"},
  106. map[string]any{"image": "https://example.com/2.png"},
  107. map[string]any{"image": "https://example.com/3.png"},
  108. map[string]any{"image": "abc"},
  109. }, got["image_list"])
  110. }
  111. func TestValidateRequestRejectsOmniModelOnTextOrImageRoutes(t *testing.T) {
  112. gin.SetMode(gin.TestMode)
  113. c, _ := gin.CreateTestContext(nil)
  114. info := &relaycommon.RelayInfo{
  115. TaskRelayInfo: &relaycommon.TaskRelayInfo{
  116. Action: ActionText2Video,
  117. },
  118. }
  119. relaycommon.StoreTaskRequest(c, info, ActionText2Video, relaycommon.TaskSubmitReq{
  120. Model: "kling-v3-omni",
  121. Metadata: map[string]any{
  122. "model_name": "kling-v3-omni",
  123. },
  124. })
  125. taskErr := (&TaskAdaptor{}).ValidateRequestAndSetAction(c, info)
  126. require.NotNil(t, taskErr)
  127. require.Equal(t, http.StatusUnprocessableEntity, taskErr.StatusCode)
  128. require.Contains(t, taskErr.Message, "/v1/videos/omni-video")
  129. }
  130. func TestSanitizeNativeTaskPayloadUsesPublicTaskIDAndAddsWatermarkURL(t *testing.T) {
  131. payload := map[string]any{
  132. "aiping_id": "internal",
  133. "data": map[string]any{
  134. "task_id": "899333358055493641",
  135. "task_result": map[string]any{
  136. "videos": []any{
  137. map[string]any{"id": "v1", "url": "https://example.com/video.mp4"},
  138. },
  139. },
  140. },
  141. }
  142. sanitizeNativeTaskPayload(payload, "task_public")
  143. require.NotContains(t, payload, "aiping_id")
  144. data := payload["data"].(map[string]any)
  145. require.Equal(t, "task_public", data["task_id"])
  146. videos := data["task_result"].(map[string]any)["videos"].([]any)
  147. require.Equal(t, "", videos[0].(map[string]any)["watermark_url"])
  148. }