Вы не можете выбрать более 25 тем Темы должны начинаться с буквы или цифры, могут содержать дефисы(-) и должны содержать не более 35 символов.
 
 
 

167 строки
5.5 KiB

  1. package openai
  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/dto"
  11. relaycommon "github.com/QuantumNous/new-api/relay/common"
  12. "github.com/QuantumNous/new-api/types"
  13. "github.com/gin-gonic/gin"
  14. "github.com/stretchr/testify/assert"
  15. "github.com/stretchr/testify/require"
  16. )
  17. func TestResponsesStreamEventErrorParsing(t *testing.T) {
  18. t.Run("standalone error event", func(t *testing.T) {
  19. data := `{"type":"error","error":{"type":"too_many_requests","code":"too_many_requests","message":"Too Many Requests","param":null}}`
  20. var streamResp dto.ResponsesStreamResponse
  21. err := common.UnmarshalJsonStr(data, &streamResp)
  22. require.NoError(t, err)
  23. assert.Equal(t, "error", streamResp.Type)
  24. assert.NotNil(t, streamResp.Error)
  25. oaiErr := dto.GetOpenAIError(streamResp.Error)
  26. require.NotNil(t, oaiErr)
  27. assert.Equal(t, "too_many_requests", oaiErr.Type)
  28. assert.Equal(t, "Too Many Requests", oaiErr.Message)
  29. })
  30. t.Run("response.failed event", func(t *testing.T) {
  31. data := `{"type":"response.failed","response":{"id":"test-id","status":"failed","error":{"type":"server_error","message":"Internal server error"}}}`
  32. var streamResp dto.ResponsesStreamResponse
  33. err := common.UnmarshalJsonStr(data, &streamResp)
  34. require.NoError(t, err)
  35. assert.Equal(t, "response.failed", streamResp.Type)
  36. require.NotNil(t, streamResp.Response)
  37. oaiErr := streamResp.Response.GetOpenAIError()
  38. require.NotNil(t, oaiErr)
  39. assert.Equal(t, "server_error", oaiErr.Type)
  40. })
  41. t.Run("response.error event with error in response object", func(t *testing.T) {
  42. data := `{"type":"response.error","response":{"id":"test-id","status":"failed","error":{"type":"invalid_request_error","message":"Invalid model"}}}`
  43. var streamResp dto.ResponsesStreamResponse
  44. err := common.UnmarshalJsonStr(data, &streamResp)
  45. require.NoError(t, err)
  46. assert.Equal(t, "response.error", streamResp.Type)
  47. oaiErr := streamResp.Response.GetOpenAIError()
  48. require.NotNil(t, oaiErr)
  49. assert.Equal(t, "invalid_request_error", oaiErr.Type)
  50. })
  51. t.Run("normal event has no error", func(t *testing.T) {
  52. data := `{"type":"response.output_text.delta","delta":"hello"}`
  53. var streamResp dto.ResponsesStreamResponse
  54. err := common.UnmarshalJsonStr(data, &streamResp)
  55. require.NoError(t, err)
  56. assert.Equal(t, "response.output_text.delta", streamResp.Type)
  57. assert.Nil(t, streamResp.Error)
  58. })
  59. }
  60. func TestOaiResponsesHandlerAttachesUpstreamBodyOnInvalidJSON(t *testing.T) {
  61. gin.SetMode(gin.TestMode)
  62. w := httptest.NewRecorder()
  63. c, _ := gin.CreateTestContext(w)
  64. resp := &http.Response{
  65. StatusCode: http.StatusOK,
  66. Body: io.NopCloser(strings.NewReader(`{"incomplete":`)),
  67. }
  68. usage, err := OaiResponsesHandler(c, &relaycommon.RelayInfo{}, resp)
  69. require.Nil(t, usage)
  70. require.Error(t, err)
  71. assert.NotEmpty(t, err.UpstreamBody)
  72. assert.Contains(t, err.UpstreamBody, `{"incomplete":`)
  73. }
  74. func TestOaiResponsesStreamHandlerDoesNotForwardErrorEvent(t *testing.T) {
  75. gin.SetMode(gin.TestMode)
  76. oldTimeout := constant.StreamingTimeout
  77. constant.StreamingTimeout = 30
  78. t.Cleanup(func() {
  79. constant.StreamingTimeout = oldTimeout
  80. })
  81. w := httptest.NewRecorder()
  82. c, _ := gin.CreateTestContext(w)
  83. c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
  84. errorEvent := `{"type":"error","error":{"type":"tokens","code":"rate_limit_exceeded","message":"Request too large","param":null},"sequence_number":2}`
  85. resp := &http.Response{
  86. StatusCode: http.StatusOK,
  87. Body: io.NopCloser(strings.NewReader("data: " + errorEvent + "\n\n")),
  88. }
  89. usage, err := OaiResponsesStreamHandler(c, &relaycommon.RelayInfo{}, resp)
  90. require.Nil(t, usage)
  91. require.Error(t, err)
  92. assert.Contains(t, err.UpstreamBody, "rate_limit_exceeded")
  93. assert.NotContains(t, w.Body.String(), "rate_limit_exceeded")
  94. assert.Empty(t, w.Body.String())
  95. assert.Empty(t, w.Header().Get("Content-Type"))
  96. }
  97. func TestOaiStreamHandlerDoesNotForwardInitialErrorEvent(t *testing.T) {
  98. gin.SetMode(gin.TestMode)
  99. oldTimeout := constant.StreamingTimeout
  100. constant.StreamingTimeout = 30
  101. t.Cleanup(func() {
  102. constant.StreamingTimeout = oldTimeout
  103. })
  104. w := httptest.NewRecorder()
  105. c, _ := gin.CreateTestContext(w)
  106. c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
  107. errorEvent := `{"error":{"type":"tokens","code":"rate_limit_exceeded","message":"Request too large","param":null}}`
  108. resp := &http.Response{
  109. StatusCode: http.StatusOK,
  110. Body: io.NopCloser(strings.NewReader("data: " + errorEvent + "\n\n")),
  111. }
  112. info := &relaycommon.RelayInfo{
  113. RelayFormat: types.RelayFormatOpenAI,
  114. ChannelMeta: &relaycommon.ChannelMeta{
  115. UpstreamModelName: "gpt-test",
  116. },
  117. }
  118. usage, err := OaiStreamHandler(c, info, resp)
  119. require.Nil(t, usage)
  120. require.Error(t, err)
  121. assert.Contains(t, err.UpstreamBody, "rate_limit_exceeded")
  122. assert.NotContains(t, w.Body.String(), "rate_limit_exceeded")
  123. assert.Empty(t, w.Body.String())
  124. assert.Empty(t, w.Header().Get("Content-Type"))
  125. }
  126. func TestOaiResponsesToChatHandlerAttachesUpstreamBodyOnInvalidJSON(t *testing.T) {
  127. gin.SetMode(gin.TestMode)
  128. w := httptest.NewRecorder()
  129. c, _ := gin.CreateTestContext(w)
  130. resp := &http.Response{
  131. StatusCode: http.StatusOK,
  132. Body: io.NopCloser(strings.NewReader(`{"incomplete":`)),
  133. }
  134. usage, err := OaiResponsesToChatHandler(c, &relaycommon.RelayInfo{}, resp)
  135. require.Nil(t, usage)
  136. require.Error(t, err)
  137. assert.NotEmpty(t, err.UpstreamBody)
  138. assert.Contains(t, err.UpstreamBody, `{"incomplete":`)
  139. }