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
14 changes: 12 additions & 2 deletions dto/audio.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"`
Expand All @@ -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 {
Expand All @@ -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) {
Expand Down
72 changes: 72 additions & 0 deletions dto/audio_test.go
Original file line number Diff line number Diff line change
@@ -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))
})
}
}
43 changes: 40 additions & 3 deletions relay/channel/openai/audio.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"io"
"math"
"net/http"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
Expand Down Expand Up @@ -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)
Expand All @@ -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"`
}
Expand All @@ -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
}
94 changes: 94 additions & 0 deletions relay/channel/openai/audio_test.go
Original file line number Diff line number Diff line change
@@ -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)),
}
}
1 change: 1 addition & 0 deletions relay/helper/valid_request.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down