Skip to content
Closed
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
4 changes: 2 additions & 2 deletions dto/gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -160,7 +160,7 @@ func (r *GeminiChatRequest) SetTools(tools []GeminiChatTool) {
}

type GeminiThinkingConfig struct {
IncludeThoughts bool `json:"includeThoughts,omitempty"`
IncludeThoughts *bool `json:"includeThoughts,omitempty"`
ThinkingBudget *int `json:"thinkingBudget,omitempty"`
// TODO Conflict with thinkingbudget.
ThinkingLevel string `json:"thinkingLevel,omitempty"`
Expand All @@ -183,7 +183,7 @@ func (c *GeminiThinkingConfig) UnmarshalJSON(data []byte) error {
*c = GeminiThinkingConfig(aux.Alias)

if aux.IncludeThoughtsSnake != nil {
c.IncludeThoughts = *aux.IncludeThoughtsSnake
c.IncludeThoughts = aux.IncludeThoughtsSnake
}

if aux.ThinkingBudgetSnake != nil {
Expand Down
51 changes: 51 additions & 0 deletions dto/gemini_generation_config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -87,3 +87,54 @@ func TestGeminiChatGenerationConfigPreservesExplicitZeroValuesSnakeCase(t *testi
assert.Equal(t, float64(0), generationConfig["seed"])
assert.Equal(t, false, generationConfig["responseLogprobs"])
}

func TestGeminiThinkingConfigPreservesExplicitFalseValues(t *testing.T) {
tests := []struct {
name string
raw []byte
}{
{
name: "camel case",
raw: []byte(`{
"contents":[{"role":"user","parts":[{"text":"hello"}]}],
"generationConfig":{
"thinkingConfig":{
"includeThoughts":false
}
}
}`),
},
{
name: "snake case",
raw: []byte(`{
"contents":[{"role":"user","parts":[{"text":"hello"}]}],
"generationConfig":{
"thinking_config":{
"include_thoughts":false
}
}
}`),
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var req GeminiChatRequest
require.NoError(t, common.Unmarshal(tt.raw, &req))

encoded, err := common.Marshal(req)
require.NoError(t, err)

var out map[string]any
require.NoError(t, common.Unmarshal(encoded, &out))

generationConfig, ok := out["generationConfig"].(map[string]any)
require.True(t, ok)
thinkingConfig, ok := generationConfig["thinkingConfig"].(map[string]any)
require.True(t, ok)

assert.Contains(t, thinkingConfig, "includeThoughts")
assert.Equal(t, false, thinkingConfig["includeThoughts"])
})
}
}
18 changes: 10 additions & 8 deletions relay/channel/gemini/relay-gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -145,7 +145,7 @@ func ThinkingAdaptor(geminiRequest *dto.GeminiChatRequest, info *relaycommon.Rel
clampedBudget := clampThinkingBudget(modelName, budgetTokens)
geminiRequest.GenerationConfig.ThinkingConfig = &dto.GeminiThinkingConfig{
ThinkingBudget: common.GetPointer(clampedBudget),
IncludeThoughts: true,
IncludeThoughts: common.GetPointer(true),
}
}
}
Expand All @@ -164,11 +164,11 @@ func ThinkingAdaptor(geminiRequest *dto.GeminiChatRequest, info *relaycommon.Rel

if isUnsupported {
geminiRequest.GenerationConfig.ThinkingConfig = &dto.GeminiThinkingConfig{
IncludeThoughts: true,
IncludeThoughts: common.GetPointer(true),
}
} else {
geminiRequest.GenerationConfig.ThinkingConfig = &dto.GeminiThinkingConfig{
IncludeThoughts: true,
IncludeThoughts: common.GetPointer(true),
}
if geminiRequest.GenerationConfig.MaxOutputTokens != nil && *geminiRequest.GenerationConfig.MaxOutputTokens > 0 {
budgetTokens := model_setting.GetGeminiSettings().ThinkingAdapterBudgetTokensPercentage * float64(*geminiRequest.GenerationConfig.MaxOutputTokens)
Expand All @@ -189,7 +189,7 @@ func ThinkingAdaptor(geminiRequest *dto.GeminiChatRequest, info *relaycommon.Rel
}
} else if _, level, ok := reasoning.TrimEffortSuffix(info.UpstreamModelName); ok && level != "" {
geminiRequest.GenerationConfig.ThinkingConfig = &dto.GeminiThinkingConfig{
IncludeThoughts: true,
IncludeThoughts: common.GetPointer(true),
ThinkingLevel: level,
}
info.ReasoningEffort = level
Expand Down Expand Up @@ -271,10 +271,10 @@ func CovertOpenAI2Gemini(c *gin.Context, textRequest dto.GeneralOpenAIRequest, i
tempThinkingConfig.ThinkingBudget = common.GetPointer(budgetInt)
if budgetInt > 0 {
// 有正数预算
tempThinkingConfig.IncludeThoughts = true
tempThinkingConfig.IncludeThoughts = common.GetPointer(true)
} else {
// 存在但为0或负数,禁用思考
tempThinkingConfig.IncludeThoughts = false
tempThinkingConfig.IncludeThoughts = common.GetPointer(false)
}
hasThinkingConfig = true
default:
Expand All @@ -284,7 +284,7 @@ func CovertOpenAI2Gemini(c *gin.Context, textRequest dto.GeneralOpenAIRequest, i

if includeThoughts, exists := thinkingConfig["include_thoughts"]; exists {
if v, ok := includeThoughts.(bool); ok {
tempThinkingConfig.IncludeThoughts = v
tempThinkingConfig.IncludeThoughts = common.GetPointer(v)
hasThinkingConfig = true
} else {
return nil, errors.New("extra_body.google.thinking_config.include_thoughts must be a boolean")
Expand All @@ -308,7 +308,9 @@ func CovertOpenAI2Gemini(c *gin.Context, textRequest dto.GeneralOpenAIRequest, i
if tempThinkingConfig.ThinkingBudget != nil {
geminiRequest.GenerationConfig.ThinkingConfig.ThinkingBudget = tempThinkingConfig.ThinkingBudget
}
geminiRequest.GenerationConfig.ThinkingConfig.IncludeThoughts = tempThinkingConfig.IncludeThoughts
if tempThinkingConfig.IncludeThoughts != nil {
geminiRequest.GenerationConfig.ThinkingConfig.IncludeThoughts = tempThinkingConfig.IncludeThoughts
}
if tempThinkingConfig.ThinkingLevel != "" {
geminiRequest.GenerationConfig.ThinkingConfig.ThinkingLevel = tempThinkingConfig.ThinkingLevel
}
Expand Down
35 changes: 35 additions & 0 deletions relay/channel/gemini/relay_gemini_usage_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -331,3 +331,38 @@ func TestGeminiTextGenerationHandlerUsesEstimatedPromptTokensWhenUsagePromptMiss
require.Equal(t, 100, usage.CompletionTokens)
require.Equal(t, 110, usage.TotalTokens)
}

func TestCovertOpenAI2GeminiPreservesExplicitFalseIncludeThoughts(t *testing.T) {
t.Parallel()

gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)

request := dto.GeneralOpenAIRequest{
Model: "gemini-2.5-flash",
Messages: []dto.Message{
{Role: "user", Content: "hello"},
},
ExtraBody: []byte(`{
"google": {
"thinking_config": {
"include_thoughts": false
}
}
}`),
}
info := &relaycommon.RelayInfo{
OriginModelName: "gemini-2.5-flash",
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeGemini,
UpstreamModelName: "gemini-2.5-flash",
},
}

converted, err := CovertOpenAI2Gemini(c, request, info)
require.NoError(t, err)
require.NotNil(t, converted.GenerationConfig.ThinkingConfig)
require.NotNil(t, converted.GenerationConfig.ThinkingConfig.IncludeThoughts)
require.False(t, *converted.GenerationConfig.ThinkingConfig.IncludeThoughts)
}