Skip to content
Open
Show file tree
Hide file tree
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
20 changes: 20 additions & 0 deletions common/body_storage.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@ import (
type BodyStorage interface {
io.ReadSeeker
io.Closer
// Open creates an independent reader for the stored body.
Open() (io.ReadCloser, error)
// Bytes 获取全部内容
Bytes() ([]byte, error)
// Size 获取数据大小
Expand Down Expand Up @@ -71,6 +73,15 @@ func (m *memoryStorage) Close() error {
return nil
}

func (m *memoryStorage) Open() (io.ReadCloser, error) {
m.mu.Lock()
defer m.mu.Unlock()
if atomic.LoadInt32(&m.closed) == 1 {
return nil, ErrStorageClosed
}
return io.NopCloser(bytes.NewReader(m.data)), nil
}

func (m *memoryStorage) Bytes() ([]byte, error) {
m.mu.Lock()
defer m.mu.Unlock()
Expand Down Expand Up @@ -195,6 +206,15 @@ func (d *diskStorage) Close() error {
return nil
}

func (d *diskStorage) Open() (io.ReadCloser, error) {
d.mu.Lock()
defer d.mu.Unlock()
if atomic.LoadInt32(&d.closed) == 1 {
return nil, ErrStorageClosed
}
return os.Open(d.filePath)
}

func (d *diskStorage) Bytes() ([]byte, error) {
d.mu.Lock()
defer d.mu.Unlock()
Expand Down
143 changes: 143 additions & 0 deletions dto/openai_response.go
Original file line number Diff line number Diff line change
Expand Up @@ -292,6 +292,149 @@ type OpenAIResponsesResponse struct {
Metadata json.RawMessage `json:"metadata"`
}

type ResponsesBillingInputTokenDetails struct {
CachedTokens int `json:"cached_tokens"`
AudioTokens int `json:"audio_tokens"`
ImageTokens int `json:"image_tokens"`
}

type ResponsesBillingOutputTokenDetails struct {
ReasoningTokens int `json:"reasoning_tokens"`
}

type ResponsesBillingUsage struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
TotalTokens int `json:"total_tokens"`

InputTokensDetails *ResponsesBillingInputTokenDetails `json:"input_tokens_details"`
OutputTokensDetails *ResponsesBillingOutputTokenDetails `json:"output_tokens_details"`
CompletionTokenDetails *ResponsesBillingOutputTokenDetails `json:"completion_tokens_details"`
}

type ResponsesBillingTool struct {
Type string `json:"type"`
}

type ResponsesBillingOutput struct {
Type string `json:"type"`
Quality string `json:"quality"`
Size string `json:"size"`
}

type ResponsesBillingMeta struct {
Error json.RawMessage `json:"error"`
Usage *ResponsesBillingUsage `json:"usage"`
Tools []ResponsesBillingTool `json:"tools"`
Output []ResponsesBillingOutput `json:"output"`
}

type ResponsesBillingStreamItem struct {
Type string `json:"type"`
}

type ResponsesBillingStreamResponse struct {
Type string `json:"type"`
Response *ResponsesBillingMeta `json:"response,omitempty"`
Delta string `json:"delta,omitempty"`
Item *ResponsesBillingStreamItem `json:"item,omitempty"`
}

type ResponsesTranslatedStreamMeta struct {
Model string `json:"model"`
CreatedAt int `json:"created_at"`
Error json.RawMessage `json:"error"`
Usage *ResponsesBillingUsage `json:"usage"`
}

type ResponsesTranslatedStreamItem struct {
Type string `json:"type"`
ID string `json:"id"`
CallId string `json:"call_id,omitempty"`
Name string `json:"name,omitempty"`
Arguments json.RawMessage `json:"arguments,omitempty"`
}

func (r *ResponsesTranslatedStreamItem) ArgumentsString() string {
if r == nil {
return ""
}
return ResponsesArgumentsString(r.Arguments)
}

type ResponsesTranslatedStreamResponse struct {
Type string `json:"type"`
Response *ResponsesTranslatedStreamMeta `json:"response,omitempty"`
Delta string `json:"delta,omitempty"`
Item *ResponsesTranslatedStreamItem `json:"item,omitempty"`
OutputIndex *int `json:"output_index,omitempty"`
ContentIndex *int `json:"content_index,omitempty"`
SummaryIndex *int `json:"summary_index,omitempty"`
ItemID string `json:"item_id,omitempty"`
Part *ResponsesReasoningSummaryPart `json:"part,omitempty"`
}

func getOpenAIErrorFromRawMessage(errorField json.RawMessage) *types.OpenAIError {
if len(errorField) == 0 || common.GetJsonType(errorField) == "null" {
return nil
}
var decoded any
if err := common.Unmarshal(errorField, &decoded); err != nil {
return nil
}
return GetOpenAIError(decoded)
}

func (m *ResponsesBillingMeta) GetOpenAIError() *types.OpenAIError {
if m == nil {
return nil
}
return getOpenAIErrorFromRawMessage(m.Error)
}

func (m *ResponsesTranslatedStreamMeta) GetOpenAIError() *types.OpenAIError {
if m == nil {
return nil
}
return getOpenAIErrorFromRawMessage(m.Error)
}

func (m *ResponsesBillingMeta) HasImageGenerationCall() bool {
if m == nil || len(m.Output) == 0 {
return false
}
for _, output := range m.Output {
if output.Type == ResponsesOutputTypeImageGenerationCall {
return true
}
}
return false
}

func (m *ResponsesBillingMeta) GetQuality() string {
if m == nil || len(m.Output) == 0 {
return ""
}
for _, output := range m.Output {
if output.Type == ResponsesOutputTypeImageGenerationCall {
return output.Quality
}
}
return ""
}

func (m *ResponsesBillingMeta) GetSize() string {
if m == nil || len(m.Output) == 0 {
return ""
}
for _, output := range m.Output {
if output.Type == ResponsesOutputTypeImageGenerationCall {
return output.Size
}
}
return ""
}

// GetOpenAIError 从动态错误类型中提取OpenAIError结构
func (o *OpenAIResponsesResponse) GetOpenAIError() *types.OpenAIError {
return GetOpenAIError(o.Error)
Expand Down
34 changes: 31 additions & 3 deletions relay/channel/api_request.go
Original file line number Diff line number Diff line change
Expand Up @@ -304,26 +304,51 @@ func applyHeaderOverrideToRequest(req *http.Request, headerOverride map[string]s
}
}

func newUpstreamRequest(method string, url string, requestBody io.Reader) (*http.Request, error) {
if replayableBody, ok := requestBody.(common.ReplayableBody); ok {
body, err := replayableBody.Open()
if err != nil {
return nil, err
}
req, err := http.NewRequest(method, url, body)
if err != nil {
_ = body.Close()
return nil, err
}
req.ContentLength = replayableBody.Size()
return req, nil
}
return http.NewRequest(method, url, requestBody)
}
Comment on lines +307 to +322

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.

🩺 Stability & Availability | 🟡 Minor

🧩 Analysis chain

🌐 Web query:

Go net/http Client when does it use Request.GetBody to retry or redirect a request body

💡 Result:

In Go's net/http package, the Request.GetBody function is used by the http.Client to re-read a request body when it needs to be sent again [1][2]. Because the standard Request.Body is an io.ReadCloser (which is consumed after being read), the client cannot automatically retry a request or follow certain redirects if the body has already been read [3][4]. The client uses GetBody in the following scenarios: 1. Redirects (307 and 308): When the server returns a 307 (Temporary Redirect) or 308 (Permanent Redirect) status code, the client is permitted to follow the redirect while preserving the original HTTP method and body [1]. To send the body again to the new location, the client invokes the GetBody function to obtain a fresh, unread copy of the request body [1][2]. 2. Request Retries: The http.Transport may automatically retry an idempotent request (such as GET, HEAD, OPTIONS, or TRACE) if it encounters a transient network error [2]. If a request has a body, the transport will only perform this retry if GetBody is defined, allowing it to re-open and send the request body again [2]. Important Considerations: - Population: The standard library's http.NewRequest function automatically populates GetBody for common types like strings, byte slices, and bytes.Buffers [1][2]. If you manually create a request and set the Body field using a custom io.ReadCloser (e.g., from an os.File), GetBody will be nil by default, and redirects or retries that require reading the body again will fail [3][4][5]. - Implementation: When creating custom requests, you should define GetBody if you want your request to support these automatic redirects and retries [1][2]. It should be a function that returns a new io.ReadCloser that reads from the start of the data each time it is called [1][2]. - Server-side: GetBody is strictly for client-side requests and is unused when processing requests on the server side [2].

Citations:


🏁 Script executed:

#!/bin/bash
set -euo pipefail

# Inspect the target file and ReplayableBody implementation.
git ls-files | rg '^relay/channel/api_request\.go$|common/.*\.go$'
printf '\n--- api_request.go outline ---\n'
ast-grep outline relay/channel/api_request.go --view expanded || true
printf '\n--- ReplayableBody search ---\n'
rg -n "type ReplayableBody|func \(.*\) Open\(|func \(.*\) Size\(|ReplayableBody" relay common -g '*.go'

Repository: QuantumNous/new-api

Length of output: 4917


🏁 Script executed:

#!/bin/bash
set -euo pipefail

# Read the relevant slices around newUpstreamRequest and ReplayableBody.
sed -n '280,340p' relay/channel/api_request.go
printf '\n---\n'
sed -n '1,220p' common/*.go 2>/dev/null | rg -n "type ReplayableBody|Open\(|Size\(|GetBody|Replayable" -n -C 2 || true

Repository: QuantumNous/new-api

Length of output: 1764


🏁 Script executed:

#!/bin/bash
set -euo pipefail

# Inspect the target file and ReplayableBody implementation.
git ls-files | rg '^relay/channel/api_request\.go$|common/.*\.go$'
printf '\n--- api_request.go outline ---\n'
ast-grep outline relay/channel/api_request.go --view expanded || true
printf '\n--- ReplayableBody search ---\n'
rg -n "type ReplayableBody|func \(.*\) Open\(|func \(.*\) Size\(|ReplayableBody" relay common -g '*.go'

Repository: QuantumNous/new-api

Length of output: 4917


🏁 Script executed:

#!/bin/bash
set -euo pipefail

# Read the relevant slices around newUpstreamRequest and ReplayableBody.
sed -n '280,340p' relay/channel/api_request.go
printf '\n---\n'
sed -n '1,220p' common/*.go 2>/dev/null | rg -n "type ReplayableBody|Open\(|Size\(|GetBody|Replayable" -n -C 2 || true

Repository: QuantumNous/new-api

Length of output: 1764


🏁 Script executed:

#!/bin/bash
set -euo pipefail

# Inspect ReplayableBody and storage implementations.
sed -n '1,140p' relay/common/outbound_body.go
printf '\n--- body_storage.go memory/disk implementations ---\n'
sed -n '1,320p' common/body_storage.go
printf '\n--- request execution path around DoRequest ---\n'
sed -n '330,560p' relay/channel/api_request.go

Repository: QuantumNous/new-api

Length of output: 15348


🏁 Script executed:

#!/bin/bash
set -euo pipefail

# Inspect ReplayableBody and storage implementations.
sed -n '1,140p' relay/common/outbound_body.go
printf '\n--- body_storage.go memory/disk implementations ---\n'
sed -n '1,320p' common/body_storage.go
printf '\n--- request execution path around DoRequest ---\n'
sed -n '330,560p' relay/channel/api_request.go

Repository: QuantumNous/new-api

Length of output: 15348


Set req.GetBody for replayable bodies

http.NewRequest only auto-populates GetBody for a few concrete reader types, so this ReplayableBody still can’t be resent on 307/308 redirects or transport retries. Wire it to replayableBody.Open() here.

♻️ Proposed change
 		req.ContentLength = replayableBody.Size()
+		req.GetBody = func() (io.ReadCloser, error) {
+			return replayableBody.Open()
+		}
 		return req, nil
📝 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
func newUpstreamRequest(method string, url string, requestBody io.Reader) (*http.Request, error) {
if replayableBody, ok := requestBody.(common.ReplayableBody); ok {
body, err := replayableBody.Open()
if err != nil {
return nil, err
}
req, err := http.NewRequest(method, url, body)
if err != nil {
_ = body.Close()
return nil, err
}
req.ContentLength = replayableBody.Size()
return req, nil
}
return http.NewRequest(method, url, requestBody)
}
func newUpstreamRequest(method string, url string, requestBody io.Reader) (*http.Request, error) {
if replayableBody, ok := requestBody.(common.ReplayableBody); ok {
body, err := replayableBody.Open()
if err != nil {
return nil, err
}
req, err := http.NewRequest(method, url, body)
if err != nil {
_ = body.Close()
return nil, err
}
req.ContentLength = replayableBody.Size()
req.GetBody = func() (io.ReadCloser, error) {
return replayableBody.Open()
}
return req, nil
}
return http.NewRequest(method, url, requestBody)
}
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@relay/channel/api_request.go` around lines 307 - 322, The newUpstreamRequest
helper handles common.ReplayableBody but does not make the request reusable for
redirects or retries. Update the replayableBody branch in newUpstreamRequest so
the returned http.Request has GetBody wired to replayableBody.Open(), and keep
ContentLength set from replayableBody.Size() to preserve resendability.


func closeUpstreamRequestBody(req *http.Request) {
if req != nil && req.Body != nil {
_ = req.Body.Close()
}
}

func DoApiRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBody io.Reader) (*http.Response, error) {
fullRequestURL, err := a.GetRequestURL(info)
if err != nil {
return nil, fmt.Errorf("get request url failed: %w", err)
}
logger.LogDebug(c, "fullRequestURL: %s", fullRequestURL)
req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody)
req, err := newUpstreamRequest(c.Request.Method, fullRequestURL, requestBody)
if err != nil {
return nil, fmt.Errorf("new request failed: %w", err)
}
applyUpstreamContentLength(req, info)
headers := req.Header
err = a.SetupRequestHeader(c, &headers, info)
if err != nil {
closeUpstreamRequestBody(req)
return nil, fmt.Errorf("setup request header failed: %w", err)
}
// 在 SetupRequestHeader 之后应用 Header Override,确保用户设置优先级最高
// 这样可以覆盖默认的 Authorization header 设置
headerOverride, err := processHeaderOverride(info, c)
if err != nil {
closeUpstreamRequestBody(req)
return nil, err
}
applyHeaderOverrideToRequest(req, headerOverride)
Expand All @@ -340,7 +365,7 @@ func DoFormRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBod
return nil, fmt.Errorf("get request url failed: %w", err)
}
logger.LogDebug(c, "fullRequestURL: %s", fullRequestURL)
req, err := http.NewRequest(c.Request.Method, fullRequestURL, requestBody)
req, err := newUpstreamRequest(c.Request.Method, fullRequestURL, requestBody)
if err != nil {
return nil, fmt.Errorf("new request failed: %w", err)
}
Expand All @@ -350,12 +375,14 @@ func DoFormRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBod
headers := req.Header
err = a.SetupRequestHeader(c, &headers, info)
if err != nil {
closeUpstreamRequestBody(req)
return nil, fmt.Errorf("setup request header failed: %w", err)
}
// 在 SetupRequestHeader 之后应用 Header Override,确保用户设置优先级最高
// 这样可以覆盖默认的 Authorization header 设置
headerOverride, err := processHeaderOverride(info, c)
if err != nil {
closeUpstreamRequestBody(req)
return nil, err
}
applyHeaderOverrideToRequest(req, headerOverride)
Expand Down Expand Up @@ -514,6 +541,8 @@ func doRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http
}
}

defer closeUpstreamRequestBody(req)

resp, err := client.Do(req)
if err != nil {
logger.LogError(c, "do request failed: "+err.Error())
Expand All @@ -527,7 +556,6 @@ func doRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http
c.Set(common2.UpstreamRequestIdKey, upID)
}

_ = req.Body.Close()
_ = c.Request.Body.Close()
return resp, nil
}
Expand Down
6 changes: 4 additions & 2 deletions relay/channel/openai/chat_via_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -302,7 +302,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
return
}

var streamResp dto.ResponsesStreamResponse
var streamResp dto.ResponsesTranslatedStreamResponse
if err := common.UnmarshalJsonStr(data, &streamResp); err != nil {
logger.LogError(c, "failed to unmarshal responses stream event: "+err.Error())
sr.Error(err)
Expand Down Expand Up @@ -469,7 +469,9 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
usage.PromptTokensDetails.ImageTokens = streamResp.Response.Usage.InputTokensDetails.ImageTokens
usage.PromptTokensDetails.AudioTokens = streamResp.Response.Usage.InputTokensDetails.AudioTokens
}
if streamResp.Response.Usage.CompletionTokenDetails.ReasoningTokens != 0 {
if streamResp.Response.Usage.OutputTokensDetails != nil && streamResp.Response.Usage.OutputTokensDetails.ReasoningTokens != 0 {
usage.CompletionTokenDetails.ReasoningTokens = streamResp.Response.Usage.OutputTokensDetails.ReasoningTokens
} else if streamResp.Response.Usage.CompletionTokenDetails != nil && streamResp.Response.Usage.CompletionTokenDetails.ReasoningTokens != 0 {
usage.CompletionTokenDetails.ReasoningTokens = streamResp.Response.Usage.CompletionTokenDetails.ReasoningTokens
}
}
Expand Down
4 changes: 2 additions & 2 deletions relay/channel/openai/helper.go
Original file line number Diff line number Diff line change
Expand Up @@ -202,9 +202,9 @@ func HandleFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, lastStream
}
}

func sendResponsesStreamData(c *gin.Context, streamResponse dto.ResponsesStreamResponse, data string) {
func sendResponsesStreamData(c *gin.Context, eventType string, data string) {
if data == "" {
return
}
helper.ResponseChunkData(c, streamResponse, data)
helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: eventType}, data)
}
34 changes: 17 additions & 17 deletions relay/channel/openai/relay_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,46 +21,46 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
defer service.CloseResponseBodyGracefully(resp)

// read response body
var responsesResponse dto.OpenAIResponsesResponse
var billingMeta dto.ResponsesBillingMeta
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
}
err = common.Unmarshal(responseBody, &responsesResponse)
err = common.Unmarshal(responseBody, &billingMeta)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if oaiError := responsesResponse.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
if oaiError := billingMeta.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
}

if responsesResponse.HasImageGenerationCall() {
if billingMeta.HasImageGenerationCall() {
c.Set("image_generation_call", true)
c.Set("image_generation_call_quality", responsesResponse.GetQuality())
c.Set("image_generation_call_size", responsesResponse.GetSize())
c.Set("image_generation_call_quality", billingMeta.GetQuality())
c.Set("image_generation_call_size", billingMeta.GetSize())
}

// 写入新的 response body
service.IOCopyBytesGracefully(c, resp, responseBody)

// compute usage
usage := dto.Usage{}
if responsesResponse.Usage != nil {
usage.PromptTokens = responsesResponse.Usage.InputTokens
usage.CompletionTokens = responsesResponse.Usage.OutputTokens
usage.TotalTokens = responsesResponse.Usage.TotalTokens
if responsesResponse.Usage.InputTokensDetails != nil {
usage.PromptTokensDetails.CachedTokens = responsesResponse.Usage.InputTokensDetails.CachedTokens
if billingMeta.Usage != nil {
usage.PromptTokens = billingMeta.Usage.InputTokens
usage.CompletionTokens = billingMeta.Usage.OutputTokens
usage.TotalTokens = billingMeta.Usage.TotalTokens
if billingMeta.Usage.InputTokensDetails != nil {
usage.PromptTokensDetails.CachedTokens = billingMeta.Usage.InputTokensDetails.CachedTokens
}
}
if info == nil || info.ResponsesUsageInfo == nil || info.ResponsesUsageInfo.BuiltInTools == nil {
return &usage, nil
}
// 解析 Tools 用量
for _, tool := range responsesResponse.Tools {
buildToolinfo, ok := info.ResponsesUsageInfo.BuiltInTools[common.Interface2String(tool["type"])]
for _, tool := range billingMeta.Tools {
buildToolinfo, ok := info.ResponsesUsageInfo.BuiltInTools[tool.Type]
if !ok || buildToolinfo == nil {
logger.LogError(c, fmt.Sprintf("BuiltInTools not found for tool type: %v", tool["type"]))
logger.LogError(c, fmt.Sprintf("BuiltInTools not found for tool type: %v", tool.Type))
continue
}
buildToolinfo.CallCount++
Expand All @@ -82,13 +82,13 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp
helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) {

// 检查当前数据是否包含 completed 状态和 usage 信息
var streamResponse dto.ResponsesStreamResponse
var streamResponse dto.ResponsesBillingStreamResponse
if err := common.UnmarshalJsonStr(data, &streamResponse); err != nil {
logger.LogError(c, "failed to unmarshal stream response: "+err.Error())
sr.Error(err)
return
}
sendResponsesStreamData(c, streamResponse, data)
sendResponsesStreamData(c, streamResponse.Type, data)
switch streamResponse.Type {
case "response.completed":
if streamResponse.Response != nil {
Expand Down
Loading