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

710 строки
24 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/constant"
  9. "github.com/QuantumNous/new-api/dto"
  10. "github.com/QuantumNous/new-api/logger"
  11. "github.com/QuantumNous/new-api/relay/channel/openrouter"
  12. relaycommon "github.com/QuantumNous/new-api/relay/common"
  13. "github.com/QuantumNous/new-api/relay/helper"
  14. "github.com/QuantumNous/new-api/service"
  15. "github.com/QuantumNous/new-api/types"
  16. "github.com/bytedance/gopkg/util/gopool"
  17. "github.com/gin-gonic/gin"
  18. "github.com/gorilla/websocket"
  19. )
  20. func sendStreamData(c *gin.Context, info *relaycommon.RelayInfo, data string, forceFormat bool, thinkToContent bool) error {
  21. if data == "" {
  22. return nil
  23. }
  24. if !forceFormat && !thinkToContent {
  25. return helper.StringData(c, data)
  26. }
  27. var lastStreamResponse dto.ChatCompletionsStreamResponse
  28. if err := common.UnmarshalJsonStr(data, &lastStreamResponse); err != nil {
  29. return err
  30. }
  31. if !thinkToContent {
  32. return helper.ObjectData(c, lastStreamResponse)
  33. }
  34. hasThinkingContent := false
  35. hasContent := false
  36. var thinkingContent strings.Builder
  37. for _, choice := range lastStreamResponse.Choices {
  38. if len(choice.Delta.GetReasoningContent()) > 0 {
  39. hasThinkingContent = true
  40. thinkingContent.WriteString(choice.Delta.GetReasoningContent())
  41. }
  42. if len(choice.Delta.GetContentString()) > 0 {
  43. hasContent = true
  44. }
  45. }
  46. // Handle think to content conversion
  47. if info.ThinkingContentInfo.IsFirstThinkingContent {
  48. if hasThinkingContent {
  49. response := lastStreamResponse.Copy()
  50. for i := range response.Choices {
  51. // send `think` tag with thinking content
  52. response.Choices[i].Delta.SetContentString("<think>\n" + thinkingContent.String())
  53. response.Choices[i].Delta.ReasoningContent = nil
  54. response.Choices[i].Delta.Reasoning = nil
  55. }
  56. info.ThinkingContentInfo.IsFirstThinkingContent = false
  57. info.ThinkingContentInfo.HasSentThinkingContent = true
  58. return helper.ObjectData(c, response)
  59. }
  60. }
  61. if lastStreamResponse.Choices == nil || len(lastStreamResponse.Choices) == 0 {
  62. return helper.ObjectData(c, lastStreamResponse)
  63. }
  64. // Process each choice
  65. for i, choice := range lastStreamResponse.Choices {
  66. // Handle transition from thinking to content
  67. // only send `</think>` tag when previous thinking content has been sent
  68. if hasContent && !info.ThinkingContentInfo.SendLastThinkingContent && info.ThinkingContentInfo.HasSentThinkingContent {
  69. response := lastStreamResponse.Copy()
  70. for j := range response.Choices {
  71. response.Choices[j].Delta.SetContentString("\n</think>\n")
  72. response.Choices[j].Delta.ReasoningContent = nil
  73. response.Choices[j].Delta.Reasoning = nil
  74. }
  75. info.ThinkingContentInfo.SendLastThinkingContent = true
  76. helper.ObjectData(c, response)
  77. }
  78. // Convert reasoning content to regular content if any
  79. if len(choice.Delta.GetReasoningContent()) > 0 {
  80. lastStreamResponse.Choices[i].Delta.SetContentString(choice.Delta.GetReasoningContent())
  81. lastStreamResponse.Choices[i].Delta.ReasoningContent = nil
  82. lastStreamResponse.Choices[i].Delta.Reasoning = nil
  83. } else if !hasThinkingContent && !hasContent {
  84. // flush thinking content
  85. lastStreamResponse.Choices[i].Delta.ReasoningContent = nil
  86. lastStreamResponse.Choices[i].Delta.Reasoning = nil
  87. }
  88. }
  89. return helper.ObjectData(c, lastStreamResponse)
  90. }
  91. func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
  92. if resp == nil || resp.Body == nil {
  93. logger.LogError(c, "invalid response or response body")
  94. return nil, types.NewOpenAIError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse, http.StatusInternalServerError)
  95. }
  96. defer service.CloseResponseBodyGracefully(resp)
  97. model := info.UpstreamModelName
  98. var responseId string
  99. var createAt int64 = 0
  100. var systemFingerprint string
  101. var containStreamUsage bool
  102. var responseTextBuilder strings.Builder
  103. var toolCount int
  104. var usage = &dto.Usage{}
  105. var streamItems []string // store stream items
  106. var lastStreamData string
  107. var streamErr *types.NewAPIError
  108. var secondLastStreamData string // 存储倒数第二个stream data,用于音频模型
  109. // 检查是否为音频模型
  110. isAudioModel := strings.Contains(strings.ToLower(model), "audio")
  111. helper.StreamScannerHandler(c, resp, info, func(data string) bool {
  112. if apiErr := streamOpenAIErrorFromData(data, resp.StatusCode); apiErr != nil {
  113. if !c.Writer.Written() {
  114. streamErr = apiErr
  115. } else {
  116. logger.LogError(c, "upstream stream error after downstream response was written: "+apiErr.Error())
  117. }
  118. return false
  119. }
  120. if lastStreamData != "" {
  121. err := HandleStreamFormat(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent)
  122. if err != nil {
  123. common.SysLog("error handling stream format: " + err.Error())
  124. }
  125. }
  126. if len(data) > 0 {
  127. // 对音频模型,保存倒数第二个stream data
  128. if isAudioModel && lastStreamData != "" {
  129. secondLastStreamData = lastStreamData
  130. }
  131. lastStreamData = data
  132. streamItems = append(streamItems, data)
  133. }
  134. return true
  135. })
  136. // 对音频模型,从倒数第二个stream data中提取usage信息
  137. if streamErr != nil {
  138. helper.ClearEventStreamHeadersIfNotWritten(c)
  139. return nil, streamErr
  140. }
  141. if isAudioModel && secondLastStreamData != "" {
  142. var streamResp struct {
  143. Usage *dto.Usage `json:"usage"`
  144. }
  145. err := common.Unmarshal([]byte(secondLastStreamData), &streamResp)
  146. if err == nil && streamResp.Usage != nil && service.ValidUsage(streamResp.Usage) {
  147. usage = streamResp.Usage
  148. containStreamUsage = true
  149. if common.DebugEnabled {
  150. logger.LogDebug(c, fmt.Sprintf("Audio model usage extracted from second last SSE: PromptTokens=%d, CompletionTokens=%d, TotalTokens=%d, InputTokens=%d, OutputTokens=%d",
  151. usage.PromptTokens, usage.CompletionTokens, usage.TotalTokens,
  152. usage.InputTokens, usage.OutputTokens))
  153. }
  154. }
  155. }
  156. // 处理最后的响应
  157. shouldSendLastResp := true
  158. if err := handleLastResponse(lastStreamData, &responseId, &createAt, &systemFingerprint, &model, &usage,
  159. &containStreamUsage, info, &shouldSendLastResp); err != nil {
  160. logger.LogError(c, fmt.Sprintf("error handling last response: %s, lastStreamData: [%s]", err.Error(), lastStreamData))
  161. }
  162. if info.RelayFormat == types.RelayFormatOpenAI {
  163. if shouldSendLastResp {
  164. _ = sendStreamData(c, info, lastStreamData, info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent)
  165. }
  166. }
  167. // 处理token计算
  168. if err := processTokens(info.RelayMode, streamItems, &responseTextBuilder, &toolCount); err != nil {
  169. logger.LogError(c, "error processing tokens: "+err.Error())
  170. }
  171. if !containStreamUsage {
  172. usage = service.ResponseText2Usage(c, responseTextBuilder.String(), info.UpstreamModelName, info.GetEstimatePromptTokens())
  173. usage.CompletionTokens += toolCount * 7
  174. }
  175. applyUsagePostProcessing(info, usage, common.StringToByteSlice(lastStreamData))
  176. HandleFinalResponse(c, info, lastStreamData, responseId, createAt, model, systemFingerprint, usage, containStreamUsage)
  177. relaycommon.SetRelayChatID(c, responseId)
  178. return usage, nil
  179. }
  180. func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
  181. defer service.CloseResponseBodyGracefully(resp)
  182. var simpleResponse dto.OpenAITextResponse
  183. responseBody, err := io.ReadAll(resp.Body)
  184. if err != nil {
  185. return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
  186. }
  187. if common.DebugEnabled {
  188. println("upstream response body:", string(responseBody))
  189. }
  190. // Unmarshal to simpleResponse
  191. if info.ChannelType == constant.ChannelTypeOpenRouter && info.ChannelOtherSettings.IsOpenRouterEnterprise() {
  192. // 尝试解析为 openrouter enterprise
  193. var enterpriseResponse openrouter.OpenRouterEnterpriseResponse
  194. err = common.Unmarshal(responseBody, &enterpriseResponse)
  195. if err != nil {
  196. return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  197. }
  198. if enterpriseResponse.Success {
  199. responseBody = enterpriseResponse.Data
  200. } else {
  201. logger.LogError(c, fmt.Sprintf("openrouter enterprise response success=false, data: %s", enterpriseResponse.Data))
  202. return nil, types.NewOpenAIError(fmt.Errorf("openrouter response success=false"), types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  203. }
  204. }
  205. err = common.Unmarshal(responseBody, &simpleResponse)
  206. if err != nil {
  207. return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  208. }
  209. relaycommon.SetRelayChatID(c, simpleResponse.Id)
  210. if oaiError := simpleResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
  211. return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
  212. }
  213. for _, choice := range simpleResponse.Choices {
  214. if choice.FinishReason == constant.FinishReasonContentFilter {
  215. common.SetContextKey(c, constant.ContextKeyAdminRejectReason, "openai_finish_reason=content_filter")
  216. break
  217. }
  218. }
  219. forceFormat := false
  220. if info.ChannelSetting.ForceFormat {
  221. forceFormat = true
  222. }
  223. usageModified := false
  224. if simpleResponse.Usage.PromptTokens == 0 {
  225. completionTokens := simpleResponse.Usage.CompletionTokens
  226. if completionTokens == 0 {
  227. for _, choice := range simpleResponse.Choices {
  228. ctkm := service.CountTextToken(choice.Message.StringContent()+choice.Message.ReasoningContent+choice.Message.Reasoning, info.UpstreamModelName)
  229. completionTokens += ctkm
  230. }
  231. }
  232. simpleResponse.Usage = dto.Usage{
  233. PromptTokens: info.GetEstimatePromptTokens(),
  234. CompletionTokens: completionTokens,
  235. TotalTokens: info.GetEstimatePromptTokens() + completionTokens,
  236. }
  237. usageModified = true
  238. }
  239. applyUsagePostProcessing(info, &simpleResponse.Usage, responseBody)
  240. switch info.RelayFormat {
  241. case types.RelayFormatOpenAI:
  242. if usageModified {
  243. var bodyMap map[string]interface{}
  244. err = common.Unmarshal(responseBody, &bodyMap)
  245. if err != nil {
  246. return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  247. }
  248. bodyMap["usage"] = simpleResponse.Usage
  249. responseBody, _ = common.Marshal(bodyMap)
  250. }
  251. if forceFormat {
  252. responseBody, err = common.Marshal(simpleResponse)
  253. if err != nil {
  254. return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
  255. }
  256. } else {
  257. break
  258. }
  259. case types.RelayFormatClaude:
  260. claudeResp := service.ResponseOpenAI2Claude(&simpleResponse, info)
  261. claudeRespStr, err := common.Marshal(claudeResp)
  262. if err != nil {
  263. return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
  264. }
  265. responseBody = claudeRespStr
  266. case types.RelayFormatGemini:
  267. geminiResp := service.ResponseOpenAI2Gemini(&simpleResponse, info)
  268. geminiRespStr, err := common.Marshal(geminiResp)
  269. if err != nil {
  270. return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
  271. }
  272. responseBody = geminiRespStr
  273. }
  274. service.IOCopyBytesGracefully(c, resp, responseBody)
  275. return &simpleResponse.Usage, nil
  276. }
  277. func streamTTSResponse(c *gin.Context, resp *http.Response) {
  278. c.Writer.WriteHeaderNow()
  279. flusher, ok := c.Writer.(http.Flusher)
  280. if !ok {
  281. logger.LogWarn(c, "streaming not supported")
  282. _, err := io.Copy(c.Writer, resp.Body)
  283. if err != nil {
  284. logger.LogWarn(c, err.Error())
  285. }
  286. return
  287. }
  288. buffer := make([]byte, 4096)
  289. for {
  290. n, err := resp.Body.Read(buffer)
  291. //logger.LogInfo(c, fmt.Sprintf("streamTTSResponse read %d bytes", n))
  292. if n > 0 {
  293. if _, writeErr := c.Writer.Write(buffer[:n]); writeErr != nil {
  294. logger.LogError(c, writeErr.Error())
  295. break
  296. }
  297. flusher.Flush()
  298. }
  299. if err != nil {
  300. if err != io.EOF {
  301. logger.LogError(c, err.Error())
  302. }
  303. break
  304. }
  305. }
  306. }
  307. func OpenaiRealtimeHandler(c *gin.Context, info *relaycommon.RelayInfo) (*types.NewAPIError, *dto.RealtimeUsage) {
  308. if info == nil || info.ClientWs == nil || info.TargetWs == nil {
  309. return types.NewError(fmt.Errorf("invalid websocket connection"), types.ErrorCodeBadResponse), nil
  310. }
  311. info.IsStream = true
  312. clientConn := info.ClientWs
  313. targetConn := info.TargetWs
  314. clientClosed := make(chan struct{})
  315. targetClosed := make(chan struct{})
  316. sendChan := make(chan []byte, 100)
  317. receiveChan := make(chan []byte, 100)
  318. errChan := make(chan error, 2)
  319. usage := &dto.RealtimeUsage{}
  320. localUsage := &dto.RealtimeUsage{}
  321. sumUsage := &dto.RealtimeUsage{}
  322. gopool.Go(func() {
  323. defer func() {
  324. if r := recover(); r != nil {
  325. errChan <- fmt.Errorf("panic in client reader: %v", r)
  326. }
  327. }()
  328. for {
  329. select {
  330. case <-c.Done():
  331. return
  332. default:
  333. _, message, err := clientConn.ReadMessage()
  334. if err != nil {
  335. if !websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) {
  336. errChan <- fmt.Errorf("error reading from client: %v", err)
  337. }
  338. close(clientClosed)
  339. return
  340. }
  341. realtimeEvent := &dto.RealtimeEvent{}
  342. err = common.Unmarshal(message, realtimeEvent)
  343. if err != nil {
  344. errChan <- fmt.Errorf("error unmarshalling message: %v", err)
  345. return
  346. }
  347. if realtimeEvent.Type == dto.RealtimeEventTypeSessionUpdate {
  348. if realtimeEvent.Session != nil {
  349. if realtimeEvent.Session.Tools != nil {
  350. info.RealtimeTools = realtimeEvent.Session.Tools
  351. }
  352. }
  353. }
  354. textToken, audioToken, err := service.CountTokenRealtime(info, *realtimeEvent, info.UpstreamModelName)
  355. if err != nil {
  356. errChan <- fmt.Errorf("error counting text token: %v", err)
  357. return
  358. }
  359. logger.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken))
  360. localUsage.TotalTokens += textToken + audioToken
  361. localUsage.InputTokens += textToken + audioToken
  362. localUsage.InputTokenDetails.TextTokens += textToken
  363. localUsage.InputTokenDetails.AudioTokens += audioToken
  364. err = helper.WssString(c, targetConn, string(message))
  365. if err != nil {
  366. errChan <- fmt.Errorf("error writing to target: %v", err)
  367. return
  368. }
  369. select {
  370. case sendChan <- message:
  371. default:
  372. }
  373. }
  374. }
  375. })
  376. gopool.Go(func() {
  377. defer func() {
  378. if r := recover(); r != nil {
  379. errChan <- fmt.Errorf("panic in target reader: %v", r)
  380. }
  381. }()
  382. for {
  383. select {
  384. case <-c.Done():
  385. return
  386. default:
  387. _, message, err := targetConn.ReadMessage()
  388. if err != nil {
  389. if !websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) {
  390. errChan <- fmt.Errorf("error reading from target: %v", err)
  391. }
  392. close(targetClosed)
  393. return
  394. }
  395. info.SetFirstResponseTime()
  396. realtimeEvent := &dto.RealtimeEvent{}
  397. err = common.Unmarshal(message, realtimeEvent)
  398. if err != nil {
  399. errChan <- fmt.Errorf("error unmarshalling message: %v", err)
  400. return
  401. }
  402. if realtimeEvent.Type == dto.RealtimeEventTypeResponseDone {
  403. realtimeUsage := realtimeEvent.Response.Usage
  404. if realtimeUsage != nil {
  405. usage.TotalTokens += realtimeUsage.TotalTokens
  406. usage.InputTokens += realtimeUsage.InputTokens
  407. usage.OutputTokens += realtimeUsage.OutputTokens
  408. usage.InputTokenDetails.AudioTokens += realtimeUsage.InputTokenDetails.AudioTokens
  409. usage.InputTokenDetails.CachedTokens += realtimeUsage.InputTokenDetails.CachedTokens
  410. usage.InputTokenDetails.TextTokens += realtimeUsage.InputTokenDetails.TextTokens
  411. usage.OutputTokenDetails.AudioTokens += realtimeUsage.OutputTokenDetails.AudioTokens
  412. usage.OutputTokenDetails.TextTokens += realtimeUsage.OutputTokenDetails.TextTokens
  413. err := preConsumeUsage(c, info, usage, sumUsage)
  414. if err != nil {
  415. errChan <- fmt.Errorf("error consume usage: %v", err)
  416. return
  417. }
  418. // 本次计费完成,清除
  419. usage = &dto.RealtimeUsage{}
  420. localUsage = &dto.RealtimeUsage{}
  421. } else {
  422. textToken, audioToken, err := service.CountTokenRealtime(info, *realtimeEvent, info.UpstreamModelName)
  423. if err != nil {
  424. errChan <- fmt.Errorf("error counting text token: %v", err)
  425. return
  426. }
  427. logger.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken))
  428. localUsage.TotalTokens += textToken + audioToken
  429. info.IsFirstRequest = false
  430. localUsage.InputTokens += textToken + audioToken
  431. localUsage.InputTokenDetails.TextTokens += textToken
  432. localUsage.InputTokenDetails.AudioTokens += audioToken
  433. err = preConsumeUsage(c, info, localUsage, sumUsage)
  434. if err != nil {
  435. errChan <- fmt.Errorf("error consume usage: %v", err)
  436. return
  437. }
  438. // 本次计费完成,清除
  439. localUsage = &dto.RealtimeUsage{}
  440. // print now usage
  441. }
  442. logger.LogInfo(c, fmt.Sprintf("realtime streaming sumUsage: %v", sumUsage))
  443. logger.LogInfo(c, fmt.Sprintf("realtime streaming localUsage: %v", localUsage))
  444. logger.LogInfo(c, fmt.Sprintf("realtime streaming localUsage: %v", localUsage))
  445. } else if realtimeEvent.Type == dto.RealtimeEventTypeSessionUpdated || realtimeEvent.Type == dto.RealtimeEventTypeSessionCreated {
  446. realtimeSession := realtimeEvent.Session
  447. if realtimeSession != nil {
  448. // update audio format
  449. info.InputAudioFormat = common.GetStringIfEmpty(realtimeSession.InputAudioFormat, info.InputAudioFormat)
  450. info.OutputAudioFormat = common.GetStringIfEmpty(realtimeSession.OutputAudioFormat, info.OutputAudioFormat)
  451. }
  452. } else {
  453. textToken, audioToken, err := service.CountTokenRealtime(info, *realtimeEvent, info.UpstreamModelName)
  454. if err != nil {
  455. errChan <- fmt.Errorf("error counting text token: %v", err)
  456. return
  457. }
  458. logger.LogInfo(c, fmt.Sprintf("type: %s, textToken: %d, audioToken: %d", realtimeEvent.Type, textToken, audioToken))
  459. localUsage.TotalTokens += textToken + audioToken
  460. localUsage.OutputTokens += textToken + audioToken
  461. localUsage.OutputTokenDetails.TextTokens += textToken
  462. localUsage.OutputTokenDetails.AudioTokens += audioToken
  463. }
  464. err = helper.WssString(c, clientConn, string(message))
  465. if err != nil {
  466. errChan <- fmt.Errorf("error writing to client: %v", err)
  467. return
  468. }
  469. select {
  470. case receiveChan <- message:
  471. default:
  472. }
  473. }
  474. }
  475. })
  476. select {
  477. case <-clientClosed:
  478. case <-targetClosed:
  479. case err := <-errChan:
  480. //return service.OpenAIErrorWrapper(err, "realtime_error", http.StatusInternalServerError), nil
  481. logger.LogError(c, "realtime error: "+err.Error())
  482. case <-c.Done():
  483. }
  484. if usage.TotalTokens != 0 {
  485. _ = preConsumeUsage(c, info, usage, sumUsage)
  486. }
  487. if localUsage.TotalTokens != 0 {
  488. _ = preConsumeUsage(c, info, localUsage, sumUsage)
  489. }
  490. // check usage total tokens, if 0, use local usage
  491. return nil, sumUsage
  492. }
  493. func preConsumeUsage(ctx *gin.Context, info *relaycommon.RelayInfo, usage *dto.RealtimeUsage, totalUsage *dto.RealtimeUsage) error {
  494. if usage == nil || totalUsage == nil {
  495. return fmt.Errorf("invalid usage pointer")
  496. }
  497. totalUsage.TotalTokens += usage.TotalTokens
  498. totalUsage.InputTokens += usage.InputTokens
  499. totalUsage.OutputTokens += usage.OutputTokens
  500. totalUsage.InputTokenDetails.CachedTokens += usage.InputTokenDetails.CachedTokens
  501. totalUsage.InputTokenDetails.TextTokens += usage.InputTokenDetails.TextTokens
  502. totalUsage.InputTokenDetails.AudioTokens += usage.InputTokenDetails.AudioTokens
  503. totalUsage.OutputTokenDetails.TextTokens += usage.OutputTokenDetails.TextTokens
  504. totalUsage.OutputTokenDetails.AudioTokens += usage.OutputTokenDetails.AudioTokens
  505. // clear usage
  506. err := service.PreWssConsumeQuota(ctx, info, usage)
  507. return err
  508. }
  509. func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
  510. defer service.CloseResponseBodyGracefully(resp)
  511. responseBody, err := io.ReadAll(resp.Body)
  512. if err != nil {
  513. return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
  514. }
  515. var usageResp dto.SimpleResponse
  516. err = common.Unmarshal(responseBody, &usageResp)
  517. if err != nil {
  518. return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
  519. }
  520. // 写入新的 response body
  521. service.IOCopyBytesGracefully(c, resp, responseBody)
  522. // Once we've written to the client, we should not return errors anymore
  523. // because the upstream has already consumed resources and returned content
  524. // We should still perform billing even if parsing fails
  525. // format
  526. if usageResp.InputTokens > 0 {
  527. usageResp.PromptTokens += usageResp.InputTokens
  528. }
  529. if usageResp.OutputTokens > 0 {
  530. usageResp.CompletionTokens += usageResp.OutputTokens
  531. }
  532. if usageResp.InputTokensDetails != nil {
  533. usageResp.PromptTokensDetails.ImageTokens += usageResp.InputTokensDetails.ImageTokens
  534. usageResp.PromptTokensDetails.TextTokens += usageResp.InputTokensDetails.TextTokens
  535. }
  536. applyUsagePostProcessing(info, &usageResp.Usage, responseBody)
  537. return &usageResp.Usage, nil
  538. }
  539. func applyUsagePostProcessing(info *relaycommon.RelayInfo, usage *dto.Usage, responseBody []byte) {
  540. if info == nil || usage == nil {
  541. return
  542. }
  543. switch info.ChannelType {
  544. case constant.ChannelTypeDeepSeek:
  545. if usage.PromptTokensDetails.CachedTokens == 0 && usage.PromptCacheHitTokens != 0 {
  546. usage.PromptTokensDetails.CachedTokens = usage.PromptCacheHitTokens
  547. }
  548. case constant.ChannelTypeZhipu_v4:
  549. // 智普的cached_tokens在标准位置: usage.prompt_tokens_details.cached_tokens
  550. if usage.PromptTokensDetails.CachedTokens == 0 {
  551. if usage.InputTokensDetails != nil && usage.InputTokensDetails.CachedTokens > 0 {
  552. usage.PromptTokensDetails.CachedTokens = usage.InputTokensDetails.CachedTokens
  553. } else if cachedTokens, ok := extractCachedTokensFromBody(responseBody); ok {
  554. usage.PromptTokensDetails.CachedTokens = cachedTokens
  555. } else if usage.PromptCacheHitTokens > 0 {
  556. usage.PromptTokensDetails.CachedTokens = usage.PromptCacheHitTokens
  557. }
  558. }
  559. case constant.ChannelTypeMoonshot:
  560. // Moonshot的cached_tokens在非标准位置: choices[].usage.cached_tokens
  561. if usage.PromptTokensDetails.CachedTokens == 0 {
  562. if usage.InputTokensDetails != nil && usage.InputTokensDetails.CachedTokens > 0 {
  563. usage.PromptTokensDetails.CachedTokens = usage.InputTokensDetails.CachedTokens
  564. } else if cachedTokens, ok := extractMoonshotCachedTokensFromBody(responseBody); ok {
  565. usage.PromptTokensDetails.CachedTokens = cachedTokens
  566. } else if cachedTokens, ok := extractCachedTokensFromBody(responseBody); ok {
  567. usage.PromptTokensDetails.CachedTokens = cachedTokens
  568. } else if usage.PromptCacheHitTokens > 0 {
  569. usage.PromptTokensDetails.CachedTokens = usage.PromptCacheHitTokens
  570. }
  571. }
  572. }
  573. }
  574. func extractCachedTokensFromBody(body []byte) (int, bool) {
  575. if len(body) == 0 {
  576. return 0, false
  577. }
  578. var payload struct {
  579. Usage struct {
  580. PromptTokensDetails struct {
  581. CachedTokens *int `json:"cached_tokens"`
  582. } `json:"prompt_tokens_details"`
  583. CachedTokens *int `json:"cached_tokens"`
  584. PromptCacheHitTokens *int `json:"prompt_cache_hit_tokens"`
  585. } `json:"usage"`
  586. }
  587. if err := common.Unmarshal(body, &payload); err != nil {
  588. return 0, false
  589. }
  590. if payload.Usage.PromptTokensDetails.CachedTokens != nil {
  591. return *payload.Usage.PromptTokensDetails.CachedTokens, true
  592. }
  593. if payload.Usage.CachedTokens != nil {
  594. return *payload.Usage.CachedTokens, true
  595. }
  596. if payload.Usage.PromptCacheHitTokens != nil {
  597. return *payload.Usage.PromptCacheHitTokens, true
  598. }
  599. return 0, false
  600. }
  601. // extractMoonshotCachedTokensFromBody 从Moonshot的非标准位置提取cached_tokens
  602. // Moonshot的流式响应格式: {"choices":[{"usage":{"cached_tokens":111}}]}
  603. func extractMoonshotCachedTokensFromBody(body []byte) (int, bool) {
  604. if len(body) == 0 {
  605. return 0, false
  606. }
  607. var payload struct {
  608. Choices []struct {
  609. Usage struct {
  610. CachedTokens *int `json:"cached_tokens"`
  611. } `json:"usage"`
  612. } `json:"choices"`
  613. }
  614. if err := common.Unmarshal(body, &payload); err != nil {
  615. return 0, false
  616. }
  617. // 遍历choices查找cached_tokens
  618. for _, choice := range payload.Choices {
  619. if choice.Usage.CachedTokens != nil && *choice.Usage.CachedTokens > 0 {
  620. return *choice.Usage.CachedTokens, true
  621. }
  622. }
  623. return 0, false
  624. }