fix(relay): preserve presence/frequency penalty in Responses conversion (#6654)

This commit is contained in:
Uzuki Shion
2026-08-11 13:40:13 +08:00
committed by GitHub
parent 9c97e78ace
commit 253a74dd1b
8 changed files with 172 additions and 1 deletions
+5
View File
@@ -867,6 +867,11 @@ type OpenAIResponsesRequest struct {
Metadata json.RawMessage `json:"metadata,omitempty"`
Moderation json.RawMessage `json:"moderation,omitempty"`
ParallelToolCalls json.RawMessage `json:"parallel_tool_calls,omitempty"`
// FrequencyPenalty/PresencePenalty are not part of the official OpenAI
// Responses API; they are forwarded verbatim for OpenAI-compatible upstreams
// (e.g. vLLM) that accept them.
FrequencyPenalty json.RawMessage `json:"frequency_penalty,omitempty"`
PresencePenalty json.RawMessage `json:"presence_penalty,omitempty"`
PreviousResponseID string `json:"previous_response_id,omitempty"`
Reasoning *Reasoning `json:"reasoning,omitempty"`
// ServiceTier specifies upstream service level and may affect billing.
@@ -123,7 +123,9 @@ func TestOpenAIResponsesRequestPreserveExplicitZeroValues(t *testing.T) {
"max_output_tokens":0,
"max_tool_calls":0,
"stream":false,
"top_p":0
"top_p":0,
"frequency_penalty":0,
"presence_penalty":0
}`)
var req OpenAIResponsesRequest
@@ -137,6 +139,8 @@ func TestOpenAIResponsesRequestPreserveExplicitZeroValues(t *testing.T) {
require.True(t, gjson.GetBytes(encoded, "max_tool_calls").Exists())
require.True(t, gjson.GetBytes(encoded, "stream").Exists())
require.True(t, gjson.GetBytes(encoded, "top_p").Exists())
require.True(t, gjson.GetBytes(encoded, "frequency_penalty").Exists())
require.True(t, gjson.GetBytes(encoded, "presence_penalty").Exists())
}
func TestOpenAIResponsesRequestPreserveQwenThinkingBudget(t *testing.T) {
@@ -372,6 +372,14 @@ func ChatCompletionsRequestToResponsesRequest(req *dto.GeneralOpenAIRequest) (*d
topP = kitutil.GetPointer(lo.FromPtr(req.TopP))
}
var frequencyPenaltyRaw, presencePenaltyRaw json.RawMessage
if req.FrequencyPenalty != nil {
frequencyPenaltyRaw, _ = kitutil.Marshal(req.FrequencyPenalty)
}
if req.PresencePenalty != nil {
presencePenaltyRaw, _ = kitutil.Marshal(req.PresencePenalty)
}
out := &dto.OpenAIResponsesRequest{
Model: req.Model,
Input: inputRaw,
@@ -382,6 +390,8 @@ func ChatCompletionsRequestToResponsesRequest(req *dto.GeneralOpenAIRequest) (*d
ToolChoice: toolChoiceRaw,
Tools: toolsRaw,
TopP: topP,
FrequencyPenalty: frequencyPenaltyRaw,
PresencePenalty: presencePenaltyRaw,
User: req.User,
ParallelToolCalls: parallelToolCallsRaw,
Store: req.Store,
@@ -84,6 +84,51 @@ func TestChatCompletionsRequestToResponsesRequestRejectsMultipleChoices(t *testi
assert.Contains(t, err.Error(), "n>1")
}
func TestChatCompletionsRequestToResponsesRequestPreservesPenalties(t *testing.T) {
tests := []struct {
name string
frequency *float64
frequencyWant json.RawMessage
presence *float64
presenceWant json.RawMessage
}{
{
name: "positive values",
frequency: lo.ToPtr(0.5),
frequencyWant: json.RawMessage(`0.5`),
presence: lo.ToPtr(1.5),
presenceWant: json.RawMessage(`1.5`),
},
{
name: "explicit zero values",
frequency: lo.ToPtr(0.0),
frequencyWant: json.RawMessage(`0`),
presence: lo.ToPtr(0.0),
presenceWant: json.RawMessage(`0`),
},
{
name: "unset stays nil",
frequency: nil,
presence: nil,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := ChatCompletionsRequestToResponsesRequest(&dto.GeneralOpenAIRequest{
Model: "gpt-test",
Messages: []dto.Message{{Role: "user", Content: "hello"}},
FrequencyPenalty: tt.frequency,
PresencePenalty: tt.presence,
})
require.NoError(t, err)
assert.Equal(t, tt.frequencyWant, got.FrequencyPenalty)
assert.Equal(t, tt.presenceWant, got.PresencePenalty)
})
}
}
func assistantMessageWithTool(content string, id string, name string, args string) dto.Message {
msg := dto.Message{Role: "assistant", Content: content}
msg.SetToolCalls([]dto.ToolCallRequest{
@@ -76,6 +76,15 @@ func ResponsesRequestToChatCompletionsRequest(req *dto.OpenAIResponsesRequest) (
ThinkingBudget: req.ThinkingBudget,
}
out.FrequencyPenalty, err = responsesRawFloat(req.FrequencyPenalty)
if err != nil {
return nil, fmt.Errorf("invalid frequency_penalty: %w", err)
}
out.PresencePenalty, err = responsesRawFloat(req.PresencePenalty)
if err != nil {
return nil, fmt.Errorf("invalid presence_penalty: %w", err)
}
if req.Reasoning != nil {
out.ReasoningEffort = req.Reasoning.Effort
}
@@ -527,6 +536,17 @@ func responseToolOutputToChatContent(value any) any {
}
}
func responsesRawFloat(raw json.RawMessage) (*float64, error) {
if !rawJSONPresent(raw) {
return nil, nil
}
var value float64
if err := kitutil.Unmarshal(raw, &value); err != nil {
return nil, err
}
return &value, nil
}
func responsesJSONString(raw json.RawMessage) (string, error) {
if kitutil.GetJsonType(raw) != "string" {
return string(raw), nil
@@ -295,6 +295,61 @@ func TestResponsesRequestToChatCompletionsRequestRejectsStatefulFields(t *testin
}
}
func TestResponsesRequestToChatCompletionsRequestPreservesPenalties(t *testing.T) {
tests := []struct {
name string
frequencyRaw json.RawMessage
frequencyWant *float64
presenceRaw json.RawMessage
presenceWant *float64
}{
{
name: "positive values",
frequencyRaw: json.RawMessage(`0.5`),
frequencyWant: lo.ToPtr(0.5),
presenceRaw: json.RawMessage(`1.5`),
presenceWant: lo.ToPtr(1.5),
},
{
name: "explicit zero values",
frequencyRaw: json.RawMessage(`0.0`),
frequencyWant: lo.ToPtr(0.0),
presenceRaw: json.RawMessage(`0.0`),
presenceWant: lo.ToPtr(0.0),
},
{
name: "unset stays nil",
frequencyRaw: nil,
presenceRaw: nil,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := ResponsesRequestToChatCompletionsRequest(&dto.OpenAIResponsesRequest{
Model: "gpt-test",
Input: mustRawMessage(t, "hello"),
FrequencyPenalty: tt.frequencyRaw,
PresencePenalty: tt.presenceRaw,
})
require.NoError(t, err)
assert.Equal(t, tt.frequencyWant, got.FrequencyPenalty)
assert.Equal(t, tt.presenceWant, got.PresencePenalty)
})
}
}
func TestResponsesRequestToChatCompletionsRequestRejectsMalformedPenalty(t *testing.T) {
_, err := ResponsesRequestToChatCompletionsRequest(&dto.OpenAIResponsesRequest{
Model: "gpt-test",
Input: mustRawMessage(t, "hello"),
FrequencyPenalty: json.RawMessage(`"not-a-number"`),
})
require.Error(t, err)
assert.Contains(t, err.Error(), "frequency_penalty")
}
func mustRawMessage(t *testing.T, value any) []byte {
t.Helper()
raw, err := kitutil.Marshal(value)