mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-14 00:01:53 +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 {
|
switch info.RelayMode {
|
||||||
default:
|
default:
|
||||||
aliReq := requestOpenAI2Ali(*request)
|
aliReq := requestOpenAI2Ali(*request, info.UpstreamModelName)
|
||||||
return aliReq, nil
|
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"
|
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)
|
topP := lo.FromPtrOr(request.TopP, 0)
|
||||||
if topP >= 1 {
|
if topP >= 1 {
|
||||||
request.TopP = lo.ToPtr(0.999)
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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 {
|
} else {
|
||||||
return difyHandler(c, info, resp)
|
return difyHandler(c, info, resp)
|
||||||
}
|
}
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) GetModelList() []string {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
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) {
|
func (a *Adaptor) ConvertClaudeRequest(*gin.Context, *relaycommon.RelayInfo, *dto.ClaudeRequest) (any, error) {
|
||||||
//TODO implement me
|
//TODO implement me
|
||||||
panic("implement me")
|
panic("implement me")
|
||||||
return nil, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) {
|
||||||
|
|||||||
@@ -89,6 +89,7 @@ type GeneralOpenAIRequest struct {
|
|||||||
// Ali Qwen Params
|
// Ali Qwen Params
|
||||||
VlHighResolutionImages json.RawMessage `json:"vl_high_resolution_images,omitempty"`
|
VlHighResolutionImages json.RawMessage `json:"vl_high_resolution_images,omitempty"`
|
||||||
EnableThinking json.RawMessage `json:"enable_thinking,omitempty"`
|
EnableThinking json.RawMessage `json:"enable_thinking,omitempty"`
|
||||||
|
ThinkingBudget json.RawMessage `json:"thinking_budget,omitempty"`
|
||||||
ChatTemplateKwargs json.RawMessage `json:"chat_template_kwargs,omitempty"`
|
ChatTemplateKwargs json.RawMessage `json:"chat_template_kwargs,omitempty"`
|
||||||
EnableSearch json.RawMessage `json:"enable_search,omitempty"`
|
EnableSearch json.RawMessage `json:"enable_search,omitempty"`
|
||||||
// ollama Params
|
// ollama Params
|
||||||
@@ -107,6 +108,14 @@ type GeneralOpenAIRequest struct {
|
|||||||
ReasoningSplit json.RawMessage `json:"reasoning_split,omitempty"`
|
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 {
|
func (r *GeneralOpenAIRequest) GetTokenCountMeta() *types.TokenCountMeta {
|
||||||
var tokenCountMeta types.TokenCountMeta
|
var tokenCountMeta types.TokenCountMeta
|
||||||
var texts = make([]string, 0)
|
var texts = make([]string, 0)
|
||||||
@@ -222,6 +231,14 @@ func IsOpenAIGPT5Model(modelName string) bool {
|
|||||||
return strings.HasPrefix(modelName, "gpt-5")
|
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 {
|
func (r *GeneralOpenAIRequest) GetSystemRoleName() string {
|
||||||
if IsOpenAIReasoningOModel(r.Model) {
|
if IsOpenAIReasoningOModel(r.Model) {
|
||||||
if !strings.HasPrefix(r.Model, "o1-mini") && !strings.HasPrefix(r.Model, "o1-preview") {
|
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"`
|
ClientMetadata json.RawMessage `json:"client_metadata,omitempty"`
|
||||||
// qwen
|
// qwen
|
||||||
EnableThinking json.RawMessage `json:"enable_thinking,omitempty"`
|
EnableThinking json.RawMessage `json:"enable_thinking,omitempty"`
|
||||||
|
ThinkingBudget json.RawMessage `json:"thinking_budget,omitempty"`
|
||||||
// perplexity
|
// perplexity
|
||||||
Preset json.RawMessage `json:"preset,omitempty"`
|
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 {
|
func (r *OpenAIResponsesRequest) GetTokenCountMeta() *types.TokenCountMeta {
|
||||||
var fileMeta = make([]*types.FileMeta, 0)
|
var fileMeta = make([]*types.FileMeta, 0)
|
||||||
var texts = make([]string, 0)
|
var texts = make([]string, 0)
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
package dto
|
package dto
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil"
|
kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"github.com/tidwall/gjson"
|
"github.com/tidwall/gjson"
|
||||||
)
|
)
|
||||||
@@ -50,6 +52,71 @@ func TestGeneralOpenAIRequestPreserveExplicitZeroValues(t *testing.T) {
|
|||||||
require.True(t, gjson.GetBytes(encoded, "return_related_questions").Exists())
|
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) {
|
func TestOpenAIResponsesRequestPreserveExplicitZeroValues(t *testing.T) {
|
||||||
raw := []byte(`{
|
raw := []byte(`{
|
||||||
"model":"gpt-4.1",
|
"model":"gpt-4.1",
|
||||||
@@ -72,6 +139,46 @@ func TestOpenAIResponsesRequestPreserveExplicitZeroValues(t *testing.T) {
|
|||||||
require.True(t, gjson.GetBytes(encoded, "top_p").Exists())
|
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) {
|
func TestGeneralOpenAIRequestGetSystemRoleName(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -386,6 +386,8 @@ func ChatCompletionsRequestToResponsesRequest(req *dto.GeneralOpenAIRequest) (*d
|
|||||||
ParallelToolCalls: parallelToolCallsRaw,
|
ParallelToolCalls: parallelToolCallsRaw,
|
||||||
Store: req.Store,
|
Store: req.Store,
|
||||||
Metadata: req.Metadata,
|
Metadata: req.Metadata,
|
||||||
|
EnableThinking: req.EnableThinking,
|
||||||
|
ThinkingBudget: req.ThinkingBudget,
|
||||||
}
|
}
|
||||||
if req.MaxTokens != nil || req.MaxCompletionTokens != nil {
|
if req.MaxTokens != nil || req.MaxCompletionTokens != nil {
|
||||||
out.MaxOutputTokens = lo.ToPtr(maxOutputTokens)
|
out.MaxOutputTokens = lo.ToPtr(maxOutputTokens)
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
package oaichat
|
package oaichat
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/QuantumNous/new-api/relaykit/dto"
|
"github.com/QuantumNous/new-api/relaykit/dto"
|
||||||
|
kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil"
|
||||||
"github.com/samber/lo"
|
"github.com/samber/lo"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"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())
|
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) {
|
func TestChatCompletionsRequestToResponsesRequestRejectsMultipleChoices(t *testing.T) {
|
||||||
_, err := ChatCompletionsRequestToResponsesRequest(&dto.GeneralOpenAIRequest{
|
_, err := ChatCompletionsRequestToResponsesRequest(&dto.GeneralOpenAIRequest{
|
||||||
Model: "gpt-test",
|
Model: "gpt-test",
|
||||||
|
|||||||
@@ -73,6 +73,7 @@ func ResponsesRequestToChatCompletionsRequest(req *dto.OpenAIResponsesRequest) (
|
|||||||
SafetyIdentifier: req.SafetyIdentifier,
|
SafetyIdentifier: req.SafetyIdentifier,
|
||||||
PromptCacheRetention: req.PromptCacheRetention,
|
PromptCacheRetention: req.PromptCacheRetention,
|
||||||
EnableThinking: req.EnableThinking,
|
EnableThinking: req.EnableThinking,
|
||||||
|
ThinkingBudget: req.ThinkingBudget,
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Reasoning != nil {
|
if req.Reasoning != nil {
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package oairesponses
|
package oairesponses
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/QuantumNous/new-api/relaykit/dto"
|
"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())
|
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) {
|
func TestResponsesRequestToChatCompletionsRequestMultimodalInput(t *testing.T) {
|
||||||
got, err := ResponsesRequestToChatCompletionsRequest(&dto.OpenAIResponsesRequest{
|
got, err := ResponsesRequestToChatCompletionsRequest(&dto.OpenAIResponsesRequest{
|
||||||
Model: "gpt-test",
|
Model: "gpt-test",
|
||||||
|
|||||||
Reference in New Issue
Block a user