From 6436d9a97b4875305a91567dc76129a743e93fe1 Mon Sep 17 00:00:00 2001 From: SHEN CE Date: Mon, 8 Jun 2026 17:33:25 +0800 Subject: [PATCH] fix: stream audio transcription responses --- dto/audio.go | 14 ++++- dto/audio_test.go | 72 +++++++++++++++++++++++ relay/channel/openai/audio.go | 43 +++++++++++++- relay/channel/openai/audio_test.go | 94 ++++++++++++++++++++++++++++++ relay/helper/valid_request.go | 1 + 5 files changed, 219 insertions(+), 5 deletions(-) create mode 100644 dto/audio_test.go create mode 100644 relay/channel/openai/audio_test.go diff --git a/dto/audio.go b/dto/audio.go index e0d4f9d07c7e..0466b3f44799 100644 --- a/dto/audio.go +++ b/dto/audio.go @@ -17,6 +17,7 @@ type AudioRequest struct { ResponseFormat string `json:"response_format,omitempty"` Speed *float64 `json:"speed,omitempty"` StreamFormat string `json:"stream_format,omitempty"` + Stream *BoolValue `json:"stream,omitempty"` Metadata json.RawMessage `json:"metadata,omitempty"` // vllm-omini TaskType json.RawMessage `json:"task_type,omitempty"` @@ -27,7 +28,6 @@ type AudioRequest struct { MaxNewTokens json.RawMessage `json:"max_new_tokens,omitempty"` InitialCodecChunkFrames json.RawMessage `json:"initial_codec_chunk_frames,omitempty"` // TODO:ensure that the logic remains correct after the stream is started. - //Stream json.RawMessage `json:"stream,omitempty"` } func (r *AudioRequest) GetTokenCountMeta() *types.TokenCountMeta { @@ -42,7 +42,17 @@ func (r *AudioRequest) GetTokenCountMeta() *types.TokenCountMeta { } func (r *AudioRequest) IsStream(c *gin.Context) bool { - return r.StreamFormat == "sse" + if r.StreamFormat == "sse" { + return true + } + if r.Stream == nil || c == nil || c.Request == nil || c.Request.URL == nil { + return false + } + path := c.Request.URL.Path + if !strings.HasSuffix(path, "/audio/transcriptions") && !strings.HasSuffix(path, "/audio/translations") { + return false + } + return bool(*r.Stream) } func (r *AudioRequest) SetModelName(modelName string) { diff --git a/dto/audio_test.go b/dto/audio_test.go new file mode 100644 index 000000000000..ab1f8cc5f1ae --- /dev/null +++ b/dto/audio_test.go @@ -0,0 +1,72 @@ +package dto + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestAudioRequestIsStream(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + path string + raw string + expected bool + }{ + { + name: "transcription parsed multipart stream true", + path: "/v1/audio/transcriptions", + raw: `{"stream":"true"}`, + expected: true, + }, + { + name: "translation json stream true", + path: "/v1/audio/translations", + raw: `{"stream":true}`, + expected: true, + }, + { + name: "transcription stream false", + path: "/v1/audio/transcriptions", + raw: `{"stream":"false"}`, + expected: false, + }, + { + name: "transcription stream missing", + path: "/v1/audio/transcriptions", + raw: `{}`, + expected: false, + }, + { + name: "speech stream true does not trigger stt stream", + path: "/v1/audio/speech", + raw: `{"stream":"true"}`, + expected: false, + }, + { + name: "stream format sse keeps existing behavior", + path: "/v1/audio/speech", + raw: `{"stream_format":"sse"}`, + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := &AudioRequest{} + require.NoError(t, common.Unmarshal([]byte(tt.raw), req)) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, tt.path, nil) + + require.Equal(t, tt.expected, req.IsStream(c)) + }) + } +} diff --git a/relay/channel/openai/audio.go b/relay/channel/openai/audio.go index 6a87d89f2ea8..02a165840f3c 100644 --- a/relay/channel/openai/audio.go +++ b/relay/channel/openai/audio.go @@ -6,6 +6,7 @@ import ( "io" "math" "net/http" + "strings" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/constant" @@ -115,6 +116,21 @@ 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) { + if shouldStreamSTTResponse(resp, info) { + usage := fallbackSTTUsage(info) + helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { + if service.SundaySearch(data, "usage") { + if parsedUsage := parseSTTUsage([]byte(data)); parsedUsage != nil { + usage = parsedUsage + } + } + if err := helper.StringData(c, data); err != nil { + sr.Error(err) + } + }) + return nil, usage + } + defer service.CloseResponseBodyGracefully(resp) responseBody, err := io.ReadAll(resp.Body) @@ -124,6 +140,22 @@ func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel // 写入新的 response body service.IOCopyBytesGracefully(c, resp, responseBody) + if usage := parseSTTUsage(responseBody); usage != nil { + return nil, usage + } + + return nil, fallbackSTTUsage(info) +} + +func shouldStreamSTTResponse(resp *http.Response, info *relaycommon.RelayInfo) bool { + if resp == nil || info == nil || !info.IsStream { + return false + } + contentType := strings.ToLower(resp.Header.Get("Content-Type")) + return strings.HasPrefix(contentType, "text/event-stream") +} + +func parseSTTUsage(responseBody []byte) *dto.Usage { var responseData struct { Usage *dto.Usage `json:"usage"` } @@ -136,13 +168,18 @@ func OpenaiSTTHandler(c *gin.Context, resp *http.Response, info *relaycommon.Rel if usage.CompletionTokens == 0 { usage.CompletionTokens = usage.OutputTokens } - return nil, usage + return usage } } + return nil +} +func fallbackSTTUsage(info *relaycommon.RelayInfo) *dto.Usage { usage := &dto.Usage{} - usage.PromptTokens = info.GetEstimatePromptTokens() + if info != nil { + usage.PromptTokens = info.GetEstimatePromptTokens() + } usage.CompletionTokens = 0 usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens - return nil, usage + return usage } diff --git a/relay/channel/openai/audio_test.go b/relay/channel/openai/audio_test.go new file mode 100644 index 000000000000..afdc96a89141 --- /dev/null +++ b/relay/channel/openai/audio_test.go @@ -0,0 +1,94 @@ +package openai + +import ( + "io" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + + "github.com/QuantumNous/new-api/constant" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestOpenaiSTTHandlerStreamsEventStreamResponse(t *testing.T) { + c, recorder := newAudioTestContext(t) + info := &relaycommon.RelayInfo{IsStream: true} + info.SetEstimatePromptTokens(7) + resp := newAudioTestResponse("text/event-stream", ""+ + "data: {\"text\":\"hello\"}\n\n"+ + "data: [DONE]\n\n") + + err, usage := OpenaiSTTHandler(c, resp, info, "json") + + require.Nil(t, err) + require.Equal(t, 7, usage.PromptTokens) + require.Equal(t, 7, usage.TotalTokens) + require.True(t, recorder.Flushed) + require.Equal(t, "text/event-stream", recorder.Header().Get("Content-Type")) + require.Contains(t, recorder.Body.String(), "data: {\"text\":\"hello\"}\n\n") +} + +func TestOpenaiSTTHandlerKeepsNonStreamJSONBehavior(t *testing.T) { + c, recorder := newAudioTestContext(t) + info := &relaycommon.RelayInfo{IsStream: true} + info.SetEstimatePromptTokens(7) + respBody := `{"text":"ok","usage":{"total_tokens":12,"input_tokens":5,"output_tokens":7}}` + resp := newAudioTestResponse("application/json", respBody) + + err, usage := OpenaiSTTHandler(c, resp, info, "json") + + require.Nil(t, err) + require.Equal(t, respBody, recorder.Body.String()) + require.Equal(t, strconv.Itoa(len(respBody)), recorder.Header().Get("Content-Length")) + require.Equal(t, 5, usage.PromptTokens) + require.Equal(t, 7, usage.CompletionTokens) + require.Equal(t, 12, usage.TotalTokens) +} + +func TestOpenaiSTTHandlerUsesStreamUsageChunk(t *testing.T) { + c, recorder := newAudioTestContext(t) + info := &relaycommon.RelayInfo{IsStream: true} + info.SetEstimatePromptTokens(7) + resp := newAudioTestResponse("text/event-stream; charset=utf-8", ""+ + "data: {\"text\":\"hello\"}\n\n"+ + "data: {\"usage\":{\"total_tokens\":9,\"input_tokens\":4,\"output_tokens\":5}}\n\n"+ + "data: [DONE]\n\n") + + err, usage := OpenaiSTTHandler(c, resp, info, "json") + + require.Nil(t, err) + require.Equal(t, 4, usage.PromptTokens) + require.Equal(t, 5, usage.CompletionTokens) + require.Equal(t, 9, usage.TotalTokens) + require.True(t, recorder.Flushed) + require.Contains(t, recorder.Body.String(), "data: {\"usage\":{\"total_tokens\":9,\"input_tokens\":4,\"output_tokens\":5}}\n\n") +} + +func newAudioTestContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + gin.SetMode(gin.TestMode) + oldTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 30 + t.Cleanup(func() { + constant.StreamingTimeout = oldTimeout + }) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/audio/transcriptions", nil) + return c, recorder +} + +func newAudioTestResponse(contentType string, body string) *http.Response { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{contentType}, + }, + Body: io.NopCloser(strings.NewReader(body)), + } +} diff --git a/relay/helper/valid_request.go b/relay/helper/valid_request.go index 2581b2812c94..072c00c8976e 100644 --- a/relay/helper/valid_request.go +++ b/relay/helper/valid_request.go @@ -64,6 +64,7 @@ func GetAndValidAudioRequest(c *gin.Context, relayMode int) (*dto.AudioRequest, if audioRequest.Model == "" { return nil, errors.New("model is required") } + audioRequest.Stream = nil default: if audioRequest.Model == "" { return nil, errors.New("model is required")