From 66ee6b8f9889050ffef1f863a4314ce4a0516fb9 Mon Sep 17 00:00:00 2001 From: Scott <908181134@qq.com> Date: Wed, 29 Jul 2026 17:45:14 +0800 Subject: [PATCH] fix: preserve Qwen thinking_budget passthrough (#5836) * fix: preserve qwen thinking budget * test: address qwen thinking budget review comments * chore: remove unreachable adaptor code * test: cover zero Qwen thinking budgets --- relay/channel/ali/adaptor.go | 2 +- relay/channel/ali/adaptor_test.go | 128 ++++++++++++++++++ relay/channel/ali/text.go | 10 +- relay/channel/baidu/adaptor.go | 1 - relay/channel/cloudflare/adaptor.go | 1 - relay/channel/cohere/adaptor.go | 1 - relay/channel/dify/adaptor.go | 2 - relay/channel/jina/adaptor.go | 1 - relay/channel/mistral/adaptor.go | 1 - relay/channel/mokaai/adaptor.go | 1 - relay/channel/palm/adaptor.go | 1 - relay/channel/tencent/adaptor.go | 1 - relay/channel/xunfei/adaptor.go | 1 - relay/channel/zhipu/adaptor.go | 1 - relaykit/dto/openai_request.go | 26 ++++ .../dto/openai_request_zero_value_test.go | 107 +++++++++++++++ .../internal/oai_chat/to_oai_responses_req.go | 2 + .../oai_chat/to_oai_responses_req_test.go | 38 ++++++ .../internal/oai_responses/to_oai_chat_req.go | 1 + .../oai_responses/to_oai_chat_req_test.go | 33 +++++ 20 files changed, 345 insertions(+), 14 deletions(-) create mode 100644 relay/channel/ali/adaptor_test.go diff --git a/relay/channel/ali/adaptor.go b/relay/channel/ali/adaptor.go index ba377659..2cf4bf96 100644 --- a/relay/channel/ali/adaptor.go +++ b/relay/channel/ali/adaptor.go @@ -176,7 +176,7 @@ func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayIn switch info.RelayMode { default: - aliReq := requestOpenAI2Ali(*request) + aliReq := requestOpenAI2Ali(*request, info.UpstreamModelName) return aliReq, nil } } diff --git a/relay/channel/ali/adaptor_test.go b/relay/channel/ali/adaptor_test.go new file mode 100644 index 00000000..a8b87140 --- /dev/null +++ b/relay/channel/ali/adaptor_test.go @@ -0,0 +1,128 @@ +package ali + +import ( + "encoding/json" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayhelper "github.com/QuantumNous/new-api/relay/helper" + "github.com/QuantumNous/new-api/relaykit/dto" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestConvertOpenAIRequestFiltersThinkingBudgetByUpstreamModel(t *testing.T) { + tests := []struct { + name string + requestModel string + upstreamModel string + budget string + wantBudget bool + wantValue int64 + }{ + { + name: "qwen", + requestModel: "qwen-plus", + upstreamModel: "qwen-plus", + budget: "128", + wantBudget: true, + wantValue: 128, + }, + { + name: "qwq explicit zero", + requestModel: "qwq-32b", + upstreamModel: "qwq-32b", + budget: "0", + wantBudget: true, + wantValue: 0, + }, + { + name: "unsupported upstream overrides qwen request", + requestModel: "qwen-plus", + upstreamModel: "deepseek-r1", + budget: "128", + wantBudget: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + request := &dto.GeneralOpenAIRequest{ + Model: tt.requestModel, + EnableThinking: json.RawMessage(`true`), + ThinkingBudget: json.RawMessage(tt.budget), + } + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: tt.upstreamModel, + }, + } + + convertedValue, err := (&Adaptor{}).ConvertOpenAIRequest(nil, info, request) + require.NoError(t, err) + converted, ok := convertedValue.(*dto.GeneralOpenAIRequest) + require.True(t, ok) + + if tt.wantBudget { + assert.Equal(t, tt.budget, string(converted.ThinkingBudget)) + } else { + assert.Nil(t, converted.ThinkingBudget) + } + + encoded, err := common.Marshal(converted) + require.NoError(t, err) + + assert.True(t, gjson.GetBytes(encoded, "enable_thinking").Bool()) + value := gjson.GetBytes(encoded, "thinking_budget") + assert.Equal(t, tt.wantBudget, value.Exists()) + if tt.wantBudget { + assert.Equal(t, tt.wantValue, value.Int()) + } + }) + } +} + +func TestConvertOpenAIRequestPreservesExplicitZeroForMappedQwenModel(t *testing.T) { + const ( + clientModel = "customer-model" + upstreamModel = "Qwen/Qwen3-235B-A22B-Thinking-2507" + ) + + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Set("model_mapping", `{"customer-model":"Qwen/Qwen3-235B-A22B-Thinking-2507"}`) + + request := &dto.GeneralOpenAIRequest{ + Model: clientModel, + EnableThinking: json.RawMessage(`true`), + ThinkingBudget: json.RawMessage(`0`), + } + info := &relaycommon.RelayInfo{ + OriginModelName: clientModel, + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: clientModel, + }, + } + + err := relayhelper.ModelMappedHelper(c, info, request) + require.NoError(t, err) + assert.True(t, info.IsModelMapped) + assert.Equal(t, upstreamModel, info.UpstreamModelName) + assert.Equal(t, upstreamModel, request.Model) + + convertedValue, err := (&Adaptor{}).ConvertOpenAIRequest(c, info, request) + require.NoError(t, err) + converted, ok := convertedValue.(*dto.GeneralOpenAIRequest) + require.True(t, ok) + assert.Equal(t, json.RawMessage(`0`), converted.ThinkingBudget) + + encoded, err := common.Marshal(converted) + require.NoError(t, err) + + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, int64(0), value.Int()) +} diff --git a/relay/channel/ali/text.go b/relay/channel/ali/text.go index 6e532f9d..eea13129 100644 --- a/relay/channel/ali/text.go +++ b/relay/channel/ali/text.go @@ -9,7 +9,15 @@ import ( const EnableSearchModelSuffix = "-internet" -func requestOpenAI2Ali(request dto.GeneralOpenAIRequest) *dto.GeneralOpenAIRequest { +func requestOpenAI2Ali(request dto.GeneralOpenAIRequest, upstreamModelName string) *dto.GeneralOpenAIRequest { + modelName := upstreamModelName + if modelName == "" { + modelName = request.Model + } + if !dto.IsQwenThinkingBudgetModel(modelName) { + request.ThinkingBudget = nil + } + topP := lo.FromPtrOr(request.TopP, 0) if topP >= 1 { request.TopP = lo.ToPtr(0.999) diff --git a/relay/channel/baidu/adaptor.go b/relay/channel/baidu/adaptor.go index fc300a94..e2cd7949 100644 --- a/relay/channel/baidu/adaptor.go +++ b/relay/channel/baidu/adaptor.go @@ -27,7 +27,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/cloudflare/adaptor.go b/relay/channel/cloudflare/adaptor.go index 172db70a..8fbbd286 100644 --- a/relay/channel/cloudflare/adaptor.go +++ b/relay/channel/cloudflare/adaptor.go @@ -28,7 +28,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) Init(info *relaycommon.RelayInfo) { diff --git a/relay/channel/cohere/adaptor.go b/relay/channel/cohere/adaptor.go index b6b68f1b..f9566f95 100644 --- a/relay/channel/cohere/adaptor.go +++ b/relay/channel/cohere/adaptor.go @@ -26,7 +26,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/dify/adaptor.go b/relay/channel/dify/adaptor.go index e88436e1..e8a274d7 100644 --- a/relay/channel/dify/adaptor.go +++ b/relay/channel/dify/adaptor.go @@ -33,7 +33,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { @@ -109,7 +108,6 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom } else { return difyHandler(c, info, resp) } - return } func (a *Adaptor) GetModelList() []string { diff --git a/relay/channel/jina/adaptor.go b/relay/channel/jina/adaptor.go index 13f8cd7b..578745f2 100644 --- a/relay/channel/jina/adaptor.go +++ b/relay/channel/jina/adaptor.go @@ -28,7 +28,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/mistral/adaptor.go b/relay/channel/mistral/adaptor.go index fa75b03b..f9573be3 100644 --- a/relay/channel/mistral/adaptor.go +++ b/relay/channel/mistral/adaptor.go @@ -25,7 +25,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/mokaai/adaptor.go b/relay/channel/mokaai/adaptor.go index 3233308f..ff061b35 100644 --- a/relay/channel/mokaai/adaptor.go +++ b/relay/channel/mokaai/adaptor.go @@ -27,7 +27,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/palm/adaptor.go b/relay/channel/palm/adaptor.go index a50c979b..bd1a3898 100644 --- a/relay/channel/palm/adaptor.go +++ b/relay/channel/palm/adaptor.go @@ -26,7 +26,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/tencent/adaptor.go b/relay/channel/tencent/adaptor.go index 69221b55..2a8964bc 100644 --- a/relay/channel/tencent/adaptor.go +++ b/relay/channel/tencent/adaptor.go @@ -34,7 +34,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/xunfei/adaptor.go b/relay/channel/xunfei/adaptor.go index d92029e2..2f8112f4 100644 --- a/relay/channel/xunfei/adaptor.go +++ b/relay/channel/xunfei/adaptor.go @@ -26,7 +26,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relay/channel/zhipu/adaptor.go b/relay/channel/zhipu/adaptor.go index d740fc5f..09a6c4a6 100644 --- a/relay/channel/zhipu/adaptor.go +++ b/relay/channel/zhipu/adaptor.go @@ -26,7 +26,6 @@ func (a *Adaptor) ConvertGeminiRequest(*gin.Context, *relaycommon.RelayInfo, *dt func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) { //TODO implement me panic("implement me") - return nil, nil } func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { diff --git a/relaykit/dto/openai_request.go b/relaykit/dto/openai_request.go index e646658f..c7d8b60d 100644 --- a/relaykit/dto/openai_request.go +++ b/relaykit/dto/openai_request.go @@ -89,6 +89,7 @@ type GeneralOpenAIRequest struct { // Ali Qwen Params VlHighResolutionImages json.RawMessage `json:"vl_high_resolution_images,omitempty"` EnableThinking json.RawMessage `json:"enable_thinking,omitempty"` + ThinkingBudget json.RawMessage `json:"thinking_budget,omitempty"` ChatTemplateKwargs json.RawMessage `json:"chat_template_kwargs,omitempty"` EnableSearch json.RawMessage `json:"enable_search,omitempty"` // ollama Params @@ -107,6 +108,14 @@ type GeneralOpenAIRequest struct { ReasoningSplit json.RawMessage `json:"reasoning_split,omitempty"` } +func (r GeneralOpenAIRequest) MarshalJSON() ([]byte, error) { + type Alias GeneralOpenAIRequest + if !IsQwenThinkingBudgetModel(r.Model) { + r.ThinkingBudget = nil + } + return kitutil.Marshal((*Alias)(&r)) +} + func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta { var tokenCountMeta types.TokenCountMeta var texts = make([]string, 0) @@ -222,6 +231,14 @@ func IsOpenAIGPT5Model(modelName string) bool { return strings.HasPrefix(modelName, "gpt-5") } +func IsQwenThinkingBudgetModel(modelName string) bool { + normalized := strings.ToLower(strings.TrimSpace(modelName)) + return strings.HasPrefix(normalized, "qwen") || + strings.Contains(normalized, "/qwen") || + strings.HasPrefix(normalized, "qwq") || + strings.Contains(normalized, "/qwq") +} + func (r *GeneralOpenAIRequest) GetSystemRoleName() string { if IsOpenAIReasoningOModel(r.Model) { if !strings.HasPrefix(r.Model, "o1-mini") && !strings.HasPrefix(r.Model, "o1-preview") { @@ -880,10 +897,19 @@ type OpenAIResponsesRequest struct { ClientMetadata json.RawMessage `json:"client_metadata,omitempty"` // qwen EnableThinking json.RawMessage `json:"enable_thinking,omitempty"` + ThinkingBudget json.RawMessage `json:"thinking_budget,omitempty"` // perplexity Preset json.RawMessage `json:"preset,omitempty"` } +func (r OpenAIResponsesRequest) MarshalJSON() ([]byte, error) { + type Alias OpenAIResponsesRequest + if !IsQwenThinkingBudgetModel(r.Model) { + r.ThinkingBudget = nil + } + return kitutil.Marshal((*Alias)(&r)) +} + func (r *OpenAIResponsesRequest) GetTokenCountMeta() *types.TokenCountMeta { var fileMeta = make([]*types.FileMeta, 0) var texts = make([]string, 0) diff --git a/relaykit/dto/openai_request_zero_value_test.go b/relaykit/dto/openai_request_zero_value_test.go index cd3c1bea..ecd83fd8 100644 --- a/relaykit/dto/openai_request_zero_value_test.go +++ b/relaykit/dto/openai_request_zero_value_test.go @@ -1,9 +1,11 @@ package dto import ( + "encoding/json" "testing" kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) @@ -50,6 +52,71 @@ func TestGeneralOpenAIRequestPreserveExplicitZeroValues(t *testing.T) { require.True(t, gjson.GetBytes(encoded, "return_related_questions").Exists()) } +func TestGeneralOpenAIRequestPreserveQwenThinkingBudget(t *testing.T) { + raw := []byte(`{ + "model":"qwen-plus", + "thinking_budget":0 + }`) + + var req GeneralOpenAIRequest + err := kitutil.Unmarshal(raw, &req) + require.NoError(t, err) + + encoded, err := kitutil.Marshal(req) + require.NoError(t, err) + + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, int64(0), value.Int()) +} + +func TestGeneralOpenAIRequestPreserveQwQThinkingBudget(t *testing.T) { + req := GeneralOpenAIRequest{ + Model: "QwQ-32B", + ThinkingBudget: json.RawMessage(`128`), + } + + encoded, err := kitutil.Marshal(req) + require.NoError(t, err) + + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, int64(128), value.Int()) +} + +func TestGeneralOpenAIRequestDropsThinkingBudgetForNonQwenModel(t *testing.T) { + req := GeneralOpenAIRequest{ + Model: "gpt-4.1", + ThinkingBudget: json.RawMessage(`128`), + } + + encoded, err := kitutil.Marshal(req) + require.NoError(t, err) + + assert.False(t, gjson.GetBytes(encoded, "thinking_budget").Exists()) +} + +func TestIsQwenThinkingBudgetModel(t *testing.T) { + tests := []struct { + model string + want bool + }{ + {model: "qwen-plus", want: true}, + {model: "Qwen/Qwen3-235B-A22B-Thinking-2507", want: true}, + {model: "qwq-32b", want: true}, + {model: "provider/qwen-plus", want: true}, + {model: "provider/qwq-32b", want: true}, + {model: "gpt-4.1", want: false}, + {model: "deepseek-r1", want: false}, + } + + for _, tt := range tests { + t.Run(tt.model, func(t *testing.T) { + assert.Equal(t, tt.want, IsQwenThinkingBudgetModel(tt.model)) + }) + } +} + func TestOpenAIResponsesRequestPreserveExplicitZeroValues(t *testing.T) { raw := []byte(`{ "model":"gpt-4.1", @@ -72,6 +139,46 @@ func TestOpenAIResponsesRequestPreserveExplicitZeroValues(t *testing.T) { require.True(t, gjson.GetBytes(encoded, "top_p").Exists()) } +func TestOpenAIResponsesRequestPreserveQwenThinkingBudget(t *testing.T) { + req := OpenAIResponsesRequest{ + Model: "qwen-plus", + ThinkingBudget: json.RawMessage(`0`), + } + + encoded, err := kitutil.Marshal(req) + require.NoError(t, err) + + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, int64(0), value.Int()) +} + +func TestOpenAIResponsesRequestPreserveQwQThinkingBudget(t *testing.T) { + req := OpenAIResponsesRequest{ + Model: "provider/QwQ-32B", + ThinkingBudget: json.RawMessage(`128`), + } + + encoded, err := kitutil.Marshal(req) + require.NoError(t, err) + + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, int64(128), value.Int()) +} + +func TestOpenAIResponsesRequestDropsThinkingBudgetForNonQwenModel(t *testing.T) { + req := OpenAIResponsesRequest{ + Model: "gpt-4.1", + ThinkingBudget: json.RawMessage(`128`), + } + + encoded, err := kitutil.Marshal(req) + require.NoError(t, err) + + assert.False(t, gjson.GetBytes(encoded, "thinking_budget").Exists()) +} + func TestGeneralOpenAIRequestGetSystemRoleName(t *testing.T) { tests := []struct { name string diff --git a/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req.go b/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req.go index b4acb511..ec51248b 100644 --- a/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req.go +++ b/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req.go @@ -386,6 +386,8 @@ func ChatCompletionsRequestToResponsesRequest(req *dto.GeneralOpenAIRequest) (*d ParallelToolCalls: parallelToolCallsRaw, Store: req.Store, Metadata: req.Metadata, + EnableThinking: req.EnableThinking, + ThinkingBudget: req.ThinkingBudget, } if req.MaxTokens != nil || req.MaxCompletionTokens != nil { out.MaxOutputTokens = lo.ToPtr(maxOutputTokens) diff --git a/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req_test.go b/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req_test.go index 095a540d..91bcae4e 100644 --- a/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req_test.go +++ b/relaykit/relayconvert/internal/oai_chat/to_oai_responses_req_test.go @@ -1,9 +1,11 @@ package oaichat import ( + "encoding/json" "testing" "github.com/QuantumNous/new-api/relaykit/dto" + kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil" "github.com/samber/lo" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -37,6 +39,42 @@ func TestChatCompletionsRequestToResponsesRequestInstructionsAndTools(t *testing assert.Equal(t, "function_call_output", gjson.GetBytes(got.Input, "3.type").String()) } +func TestChatCompletionsRequestToResponsesRequestPreservesQwenThinkingBudget(t *testing.T) { + tests := []struct { + name string + budget json.RawMessage + want int64 + }{ + {name: "positive budget", budget: json.RawMessage(`128`), want: 128}, + {name: "zero budget", budget: json.RawMessage(`0`), want: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := &dto.GeneralOpenAIRequest{ + Model: "qwen-plus", + EnableThinking: json.RawMessage(`true`), + ThinkingBudget: tt.budget, + Messages: []dto.Message{ + {Role: "user", Content: "hello"}, + }, + } + + got, err := ChatCompletionsRequestToResponsesRequest(req) + require.NoError(t, err) + assert.Equal(t, tt.budget, got.ThinkingBudget) + + encoded, err := kitutil.Marshal(got) + require.NoError(t, err) + + assert.True(t, gjson.GetBytes(encoded, "enable_thinking").Bool()) + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, tt.want, value.Int()) + }) + } +} + func TestChatCompletionsRequestToResponsesRequestRejectsMultipleChoices(t *testing.T) { _, err := ChatCompletionsRequestToResponsesRequest(&dto.GeneralOpenAIRequest{ Model: "gpt-test", diff --git a/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req.go b/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req.go index 263887a4..23b4bc70 100644 --- a/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req.go +++ b/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req.go @@ -73,6 +73,7 @@ func ResponsesRequestToChatCompletionsRequest(req *dto.OpenAIResponsesRequest) ( SafetyIdentifier: req.SafetyIdentifier, PromptCacheRetention: req.PromptCacheRetention, EnableThinking: req.EnableThinking, + ThinkingBudget: req.ThinkingBudget, } if req.Reasoning != nil { diff --git a/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req_test.go b/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req_test.go index a6f778f6..d27f2829 100644 --- a/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req_test.go +++ b/relaykit/relayconvert/internal/oai_responses/to_oai_chat_req_test.go @@ -1,6 +1,7 @@ package oairesponses import ( + "encoding/json" "testing" "github.com/QuantumNous/new-api/relaykit/dto" @@ -55,6 +56,38 @@ func TestResponsesRequestToChatCompletionsRequestInstructionsAndScalarInput(t *t assert.Equal(t, "abc", gjson.GetBytes(got.Metadata, "trace").String()) } +func TestResponsesRequestToChatCompletionsRequestPreservesQwenThinkingBudget(t *testing.T) { + tests := []struct { + name string + budget json.RawMessage + want int64 + }{ + {name: "positive budget", budget: json.RawMessage(`128`), want: 128}, + {name: "zero budget", budget: json.RawMessage(`0`), want: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := ResponsesRequestToChatCompletionsRequest(&dto.OpenAIResponsesRequest{ + Model: "qwen-plus", + Input: mustRawMessage(t, "hello"), + EnableThinking: json.RawMessage(`true`), + ThinkingBudget: tt.budget, + }) + require.NoError(t, err) + assert.Equal(t, tt.budget, got.ThinkingBudget) + + encoded, err := kitutil.Marshal(got) + require.NoError(t, err) + + assert.True(t, gjson.GetBytes(encoded, "enable_thinking").Bool()) + value := gjson.GetBytes(encoded, "thinking_budget") + assert.True(t, value.Exists()) + assert.Equal(t, tt.want, value.Int()) + }) + } +} + func TestResponsesRequestToChatCompletionsRequestMultimodalInput(t *testing.T) { got, err := ResponsesRequestToChatCompletionsRequest(&dto.OpenAIResponsesRequest{ Model: "gpt-test",