fix(relay): preserve presence/frequency penalty in Responses conversion (#6654)
This commit is contained in:
@@ -104,6 +104,8 @@ func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommo
|
||||
// rm max_output_tokens
|
||||
request.MaxOutputTokens = nil
|
||||
request.Temperature = nil
|
||||
request.FrequencyPenalty = nil
|
||||
request.PresencePenalty = nil
|
||||
return request, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
package codex
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/QuantumNous/new-api/constant"
|
||||
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
||||
relayconstant "github.com/QuantumNous/new-api/relay/constant"
|
||||
"github.com/QuantumNous/new-api/relaykit/dto"
|
||||
"github.com/samber/lo"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
@@ -24,3 +27,30 @@ func TestGetRequestURLAlphaSearch(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "https://chatgpt.com/backend-api/codex/alpha/search", url)
|
||||
}
|
||||
|
||||
// The Codex backend rejects these fields, so the adaptor clears them rather
|
||||
// than forwarding what the client sent.
|
||||
func TestConvertOpenAIResponsesRequestDropsPenalties(t *testing.T) {
|
||||
adaptor := &Adaptor{}
|
||||
info := &relaycommon.RelayInfo{
|
||||
ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeCodex},
|
||||
RelayMode: relayconstant.RelayModeResponses,
|
||||
}
|
||||
|
||||
converted, err := adaptor.ConvertOpenAIResponsesRequest(nil, info, dto.OpenAIResponsesRequest{
|
||||
Model: "gpt-5-codex",
|
||||
Input: json.RawMessage(`"hello"`),
|
||||
MaxOutputTokens: lo.ToPtr(uint(128)),
|
||||
Temperature: lo.ToPtr(1.0),
|
||||
FrequencyPenalty: json.RawMessage(`1.5`),
|
||||
PresencePenalty: json.RawMessage(`1.5`),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
request, ok := converted.(dto.OpenAIResponsesRequest)
|
||||
require.True(t, ok)
|
||||
assert.Nil(t, request.MaxOutputTokens)
|
||||
assert.Nil(t, request.Temperature)
|
||||
assert.Nil(t, request.FrequencyPenalty)
|
||||
assert.Nil(t, request.PresencePenalty)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user