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.
 
 
 

160 lines
5.6 KiB

  1. package openai
  2. import (
  3. "fmt"
  4. "io"
  5. "net/http"
  6. "strings"
  7. "github.com/QuantumNous/new-api/common"
  8. "github.com/QuantumNous/new-api/dto"
  9. "github.com/QuantumNous/new-api/logger"
  10. relaycommon "github.com/QuantumNous/new-api/relay/common"
  11. "github.com/QuantumNous/new-api/relay/helper"
  12. "github.com/QuantumNous/new-api/service"
  13. "github.com/QuantumNous/new-api/types"
  14. "github.com/gin-gonic/gin"
  15. )
  16. func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
  17. defer service.CloseResponseBodyGracefully(resp)
  18. // read response body
  19. var responsesResponse dto.OpenAIResponsesResponse
  20. responseBody, err := io.ReadAll(resp.Body)
  21. if err != nil {
  22. return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
  23. }
  24. err = common.Unmarshal(responseBody, &responsesResponse)
  25. if err != nil {
  26. apiErr := types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  27. apiErr.UpstreamBody = service.TruncateBody(string(responseBody))
  28. return nil, apiErr
  29. }
  30. if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
  31. apiErr := types.WithOpenAIError(*oaiError, resp.StatusCode)
  32. apiErr.UpstreamBody = service.TruncateBody(string(responseBody))
  33. return nil, apiErr
  34. }
  35. if responsesResponse.HasImageGenerationCall() {
  36. c.Set("image_generation_call", true)
  37. c.Set("image_generation_call_quality", responsesResponse.GetQuality())
  38. c.Set("image_generation_call_size", responsesResponse.GetSize())
  39. }
  40. // 写入新的 response body
  41. service.IOCopyBytesGracefully(c, resp, responseBody)
  42. // compute usage
  43. usage := dto.Usage{}
  44. if responsesResponse.Usage != nil {
  45. usage.PromptTokens = responsesResponse.Usage.InputTokens
  46. usage.CompletionTokens = responsesResponse.Usage.OutputTokens
  47. usage.TotalTokens = responsesResponse.Usage.TotalTokens
  48. if responsesResponse.Usage.InputTokensDetails != nil {
  49. usage.PromptTokensDetails.CachedTokens = responsesResponse.Usage.InputTokensDetails.CachedTokens
  50. }
  51. }
  52. if info == nil || info.ResponsesUsageInfo == nil || info.ResponsesUsageInfo.BuiltInTools == nil {
  53. return &usage, nil
  54. }
  55. // 解析 Tools 用量
  56. for _, tool := range responsesResponse.Tools {
  57. buildToolinfo, ok := info.ResponsesUsageInfo.BuiltInTools[common.Interface2String(tool["type"])]
  58. if !ok || buildToolinfo == nil {
  59. logger.LogError(c, fmt.Sprintf("BuiltInTools not found for tool type: %v", tool["type"]))
  60. continue
  61. }
  62. buildToolinfo.CallCount++
  63. }
  64. return &usage, nil
  65. }
  66. func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
  67. if resp == nil || resp.Body == nil {
  68. logger.LogError(c, "invalid response or response body")
  69. return nil, types.NewError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse)
  70. }
  71. defer service.CloseResponseBodyGracefully(resp)
  72. var usage = &dto.Usage{}
  73. var responseTextBuilder strings.Builder
  74. var streamErr *types.NewAPIError
  75. helper.StreamScannerHandler(c, resp, info, func(data string) bool {
  76. var streamResponse dto.ResponsesStreamResponse
  77. if err := common.UnmarshalJsonStr(data, &streamResponse); err == nil {
  78. sendResponsesStreamData(c, streamResponse, data)
  79. switch streamResponse.Type {
  80. case "response.completed":
  81. if streamResponse.Response != nil {
  82. if streamResponse.Response.Usage != nil {
  83. if streamResponse.Response.Usage.InputTokens != 0 {
  84. usage.PromptTokens = streamResponse.Response.Usage.InputTokens
  85. }
  86. if streamResponse.Response.Usage.OutputTokens != 0 {
  87. usage.CompletionTokens = streamResponse.Response.Usage.OutputTokens
  88. }
  89. if streamResponse.Response.Usage.TotalTokens != 0 {
  90. usage.TotalTokens = streamResponse.Response.Usage.TotalTokens
  91. }
  92. if streamResponse.Response.Usage.InputTokensDetails != nil {
  93. usage.PromptTokensDetails.CachedTokens = streamResponse.Response.Usage.InputTokensDetails.CachedTokens
  94. }
  95. }
  96. if streamResponse.Response.HasImageGenerationCall() {
  97. c.Set("image_generation_call", true)
  98. c.Set("image_generation_call_quality", streamResponse.Response.GetQuality())
  99. c.Set("image_generation_call_size", streamResponse.Response.GetSize())
  100. }
  101. }
  102. case "response.output_text.delta":
  103. responseTextBuilder.WriteString(streamResponse.Delta)
  104. case dto.ResponsesOutputTypeItemDone:
  105. if streamResponse.Item != nil {
  106. switch streamResponse.Item.Type {
  107. case dto.BuildInCallWebSearchCall:
  108. if info != nil && info.ResponsesUsageInfo != nil && info.ResponsesUsageInfo.BuiltInTools != nil {
  109. if webSearchTool, exists := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview]; exists && webSearchTool != nil {
  110. webSearchTool.CallCount++
  111. }
  112. }
  113. }
  114. }
  115. case "response.error", "response.failed", "error":
  116. streamErr = handleResponsesStreamError(streamResponse, data)
  117. return false
  118. }
  119. } else {
  120. logger.LogError(c, "failed to unmarshal stream response: "+err.Error())
  121. }
  122. return true
  123. })
  124. if streamErr != nil {
  125. return nil, streamErr
  126. }
  127. if usage.CompletionTokens == 0 {
  128. // 计算输出文本的 token 数量
  129. tempStr := responseTextBuilder.String()
  130. if len(tempStr) > 0 {
  131. // 非正常结束,使用输出文本的 token 数量
  132. completionTokens := service.CountTextToken(tempStr, info.UpstreamModelName)
  133. usage.CompletionTokens = completionTokens
  134. }
  135. }
  136. if usage.PromptTokens == 0 && usage.CompletionTokens != 0 {
  137. usage.PromptTokens = info.GetEstimatePromptTokens()
  138. }
  139. usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
  140. return usage, nil
  141. }