Skip to content
Merged
2 changes: 1 addition & 1 deletion relay/channel/ali/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
}
Expand Down
128 changes: 128 additions & 0 deletions relay/channel/ali/adaptor_test.go
Original file line number Diff line number Diff line change
@@ -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())
}
10 changes: 9 additions & 1 deletion relay/channel/ali/text.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
1 change: 0 additions & 1 deletion relay/channel/baidu/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/cloudflare/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/cohere/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
2 changes: 0 additions & 2 deletions relay/channel/dify/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -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 {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/jina/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/mistral/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/mokaai/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/palm/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/tencent/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/xunfei/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
1 change: 0 additions & 1 deletion relay/channel/zhipu/adaptor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
26 changes: 26 additions & 0 deletions relaykit/dto/openai_request.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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") {
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading