mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-12 15:21:09 +00:00
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
This commit is contained in:
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user