Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 21 additions & 5 deletions relay/channel/openai/relay-openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package openai

import (
"bytes"
"encoding/json"
"fmt"
"io"
"math"
Expand Down Expand Up @@ -280,18 +281,33 @@ func OpenaiTTSHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel
func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo, responseFormat string) (*types.NewAPIError, *dto.Usage) {
defer service.CloseResponseBodyGracefully(resp)

// count tokens by audio file duration
audioTokens, err := countAudioTokens(c)
if err != nil {
return types.NewError(err, types.ErrorCodeCountTokenFailed), nil
}
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
return types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError), nil
}
// 写入新的 response body
service.IOCopyBytesGracefully(c, resp, responseBody)

var responseData struct {
Usage *dto.Usage `json:"usage"`
}
if err := json.Unmarshal(responseBody, &responseData); err == nil && responseData.Usage != nil {
if responseData.Usage.TotalTokens > 0 {
usage := responseData.Usage
if usage.PromptTokens == 0 {
usage.PromptTokens = usage.InputTokens
}
if usage.CompletionTokens == 0 {
usage.CompletionTokens = usage.OutputTokens
}
return nil, usage
}
}
Comment on lines +291 to +305

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🛠️ Refactor suggestion

Normalize upstream usage even when total_tokens is missing; also propagate token details

Many providers return only input/output tokens (and set total_tokens to 0). Your current gate requires TotalTokens > 0, which will incorrectly fall back to local estimation and under/over bill. Normalize as soon as any of Input/Output/Prompt/Completion is present, compute TotalTokens if absent, and map InputTokensDetails into PromptTokensDetails (mirrors OpenaiHandlerWithUsage).

Apply:

 var responseData struct {
   Usage *dto.Usage `json:"usage"`
 }
-if err := json.Unmarshal(responseBody, &responseData); err == nil && responseData.Usage != nil {
-    if responseData.Usage.TotalTokens > 0 {
-        usage := responseData.Usage
-        if usage.PromptTokens == 0 {
-            usage.PromptTokens = usage.InputTokens
-        }
-        if usage.CompletionTokens == 0 {
-            usage.CompletionTokens = usage.OutputTokens
-        }
-        return nil, usage
-    }
-}
+if err := json.Unmarshal(responseBody, &responseData); err == nil && responseData.Usage != nil {
+    u := responseData.Usage
+    // Fill prompt/completion from input/output if missing
+    if u.PromptTokens == 0 && u.InputTokens > 0 {
+        u.PromptTokens = u.InputTokens
+    }
+    if u.CompletionTokens == 0 && u.OutputTokens > 0 {
+        u.CompletionTokens = u.OutputTokens
+    }
+    // Compute total if missing but components exist
+    if u.TotalTokens == 0 && (u.PromptTokens > 0 || u.CompletionTokens > 0) {
+        u.TotalTokens = u.PromptTokens + u.CompletionTokens
+    }
+    // Propagate input token details to prompt details if provided
+    if u.InputTokensDetails != nil {
+        u.PromptTokensDetails.TextTokens += u.InputTokensDetails.TextTokens
+        u.PromptTokensDetails.ImageTokens += u.InputTokensDetails.ImageTokens
+        u.PromptTokensDetails.AudioTokens += u.InputTokensDetails.AudioTokens
+    }
+    if u.TotalTokens > 0 || u.PromptTokens > 0 || u.CompletionTokens > 0 {
+        return nil, u
+    }
+}
📝 Committable suggestion

‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.

Suggested change
var responseData struct {
Usage *dto.Usage `json:"usage"`
}
if err := json.Unmarshal(responseBody, &responseData); err == nil && responseData.Usage != nil {
if responseData.Usage.TotalTokens > 0 {
usage := responseData.Usage
if usage.PromptTokens == 0 {
usage.PromptTokens = usage.InputTokens
}
if usage.CompletionTokens == 0 {
usage.CompletionTokens = usage.OutputTokens
}
return nil, usage
}
}
var responseData struct {
Usage *dto.Usage `json:"usage"`
}
if err := json.Unmarshal(responseBody, &responseData); err == nil && responseData.Usage != nil {
u := responseData.Usage
// Fill prompt/completion from input/output if missing
if u.PromptTokens == 0 && u.InputTokens > 0 {
u.PromptTokens = u.InputTokens
}
if u.CompletionTokens == 0 && u.OutputTokens > 0 {
u.CompletionTokens = u.OutputTokens
}
// Compute total if missing but components exist
if u.TotalTokens == 0 && (u.PromptTokens > 0 || u.CompletionTokens > 0) {
u.TotalTokens = u.PromptTokens + u.CompletionTokens
}
// Propagate input token details to prompt details if provided
if u.InputTokensDetails != nil {
u.PromptTokensDetails.TextTokens += u.InputTokensDetails.TextTokens
u.PromptTokensDetails.ImageTokens += u.InputTokensDetails.ImageTokens
u.PromptTokensDetails.AudioTokens += u.InputTokensDetails.AudioTokens
}
// Return as long as any token data is present
if u.TotalTokens > 0 || u.PromptTokens > 0 || u.CompletionTokens > 0 {
return nil, u
}
}


audioTokens, err := countAudioTokens(c)
if err != nil {
return types.NewError(err, types.ErrorCodeCountTokenFailed), nil
}
Comment on lines +307 to +310

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue

Avoid returning an error after writing the upstream body to the client

By Line 289 you’ve already flushed the upstream response. Returning a non-nil error here risks double-send/error paths and inconsistent retries. Align with the TTS handler: log and return a zero-usage fallback instead of surfacing an error.

Apply:

 audioTokens, err := countAudioTokens(c)
 if err != nil {
-    return types.NewError(err, types.ErrorCodeCountTokenFailed), nil
+    logger.LogError(c, fmt.Sprintf("count audio tokens failed: %v", err))
+    // After body is sent, do not bubble errors; return zero-usage fallback.
+    return nil, &dto.Usage{}
 }
📝 Committable suggestion

‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.

Suggested change
audioTokens, err := countAudioTokens(c)
if err != nil {
return types.NewError(err, types.ErrorCodeCountTokenFailed), nil
}
audioTokens, err := countAudioTokens(c)
if err != nil {
logger.LogError(c, fmt.Sprintf("count audio tokens failed: %v", err))
// After body is sent, do not bubble errors; return zero-usage fallback.
return nil, &dto.Usage{}
}
🤖 Prompt for AI Agents
In relay/channel/openai/relay-openai.go around lines 307–310, do not return a
non-nil error after the upstream response has already been flushed; instead,
catch the countAudioTokens error, log the failure with context, set audioTokens
(or the usage result) to a zero-usage fallback value, and continue execution
returning nil error so we avoid double-send/retry paths (mirror the TTS handler
behavior).

usage := &dto.Usage{}
usage.PromptTokens = audioTokens
usage.CompletionTokens = 0
Expand Down