-
Notifications
You must be signed in to change notification settings - Fork 11.1k
feat: use audio token usage if return #1721
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
|
@@ -2,6 +2,7 @@ package openai | |||||||||||||||||||||
|
|
||||||||||||||||||||||
| import ( | ||||||||||||||||||||||
| "bytes" | ||||||||||||||||||||||
| "encoding/json" | ||||||||||||||||||||||
| "fmt" | ||||||||||||||||||||||
| "io" | ||||||||||||||||||||||
| "math" | ||||||||||||||||||||||
|
|
@@ -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 | ||||||||||||||||||||||
| } | ||||||||||||||||||||||
| } | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| audioTokens, err := countAudioTokens(c) | ||||||||||||||||||||||
| if err != nil { | ||||||||||||||||||||||
| return types.NewError(err, types.ErrorCodeCountTokenFailed), nil | ||||||||||||||||||||||
| } | ||||||||||||||||||||||
|
Comment on lines
+307
to
+310
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 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
Suggested change
🤖 Prompt for AI Agents |
||||||||||||||||||||||
| usage := &dto.Usage{} | ||||||||||||||||||||||
| usage.PromptTokens = audioTokens | ||||||||||||||||||||||
| usage.CompletionTokens = 0 | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
There was a problem hiding this comment.
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