relay_responses.go 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132
  1. package openai
  2. import (
  3. "fmt"
  4. "io"
  5. "net/http"
  6. "one-api/common"
  7. "one-api/dto"
  8. "one-api/logger"
  9. relaycommon "one-api/relay/common"
  10. "one-api/relay/helper"
  11. "one-api/service"
  12. "one-api/types"
  13. "strings"
  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. return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  27. }
  28. if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
  29. return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
  30. }
  31. // 写入新的 response body
  32. service.IOCopyBytesGracefully(c, resp, responseBody)
  33. // compute usage
  34. usage := dto.Usage{}
  35. if responsesResponse.Usage != nil {
  36. usage.PromptTokens = responsesResponse.Usage.InputTokens
  37. usage.CompletionTokens = responsesResponse.Usage.OutputTokens
  38. usage.TotalTokens = responsesResponse.Usage.TotalTokens
  39. if responsesResponse.Usage.InputTokensDetails != nil {
  40. usage.PromptTokensDetails.CachedTokens = responsesResponse.Usage.InputTokensDetails.CachedTokens
  41. }
  42. }
  43. if info == nil || info.ResponsesUsageInfo == nil || info.ResponsesUsageInfo.BuiltInTools == nil {
  44. return &usage, nil
  45. }
  46. // 解析 Tools 用量
  47. for _, tool := range responsesResponse.Tools {
  48. buildToolinfo, ok := info.ResponsesUsageInfo.BuiltInTools[common.Interface2String(tool["type"])]
  49. if !ok || buildToolinfo == nil {
  50. logger.LogError(c, fmt.Sprintf("BuiltInTools not found for tool type: %v", tool["type"]))
  51. continue
  52. }
  53. buildToolinfo.CallCount++
  54. }
  55. return &usage, nil
  56. }
  57. func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
  58. if resp == nil || resp.Body == nil {
  59. logger.LogError(c, "invalid response or response body")
  60. return nil, types.NewError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse)
  61. }
  62. defer service.CloseResponseBodyGracefully(resp)
  63. var usage = &dto.Usage{}
  64. var responseTextBuilder strings.Builder
  65. helper.StreamScannerHandler(c, resp, info, func(data string) bool {
  66. // 检查当前数据是否包含 completed 状态和 usage 信息
  67. var streamResponse dto.ResponsesStreamResponse
  68. if err := common.UnmarshalJsonStr(data, &streamResponse); err == nil {
  69. sendResponsesStreamData(c, streamResponse, data)
  70. switch streamResponse.Type {
  71. case "response.completed":
  72. if streamResponse.Response != nil && streamResponse.Response.Usage != nil {
  73. if streamResponse.Response.Usage.InputTokens != 0 {
  74. usage.PromptTokens = streamResponse.Response.Usage.InputTokens
  75. }
  76. if streamResponse.Response.Usage.OutputTokens != 0 {
  77. usage.CompletionTokens = streamResponse.Response.Usage.OutputTokens
  78. }
  79. if streamResponse.Response.Usage.TotalTokens != 0 {
  80. usage.TotalTokens = streamResponse.Response.Usage.TotalTokens
  81. }
  82. if streamResponse.Response.Usage.InputTokensDetails != nil {
  83. usage.PromptTokensDetails.CachedTokens = streamResponse.Response.Usage.InputTokensDetails.CachedTokens
  84. }
  85. }
  86. case "response.output_text.delta":
  87. // 处理输出文本
  88. responseTextBuilder.WriteString(streamResponse.Delta)
  89. case dto.ResponsesOutputTypeItemDone:
  90. // 函数调用处理
  91. if streamResponse.Item != nil {
  92. switch streamResponse.Item.Type {
  93. case dto.BuildInCallWebSearchCall:
  94. info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview].CallCount++
  95. }
  96. }
  97. }
  98. } else {
  99. logger.LogError(c, "failed to unmarshal stream response: "+err.Error())
  100. }
  101. return true
  102. })
  103. if usage.CompletionTokens == 0 {
  104. // 计算输出文本的 token 数量
  105. tempStr := responseTextBuilder.String()
  106. if len(tempStr) > 0 {
  107. // 非正常结束,使用输出文本的 token 数量
  108. completionTokens := service.CountTextToken(tempStr, info.UpstreamModelName)
  109. usage.CompletionTokens = completionTokens
  110. }
  111. }
  112. if usage.PromptTokens == 0 && usage.CompletionTokens != 0 {
  113. usage.PromptTokens = info.PromptTokens
  114. }
  115. usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
  116. return usage, nil
  117. }